|
| 1 | +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. |
| 2 | +// SPDX-License-Identifier: MIT |
| 3 | + |
| 4 | +#pragma once |
| 5 | + |
| 6 | +template <typename ProblemType> |
| 7 | +bool run_gemm(const ProblemType& problem_size, const ExecutionConfig& config) |
| 8 | +{ |
| 9 | + using namespace ck::literals; |
| 10 | + |
| 11 | + auto M = problem_size.M; |
| 12 | + auto N = problem_size.N; |
| 13 | + auto K = problem_size.K; |
| 14 | + auto StrideA = problem_size.StrideA; |
| 15 | + auto StrideB = problem_size.StrideB; |
| 16 | + auto StrideC = problem_size.StrideC; |
| 17 | + auto KBatch = problem_size.KBatch; |
| 18 | + |
| 19 | + auto f_host_tensor_descriptor = |
| 20 | + [](std::size_t row, std::size_t col, std::size_t stride, auto layout) { |
| 21 | + if constexpr(std::is_same_v<decltype(layout), ck::tensor_layout::gemm::RowMajor>) |
| 22 | + { |
| 23 | + return HostTensorDescriptor({row, col}, {stride, 1_uz}); |
| 24 | + } |
| 25 | + else |
| 26 | + { |
| 27 | + return HostTensorDescriptor({row, col}, {1_uz, stride}); |
| 28 | + } |
| 29 | + }; |
| 30 | + |
| 31 | + auto f_get_default_stride = |
| 32 | + [](std::size_t row, std::size_t col, ck::index_t stride, auto layout) { |
| 33 | + if(stride == -1) |
| 34 | + { |
| 35 | + // give a chance if stride is -1, return a default packed stride |
| 36 | + if constexpr(std::is_same_v<decltype(layout), ck::tensor_layout::gemm::RowMajor>) |
| 37 | + { |
| 38 | + return static_cast<std::size_t>(col); |
| 39 | + } |
| 40 | + else |
| 41 | + { |
| 42 | + return static_cast<std::size_t>(row); |
| 43 | + } |
| 44 | + } |
| 45 | + else |
| 46 | + return static_cast<std::size_t>(stride); |
| 47 | + }; |
| 48 | + |
| 49 | + StrideA = f_get_default_stride(M, K, StrideA, ALayout{}); |
| 50 | + StrideB = f_get_default_stride(K, N, StrideB, BLayout{}); |
| 51 | + StrideC = f_get_default_stride(M, N, StrideC, CLayout{}); |
| 52 | + |
| 53 | + Tensor<ADataType> a_m_k(f_host_tensor_descriptor(M, K, StrideA, ALayout{})); |
| 54 | + Tensor<BDataType> b_k_n(f_host_tensor_descriptor(K, N, StrideB, BLayout{})); |
| 55 | + Tensor<BDataType> b_k_n_preshuffled(f_host_tensor_descriptor(K, N, StrideB, BLayout{})); |
| 56 | + |
| 57 | + switch(config.init_method) |
| 58 | + { |
| 59 | + case 0: break; |
| 60 | + case 1: |
| 61 | + a_m_k.GenerateTensorValue(GeneratorTensor_2<ADataType>{-2, 2}); |
| 62 | + b_k_n.GenerateTensorValue(GeneratorTensor_2<BDataType>{0, 2}); |
| 63 | + break; |
| 64 | + case 2: |
| 65 | + a_m_k.GenerateTensorValue(GeneratorTensor_1<ADataType>{}); |
| 66 | + b_k_n.GenerateTensorValue(GeneratorTensor_1<BDataType>{}); |
| 67 | + break; |
| 68 | + default: |
| 69 | + a_m_k.GenerateTensorValue(GeneratorTensor_3<ADataType>{0.0, 1.0}); |
| 70 | + b_k_n.GenerateTensorValue(GeneratorTensor_3<BDataType>{-0.5, 0.5}); |
| 71 | + } |
| 72 | + |
| 73 | + Tensor<CDataType> c_m_n_host_result(f_host_tensor_descriptor(M, N, StrideC, CLayout{})); |
| 74 | + Tensor<CDataType> c_m_n_device_result(f_host_tensor_descriptor(M, N, StrideC, CLayout{})); |
| 75 | + |
| 76 | + std::cout << "a_m_k: " << a_m_k.mDesc << std::endl; |
| 77 | + std::cout << "b_k_n: " << b_k_n.mDesc << std::endl; |
| 78 | + std::cout << "b_k_n_preshuffled: " << b_k_n_preshuffled.mDesc << std::endl; |
| 79 | + std::cout << "c_m_n: " << c_m_n_host_result.mDesc << std::endl; |
| 80 | + |
| 81 | + DeviceMem a_m_k_device_buf(sizeof(ADataType) * a_m_k.mDesc.GetElementSpaceSize()); |
| 82 | + DeviceMem b_k_n_device_buf(sizeof(BDataType) * b_k_n.mDesc.GetElementSpaceSize()); |
| 83 | + DeviceMem c_m_n_device_buf(sizeof(CDataType) * c_m_n_device_result.mDesc.GetElementSpaceSize()); |
| 84 | + |
| 85 | + // do GEMM |
| 86 | + auto device_op = DeviceOpInstance{}; |
| 87 | + |
| 88 | + // weight pre-shuffle |
| 89 | + int NPerWmma = device_op.GetPreShuffleParameters(); |
| 90 | + int KLane = ck::get_warp_size() / NPerWmma; |
| 91 | + |
| 92 | + int K0 = K / (KLane * KPack); |
| 93 | + // K -> K0 KLane KPack |
| 94 | + // N -> N0 NPerWmma |
| 95 | + // N, K -> N0 K0 KLane NPerWmma KPack |
| 96 | + int tempk; |
| 97 | + for(int n = 0; n < N; ++n) |
| 98 | + { |
| 99 | + for(int k = 0; k < K; ++k) |
| 100 | + { |
| 101 | + int n0 = n / NPerWmma; |
| 102 | + int n1 = n % NPerWmma; |
| 103 | + |
| 104 | + int k0 = k / (KLane * KPack); |
| 105 | + tempk = k % (KLane * KPack); |
| 106 | + int k1 = tempk / KPack; |
| 107 | + int k2 = tempk % KPack; |
| 108 | + |
| 109 | + int outputIndex = n0 * KPack * NPerWmma * KLane * K0 + k0 * KPack * NPerWmma * KLane + |
| 110 | + k1 * KPack * NPerWmma + n1 * KPack + k2; |
| 111 | + |
| 112 | + b_k_n_preshuffled(outputIndex) = b_k_n(n * K + k); |
| 113 | + } |
| 114 | + } |
| 115 | + |
| 116 | + a_m_k_device_buf.ToDevice(a_m_k.mData.data()); |
| 117 | + b_k_n_device_buf.ToDevice(b_k_n_preshuffled.mData.data()); |
| 118 | + c_m_n_device_buf.ToDevice(c_m_n_device_result.mData.data()); |
| 119 | + |
| 120 | + auto a_element_op = AElementOp{}; |
| 121 | + auto b_element_op = BElementOp{}; |
| 122 | + auto c_element_op = CElementOp{}; |
| 123 | + |
| 124 | + auto invoker = device_op.MakeInvoker(); |
| 125 | + |
| 126 | + auto argument = |
| 127 | + device_op.MakeArgument(static_cast<ADataType*>(a_m_k_device_buf.GetDeviceBuffer()), |
| 128 | + static_cast<BDataType*>(b_k_n_device_buf.GetDeviceBuffer()), |
| 129 | + static_cast<CDataType*>(c_m_n_device_buf.GetDeviceBuffer()), |
| 130 | + M, |
| 131 | + N, |
| 132 | + K, |
| 133 | + StrideA, |
| 134 | + StrideB, |
| 135 | + StrideC, |
| 136 | + KBatch, |
| 137 | + a_element_op, |
| 138 | + b_element_op, |
| 139 | + c_element_op); |
| 140 | + |
| 141 | + if(!device_op.IsSupportedArgument(argument)) |
| 142 | + { |
| 143 | + std::cerr << device_op.GetTypeString() << " does not support this problem" << std::endl; |
| 144 | + |
| 145 | + return true; |
| 146 | + } |
| 147 | + |
| 148 | + float ave_time = |
| 149 | + invoker.Run(argument, StreamConfig{nullptr, config.time_kernel, 0, 50, 50, false, 1}); |
| 150 | + |
| 151 | + bool pass = true; |
| 152 | + if(config.do_verification) |
| 153 | + { |
| 154 | + using ReferenceGemmInstance = ck::tensor_operation::host::ReferenceGemm<ADataType, |
| 155 | + BDataType, |
| 156 | + CDataType, |
| 157 | + AccDataType, |
| 158 | + PassThrough, |
| 159 | + PassThrough, |
| 160 | + PassThrough>; |
| 161 | + |
| 162 | + auto ref_gemm = ReferenceGemmInstance{}; |
| 163 | + auto ref_invoker = ref_gemm.MakeInvoker(); |
| 164 | + |
| 165 | + auto ref_argument = ref_gemm.MakeArgument( |
| 166 | + a_m_k, b_k_n, c_m_n_host_result, PassThrough{}, PassThrough{}, PassThrough{}); |
| 167 | + |
| 168 | + ref_invoker.Run(ref_argument); |
| 169 | + |
| 170 | + invoker.Run(argument, StreamConfig{nullptr, false, 0}); |
| 171 | + c_m_n_device_buf.FromDevice(c_m_n_device_result.mData.data()); |
| 172 | + |
| 173 | + pass &= ck::utils::check_err(c_m_n_device_result, |
| 174 | + c_m_n_host_result, |
| 175 | + "Error: Incorrect results!", |
| 176 | + get_rtol<CDataType>(), |
| 177 | + get_atol<CDataType>()); |
| 178 | + } |
| 179 | + |
| 180 | + if(config.time_kernel) |
| 181 | + { |
| 182 | + ave_time = |
| 183 | + invoker.Run(argument, StreamConfig{nullptr, config.time_kernel, 0, 20, 50, true, 50}); |
| 184 | + |
| 185 | + std::size_t flop = 2_uz * M * N * K; |
| 186 | + std::size_t num_btype = |
| 187 | + sizeof(ADataType) * M * K + sizeof(BDataType) * K * N + sizeof(CDataType) * M * N; |
| 188 | + |
| 189 | + float tflops = static_cast<float>(flop) / 1.E9 / ave_time; |
| 190 | + |
| 191 | + float gb_per_sec = num_btype / 1.E6 / ave_time; |
| 192 | + |
| 193 | + std::cout << "Perf: " << ave_time << " ms, " << tflops << " TFlops, " << gb_per_sec |
| 194 | + << " GB/s, " << device_op.GetTypeString() << std::endl; |
| 195 | + } |
| 196 | + |
| 197 | + return pass; |
| 198 | +} |
| 199 | + |
| 200 | +bool run_gemm_splitk_example(int argc, char* argv[]) |
| 201 | +{ |
| 202 | + ProblemSizeSplitK problem_size{3840, 4096, 4096, 4096, 4096, 4096, 1}; |
| 203 | + ExecutionConfig config; |
| 204 | + |
| 205 | + return parse_cmd_args(argc, argv, problem_size, config) && run_gemm(problem_size, config); |
| 206 | +} |
0 commit comments