已合并
refactor(matmul_recipes): rename QuantMatmulTilingData to QuantMatmulMxTilingData #393
tangweiwei2创建于 21 天前
refactor(matmul_recipes): rename QuantMatmulTilingData to QuantMatmulMxTilingData #393
已合并
共 14 个文件变更+38-38
| @@ -85,7 +85,7 @@ inline uint64_t ParsePositiveUint64(const char* arg, const char* name) | |||
| 85 | 85 | ||
| 86 | inline void CheckUint32Shape(uint64_t value, const char* name) | 86 | inline void CheckUint32Shape(uint64_t value, const char* name) |
| 87 | { | 87 | { |
| 88 | - // QuantMatmulTilingData serializes public shape fields as uint32_t. | 88 | + // QuantMatmulMxTilingData serializes public shape fields as uint32_t. |
| 89 | constexpr uint64_t uint32Max = static_cast<uint64_t>(std::numeric_limits<uint32_t>::max()); | 89 | constexpr uint64_t uint32Max = static_cast<uint64_t>(std::numeric_limits<uint32_t>::max()); |
| 90 | if (value > uint32Max) { | 90 | if (value > uint32Max) { |
| 91 | throw std::invalid_argument(std::string("ERROR: ") + name + " must not exceed UINT32_MAX"); | 91 | throw std::invalid_argument(std::string("ERROR: ") + name + " must not exceed UINT32_MAX"); |
| @@ -32,12 +32,12 @@ | |||
| 32 | #include "host_utils/io_utils.h" | 32 | #include "host_utils/io_utils.h" |
| 33 | #include "kernel/quant_matmul_mx_kernel_a_full_load.h" | 33 | #include "kernel/quant_matmul_mx_kernel_a_full_load.h" |
| 34 | #include "tiling/quant_matmul_mx_tiling_a_full_load.h" | 34 | #include "tiling/quant_matmul_mx_tiling_a_full_load.h" |
| 35 | -#include "tiling/quant_matmul_tiling_data.h" | 35 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 36 | 36 | ||
| 37 | template <bool TransA, bool TransB> | 37 | template <bool TransA, bool TransB> |
| 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4AFullLoadTransKernel( | 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4AFullLoadTransKernel( |
| 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 40 | - const QuantMatmulTilingData quantMatmulTilingData) | 40 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 41 | { | 41 | { |
| 42 | using TypeA = fp4x2_e2m1_t; | 42 | using TypeA = fp4x2_e2m1_t; |
| 43 | using TypeB = fp4x2_e2m1_t; | 43 | using TypeB = fp4x2_e2m1_t; |
| @@ -114,7 +114,7 @@ int main(int argc, char* argv[]) | |||
| 114 | constexpr int32_t deviceId = 0; | 114 | constexpr int32_t deviceId = 0; |
| 115 | 115 | ||
| 116 | try { | 116 | try { |
| 117 | - QuantMatmulTilingData tilingData; | 117 | + QuantMatmulMxTilingData tilingData; |
| 118 | QuantMatmulTilingAFullLoad<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; | 118 | QuantMatmulTilingAFullLoad<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; |
| 119 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 119 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 120 | 120 | ||
| @@ -32,12 +32,12 @@ | |||
| 32 | #include "host_utils/io_utils.h" | 32 | #include "host_utils/io_utils.h" |
| 33 | #include "kernel/quant_matmul_mx_kernel_swat.h" | 33 | #include "kernel/quant_matmul_mx_kernel_swat.h" |
| 34 | #include "tiling/quant_matmul_mx_tiling_swat.h" | 34 | #include "tiling/quant_matmul_mx_tiling_swat.h" |
| 35 | -#include "tiling/quant_matmul_tiling_data.h" | 35 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 36 | 36 | ||
| 37 | template <bool TransA, bool TransB> | 37 | template <bool TransA, bool TransB> |
| 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4SwatTransKernel( | 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4SwatTransKernel( |
| 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 40 | - const QuantMatmulTilingData quantMatmulTilingData) | 40 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 41 | { | 41 | { |
| 42 | using TypeA = fp4x2_e2m1_t; | 42 | using TypeA = fp4x2_e2m1_t; |
| 43 | using TypeB = fp4x2_e2m1_t; | 43 | using TypeB = fp4x2_e2m1_t; |
| @@ -115,7 +115,7 @@ int main(int argc, char* argv[]) | |||
| 115 | constexpr int32_t deviceId = 0; | 115 | constexpr int32_t deviceId = 0; |
| 116 | 116 | ||
| 117 | try { | 117 | try { |
| 118 | - QuantMatmulTilingData tilingData; | 118 | + QuantMatmulMxTilingData tilingData; |
| 119 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; | 119 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; |
| 120 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 120 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 121 | 121 | ||
| @@ -31,12 +31,12 @@ | |||
| 31 | #include "host_utils/io_utils.h" | 31 | #include "host_utils/io_utils.h" |
| 32 | #include "kernel/quant_matmul_mx_kernel_swat_4_buffer.h" | 32 | #include "kernel/quant_matmul_mx_kernel_swat_4_buffer.h" |
| 33 | #include "tiling/quant_matmul_mx_tiling_swat_4_buffer.h" | 33 | #include "tiling/quant_matmul_mx_tiling_swat_4_buffer.h" |
| 34 | -#include "tiling/quant_matmul_tiling_data.h" | 34 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 35 | 35 | ||
| 36 | template <bool TransA, bool TransB> | 36 | template <bool TransA, bool TransB> |
| 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4Swat4BufferKernel( | 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4Swat4BufferKernel( |
| 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 39 | - const QuantMatmulTilingData quantMatmulTilingData) | 39 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 40 | { | 40 | { |
| 41 | using TypeA = fp4x2_e2m1_t; | 41 | using TypeA = fp4x2_e2m1_t; |
| 42 | using TypeB = fp4x2_e2m1_t; | 42 | using TypeB = fp4x2_e2m1_t; |
| @@ -114,7 +114,7 @@ int main(int argc, char* argv[]) | |||
| 114 | constexpr int32_t deviceId = 0; | 114 | constexpr int32_t deviceId = 0; |
| 115 | 115 | ||
| 116 | try { | 116 | try { |
| 117 | - QuantMatmulTilingData tilingData; | 117 | + QuantMatmulMxTilingData tilingData; |
| 118 | QuantMatmulTilingSwat4Buffer<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; | 118 | QuantMatmulTilingSwat4Buffer<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; |
| 119 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 119 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 120 | 120 | ||
| @@ -32,13 +32,13 @@ | |||
| 32 | #include "host_utils/io_utils.h" | 32 | #include "host_utils/io_utils.h" |
| 33 | #include "kernel/quant_matmul_mx_kernel_swat.h" | 33 | #include "kernel/quant_matmul_mx_kernel_swat.h" |
| 34 | #include "tiling/quant_matmul_mx_tiling_swat.h" | 34 | #include "tiling/quant_matmul_mx_tiling_swat.h" |
| 35 | -#include "tiling/quant_matmul_tiling_data.h" | 35 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 36 | #include "utils/constant.h" | 36 | #include "utils/constant.h" |
| 37 | 37 | ||
| 38 | template <bool TransA, bool TransB> | 38 | template <bool TransA, bool TransB> |
| 39 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4SwatWeightNzKernel( | 39 | __global__ __aicore__ __cube__ void QuantMatmulMxfp4SwatWeightNzKernel( |
| 40 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 40 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 41 | - const QuantMatmulTilingData quantMatmulTilingData) | 41 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 42 | { | 42 | { |
| 43 | using TypeA = fp4x2_e2m1_t; | 43 | using TypeA = fp4x2_e2m1_t; |
| 44 | using TypeB = fp4x2_e2m1_t; | 44 | using TypeB = fp4x2_e2m1_t; |
| @@ -118,7 +118,7 @@ int main(int argc, char* argv[]) | |||
| 118 | constexpr int32_t deviceId = 0; | 118 | constexpr int32_t deviceId = 0; |
| 119 | 119 | ||
| 120 | try { | 120 | try { |
| 121 | - QuantMatmulTilingData tilingData; | 121 | + QuantMatmulMxTilingData tilingData; |
| 122 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; | 122 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT4_E2M1, mm::DataType::DT_FLOAT4_E2M1> tilingEngine; |
| 123 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 123 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 124 | 124 | ||
| @@ -31,12 +31,12 @@ | |||
| 31 | #include "host_utils/io_utils.h" | 31 | #include "host_utils/io_utils.h" |
| 32 | #include "kernel/quant_matmul_mx_kernel_a_full_load.h" | 32 | #include "kernel/quant_matmul_mx_kernel_a_full_load.h" |
| 33 | #include "tiling/quant_matmul_mx_tiling_a_full_load.h" | 33 | #include "tiling/quant_matmul_mx_tiling_a_full_load.h" |
| 34 | -#include "tiling/quant_matmul_tiling_data.h" | 34 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 35 | 35 | ||
| 36 | template <bool TransA, bool TransB> | 36 | template <bool TransA, bool TransB> |
| 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8AFullLoadTransKernel( | 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8AFullLoadTransKernel( |
| 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 39 | - const QuantMatmulTilingData quantMatmulTilingData) | 39 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 40 | { | 40 | { |
| 41 | using TypeA = fp8_e4m3fn_t; | 41 | using TypeA = fp8_e4m3fn_t; |
| 42 | using TypeB = fp8_e4m3fn_t; | 42 | using TypeB = fp8_e4m3fn_t; |
| @@ -110,7 +110,7 @@ int main(int argc, char* argv[]) | |||
| 110 | constexpr int32_t deviceId = 0; | 110 | constexpr int32_t deviceId = 0; |
| 111 | 111 | ||
| 112 | try { | 112 | try { |
| 113 | - QuantMatmulTilingData tilingData; | 113 | + QuantMatmulMxTilingData tilingData; |
| 114 | QuantMatmulTilingAFullLoad<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; | 114 | QuantMatmulTilingAFullLoad<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; |
| 115 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 115 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 116 | 116 | ||
| @@ -31,12 +31,12 @@ | |||
| 31 | #include "host_utils/io_utils.h" | 31 | #include "host_utils/io_utils.h" |
| 32 | #include "kernel/quant_matmul_mx_kernel_swat.h" | 32 | #include "kernel/quant_matmul_mx_kernel_swat.h" |
| 33 | #include "tiling/quant_matmul_mx_tiling_swat.h" | 33 | #include "tiling/quant_matmul_mx_tiling_swat.h" |
| 34 | -#include "tiling/quant_matmul_tiling_data.h" | 34 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 35 | 35 | ||
| 36 | template <bool TransA, bool TransB> | 36 | template <bool TransA, bool TransB> |
| 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8SwatKernel( | 37 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8SwatKernel( |
| 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 38 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 39 | - const QuantMatmulTilingData quantMatmulTilingData) | 39 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 40 | { | 40 | { |
| 41 | // Keep the sample explicit about the datatype/layout combination that this | 41 | // Keep the sample explicit about the datatype/layout combination that this |
| 42 | // executable demonstrates so host tiling and kernel traits stay in sync. | 42 | // executable demonstrates so host tiling and kernel traits stay in sync. |
| @@ -116,7 +116,7 @@ int main(int argc, char* argv[]) | |||
| 116 | constexpr int32_t deviceId = 0; | 116 | constexpr int32_t deviceId = 0; |
| 117 | 117 | ||
| 118 | try { | 118 | try { |
| 119 | - QuantMatmulTilingData tilingData; | 119 | + QuantMatmulMxTilingData tilingData; |
| 120 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; | 120 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; |
| 121 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 121 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 122 | 122 | ||
| @@ -30,12 +30,12 @@ | |||
| 30 | #include "host_utils/io_utils.h" | 30 | #include "host_utils/io_utils.h" |
| 31 | #include "kernel/quant_matmul_mx_kernel_swat_4_buffer.h" | 31 | #include "kernel/quant_matmul_mx_kernel_swat_4_buffer.h" |
| 32 | #include "tiling/quant_matmul_mx_tiling_swat_4_buffer.h" | 32 | #include "tiling/quant_matmul_mx_tiling_swat_4_buffer.h" |
| 33 | -#include "tiling/quant_matmul_tiling_data.h" | 33 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 34 | 34 | ||
| 35 | template <bool TransA, bool TransB> | 35 | template <bool TransA, bool TransB> |
| 36 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8Swat4BufferKernel( | 36 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8Swat4BufferKernel( |
| 37 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 37 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 38 | - const QuantMatmulTilingData quantMatmulTilingData) | 38 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 39 | { | 39 | { |
| 40 | using TypeA = fp8_e4m3fn_t; | 40 | using TypeA = fp8_e4m3fn_t; |
| 41 | using TypeB = fp8_e4m3fn_t; | 41 | using TypeB = fp8_e4m3fn_t; |
| @@ -110,7 +110,7 @@ int main(int argc, char* argv[]) | |||
| 110 | constexpr int32_t deviceId = 0; | 110 | constexpr int32_t deviceId = 0; |
| 111 | 111 | ||
| 112 | try { | 112 | try { |
| 113 | - QuantMatmulTilingData tilingData; | 113 | + QuantMatmulMxTilingData tilingData; |
| 114 | QuantMatmulTilingSwat4Buffer<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; | 114 | QuantMatmulTilingSwat4Buffer<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; |
| 115 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 115 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 116 | 116 | ||
| @@ -31,13 +31,13 @@ | |||
| 31 | #include "host_utils/io_utils.h" | 31 | #include "host_utils/io_utils.h" |
| 32 | #include "kernel/quant_matmul_mx_kernel_swat.h" | 32 | #include "kernel/quant_matmul_mx_kernel_swat.h" |
| 33 | #include "tiling/quant_matmul_mx_tiling_swat.h" | 33 | #include "tiling/quant_matmul_mx_tiling_swat.h" |
| 34 | -#include "tiling/quant_matmul_tiling_data.h" | 34 | +#include "tiling/quant_matmul_mx_tiling_data.h" |
| 35 | #include "utils/constant.h" | 35 | #include "utils/constant.h" |
| 36 | 36 | ||
| 37 | template <bool TransA, bool TransB> | 37 | template <bool TransA, bool TransB> |
| 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8SwatWeightNzKernel( | 38 | __global__ __aicore__ __cube__ void QuantMatmulMxfp8SwatWeightNzKernel( |
| 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, | 39 | GM_ADDR dA, GM_ADDR dB, GM_ADDR dScaleA, GM_ADDR dScaleB, GM_ADDR dC, |
| 40 | - const QuantMatmulTilingData quantMatmulTilingData) | 40 | + const QuantMatmulMxTilingData quantMatmulTilingData) |
| 41 | { | 41 | { |
| 42 | // A / scaleA / scaleB follow quant_matmul_mxfp8_swat.asc (Frame + transpose). B is Zn GM (weight NZ). | 42 | // A / scaleA / scaleB follow quant_matmul_mxfp8_swat.asc (Frame + transpose). B is Zn GM (weight NZ). |
| 43 | using TypeA = fp8_e4m3fn_t; | 43 | using TypeA = fp8_e4m3fn_t; |
| @@ -115,7 +115,7 @@ int main(int argc, char* argv[]) | |||
| 115 | constexpr int32_t deviceId = 0; | 115 | constexpr int32_t deviceId = 0; |
| 116 | 116 | ||
| 117 | try { | 117 | try { |
| 118 | - QuantMatmulTilingData tilingData; | 118 | + QuantMatmulMxTilingData tilingData; |
| 119 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; | 119 | QuantMatmulTilingSwat<mm::DataType::DT_FLOAT8_E4M3FN, mm::DataType::DT_FLOAT8_E4M3FN> tilingEngine; |
| 120 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); | 120 | tilingEngine.GetTilingData(m, n, k, transA, transB, tilingData); |
| 121 | 121 | ||
| @@ -37,7 +37,7 @@ protected: | |||
| 37 | return "a_full_load"; | 37 | return "a_full_load"; |
| 38 | } | 38 | } |
| 39 | 39 | ||
| 40 | - void DoOpTiling(QuantMatmulTilingData& tilingData) override | 40 | + void DoOpTiling(QuantMatmulMxTilingData& tilingData) override |
| 41 | { | 41 | { |
| 42 | // The A-full-load path validates eligibility before computing its | 42 | // The A-full-load path validates eligibility before computing its |
| 43 | // L1 layout because later calculations assume A stays resident in L1. | 43 | // L1 layout because later calculations assume A stays resident in L1. |
| @@ -61,7 +61,7 @@ private: | |||
| 61 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); | 61 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); |
| 62 | } | 62 | } |
| 63 | 63 | ||
| 64 | - void BuildTilingData(QuantMatmulTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const | 64 | + void BuildTilingData(QuantMatmulMxTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const |
| 65 | { | 65 | { |
| 66 | // Flatten the host-side search result into the POD payload consumed by | 66 | // Flatten the host-side search result into the POD payload consumed by |
| 67 | // the launcher and device kernel. | 67 | // the launcher and device kernel. |
RSamples/2_Performance/matmul_story/matmul_recipes/include/tiling/quant_matmul_tiling_data.h→Samples/2_Performance/matmul_story/matmul_recipes/include/tiling/quant_matmul_mx_tiling_data.h+2-2
| @@ -9,7 +9,7 @@ | |||
| 9 | */ | 9 | */ |
| 10 | 10 | ||
| 11 | /*! | 11 | /*! |
| 12 | - * \file quant_matmul_tiling_data.h | 12 | + * \file quant_matmul_mx_tiling_data.h |
| 13 | * \brief Serialized tiling data passed from the host launcher to the kernel. | 13 | * \brief Serialized tiling data passed from the host launcher to the kernel. |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| @@ -24,7 +24,7 @@ | |||
| 24 | // The field order is part of the host-device contract, so layout stability is | 24 | // The field order is part of the host-device contract, so layout stability is |
| 25 | // more important here than convenience of reordering members. | 25 | // more important here than convenience of reordering members. |
| 26 | 26 | ||
| 27 | -struct alignas(8) QuantMatmulTilingData { | 27 | +struct alignas(8) QuantMatmulMxTilingData { |
| 28 | // Original problem shape. | 28 | // Original problem shape. |
| 29 | uint32_t m{0}; | 29 | uint32_t m{0}; |
| 30 | uint32_t n{0}; | 30 | uint32_t n{0}; |
| @@ -37,7 +37,7 @@ protected: | |||
| 37 | return "swat"; | 37 | return "swat"; |
| 38 | } | 38 | } |
| 39 | 39 | ||
| 40 | - void DoOpTiling(QuantMatmulTilingData& tilingData) override | 40 | + void DoOpTiling(QuantMatmulMxTilingData& tilingData) override |
| 41 | { | 41 | { |
| 42 | // The streaming path can reuse the common base block search directly, | 42 | // The streaming path can reuse the common base block search directly, |
| 43 | // then specializes only the tail split and L1-depth decisions. | 43 | // then specializes only the tail split and L1-depth decisions. |
| @@ -60,7 +60,7 @@ private: | |||
| 60 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); | 60 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); |
| 61 | } | 61 | } |
| 62 | 62 | ||
| 63 | - void BuildTilingData(QuantMatmulTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const | 63 | + void BuildTilingData(QuantMatmulMxTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const |
| 64 | { | 64 | { |
| 65 | // Flatten the host-side search result into the POD payload consumed by | 65 | // Flatten the host-side search result into the POD payload consumed by |
| 66 | // the launcher and device kernel. | 66 | // the launcher and device kernel. |
| @@ -37,7 +37,7 @@ protected: | |||
| 37 | return "swat_4_buffer"; | 37 | return "swat_4_buffer"; |
| 38 | } | 38 | } |
| 39 | 39 | ||
| 40 | - void DoOpTiling(QuantMatmulTilingData& tilingData) override | 40 | + void DoOpTiling(QuantMatmulMxTilingData& tilingData) override |
| 41 | { | 41 | { |
| 42 | // The streaming path can reuse the common base block search directly, | 42 | // The streaming path can reuse the common base block search directly, |
| 43 | // then specializes only the tail split and L1-depth decisions. | 43 | // then specializes only the tail split and L1-depth decisions. |
| @@ -77,7 +77,7 @@ private: | |||
| 77 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); | 77 | runInfo_.scaleFactorB * runInfo_.stepKb * runInfo_.baseK)); |
| 78 | } | 78 | } |
| 79 | 79 | ||
| 80 | - void BuildTilingData(QuantMatmulTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const | 80 | + void BuildTilingData(QuantMatmulMxTilingData& tilingData, uint32_t scaleKL1, uint8_t nBufferNum) const |
| 81 | { | 81 | { |
| 82 | // Flatten the host-side search result into the POD payload consumed by | 82 | // Flatten the host-side search result into the POD payload consumed by |
| 83 | // the launcher and device kernel. | 83 | // the launcher and device kernel. |
| @@ -20,7 +20,7 @@ | |||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | 22 | ||
| 23 | -#include "quant_matmul_tiling_data.h" | 23 | +#include "quant_matmul_mx_tiling_data.h" |
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | template <mm::DataType aDataType, mm::DataType bDataType> | 26 | template <mm::DataType aDataType, mm::DataType bDataType> |
| @@ -29,7 +29,7 @@ public: | |||
| 29 | QuantMatmulTilingBase() = default; | 29 | QuantMatmulTilingBase() = default; |
| 30 | virtual ~QuantMatmulTilingBase() = default; | 30 | virtual ~QuantMatmulTilingBase() = default; |
| 31 | 31 | ||
| 32 | - void GetTilingData(uint64_t m, uint64_t n, uint64_t k, bool transA, bool transB, QuantMatmulTilingData& tilingData) | 32 | + void GetTilingData(uint64_t m, uint64_t n, uint64_t k, bool transA, bool transB, QuantMatmulMxTilingData& tilingData) |
| 33 | { | 33 | { |
| 34 | // Clear the cached state so one tiling object can safely be reused for | 34 | // Clear the cached state so one tiling object can safely be reused for |
| 35 | // multiple shapes without leaking the previous decision. | 35 | // multiple shapes without leaking the previous decision. |
| @@ -46,7 +46,7 @@ public: | |||
| 46 | PrintTilingData(tilingData); | 46 | PrintTilingData(tilingData); |
| 47 | } | 47 | } |
| 48 | 48 | ||
| 49 | - void GetTilingData(uint64_t m, uint64_t n, uint64_t k, QuantMatmulTilingData& tilingData) | 49 | + void GetTilingData(uint64_t m, uint64_t n, uint64_t k, QuantMatmulMxTilingData& tilingData) |
| 50 | { | 50 | { |
| 51 | // Keep compatibility with the common sample default: | 51 | // Keep compatibility with the common sample default: |
| 52 | // A is not transposed and B is transposed. | 52 | // A is not transposed and B is transposed. |
| @@ -60,10 +60,10 @@ protected: | |||
| 60 | 60 | ||
| 61 | virtual const char* TilingName() const = 0; | 61 | virtual const char* TilingName() const = 0; |
| 62 | 62 | ||
| 63 | - virtual void DoOpTiling(QuantMatmulTilingData& tilingData) = 0; | 63 | + virtual void DoOpTiling(QuantMatmulMxTilingData& tilingData) = 0; |
| 64 | 64 | ||
| 65 | private: | 65 | private: |
| 66 | - void PrintTilingData(const QuantMatmulTilingData& tilingData) const | 66 | + void PrintTilingData(const QuantMatmulMxTilingData& tilingData) const |
| 67 | { | 67 | { |
| 68 | printf("[QuantMatmul Strategy]\n"); | 68 | printf("[QuantMatmul Strategy]\n"); |
| 69 | printf(" strategy : %s\n", TilingName()); | 69 | printf(" strategy : %s\n", TilingName()); |