已合并
refactor(matmul_recipes): rename QuantMatmulTilingData to QuantMatmulMxTilingData #393
refactor(matmul_recipes): rename QuantMatmulTilingData to QuantMatmulMxTilingData #393
已合并
tangweiwei2创建于 21 天前
14 个文件变更+38-38
@@ -85,7 +85,7 @@ inline uint64_t ParsePositiveUint64(const char* arg, const char* name)
85 85 
86inline void CheckUint32Shape(uint64_t value, const char* name)86inline 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 
37template <bool TransA, bool TransB>37template <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 
37template <bool TransA, bool TransB>37template <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 
36template <bool TransA, bool TransB>36template <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 
38template <bool TransA, bool TransB>38template <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 
36template <bool TransA, bool TransB>36template <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 
36template <bool TransA, bool TransB>36template <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 this41 // 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 
35template <bool TransA, bool TransB>35template <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 
37template <bool TransA, bool TransB>37template <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) override40+ void DoOpTiling(QuantMatmulMxTilingData& tilingData) override
41 {41 {
42 // The A-full-load path validates eligibility before computing its42 // 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) const64+ 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 by66 // 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.hSamples/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.h12+ * \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 is24// 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#pragma pack(push, 8)26#pragma pack(push, 8)
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) override40+ 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) const63+ 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 by65 // 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) override40+ 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) const80+ 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 by82 // 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#include "host_utils/common_utils.h"21#include "host_utils/common_utils.h"
22#include "quant_matmul_tiling_common.h"22#include "quant_matmul_tiling_common.h"
23-#include "quant_matmul_tiling_data.h"23+#include "quant_matmul_mx_tiling_data.h"
24#include "utils/constant.h"24#include "utils/constant.h"
25 25 
26template <mm::DataType aDataType, mm::DataType bDataType>26template <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 for34 // 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 
65private:65private:
66- void PrintTilingData(const QuantMatmulTilingData& tilingData) const66+ 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());