@@ -59,21 +59,39 @@ From https://github.com/RadeonOpenCompute/amd_matrix_instruction_calculator
5959./matrix_calculator.py --architecture cdna1 --instruction v_mfma_f32_16x16x16f16
6060--detail-instruction
6161*/
62- Fragment makeGemmFragmentAB16x16CDNA () {
62+ Fragment makeGemmFragmentAB16x16CDNA (const int k_pack ) {
6363 IterVar i = make_itervar (" i" , 16 );
64+ IterVar j = make_itervar (" j" , 16 * k_pack);
65+ IterVar rep = make_itervar (" rep" , 1 );
66+ PrimExpr forward_thread = 16 * FloorDiv (j->var , 4 * k_pack) + i;
67+ PrimExpr index = FloorMod (j->var , 4 * k_pack);
68+ return Fragment ({i, j}, {index}, forward_thread, rep);
69+ }
70+
71+ Fragment makeGemmFragmentAB16x16CDNATransposed (const int k_pack) {
72+ IterVar i = make_itervar (" i" , 16 * k_pack);
6473 IterVar j = make_itervar (" j" , 16 );
6574 IterVar rep = make_itervar (" rep" , 1 );
66- PrimExpr forward_thread = 16 * FloorDiv (j ->var , 4 ) + i ;
67- PrimExpr index = FloorMod (j ->var , 4 );
75+ PrimExpr forward_thread = 16 * FloorDiv (i ->var , 4 * k_pack ) + j ;
76+ PrimExpr index = FloorMod (i ->var , 4 * k_pack );
6877 return Fragment ({i, j}, {index}, forward_thread, rep);
6978}
7079
71- Fragment makeGemmFragmentAB16x16CDNATransposed ( ) {
80+ Fragment makeGemmFragmentAB16x32CDNA ( const int k_pack ) {
7281 IterVar i = make_itervar (" i" , 16 );
82+ IterVar j = make_itervar (" j" , 32 * k_pack);
83+ IterVar rep = make_itervar (" rep" , 1 );
84+ PrimExpr forward_thread = 16 * FloorDiv (j->var , 8 * k_pack) + i;
85+ PrimExpr index = FloorMod (j->var , 8 * k_pack);
86+ return Fragment ({i, j}, {index}, forward_thread, rep);
87+ }
88+
89+ Fragment makeGemmFragmentAB16x32CDNATransposed (const int k_pack) {
90+ IterVar i = make_itervar (" i" , 32 * k_pack);
7391 IterVar j = make_itervar (" j" , 16 );
7492 IterVar rep = make_itervar (" rep" , 1 );
75- PrimExpr forward_thread = 16 * FloorDiv (i->var , 4 ) + j;
76- PrimExpr index = FloorMod (i->var , 4 );
93+ PrimExpr forward_thread = 16 * FloorDiv (i->var , 8 * k_pack ) + j;
94+ PrimExpr index = FloorMod (i->var , 8 * k_pack );
7795 return Fragment ({i, j}, {index}, forward_thread, rep);
7896}
7997
@@ -224,27 +242,34 @@ Fragment makeGemmFragmentB(const int block_m, const int block_n,
224242Fragment makeGemmFragmentACDNA (const int block_m, const int block_n,
225243 const int block_k, const int warp_m,
226244 const int warp_n, const int element_size,
227- bool transposed) {
245+ const int k_pack, bool transposed) {
228246 // assume not transposed
229247 ICHECK (block_m % warp_m == 0 );
230248 ICHECK (block_n % warp_n == 0 );
231249 ICHECK (warp_m % 16 == 0 );
232- ICHECK (block_k % 16 == 0 );
250+ const int mfma_k = k_pack * (element_size == 16 ? 16 : 32 );
251+ ICHECK (block_k % mfma_k == 0 );
233252 ICHECK (element_size == 8 || element_size == 16 )
234253 << " element bitwidth=" << element_size;
235254 if (transposed) {
236255 auto base_layout =
237- makeGemmFragmentAB16x16CDNATransposed ()->Repeat ({1 , 1 }, false , false );
256+ element_size == 16
257+ ? makeGemmFragmentAB16x16CDNATransposed (k_pack)->Repeat (
258+ {1 , 1 }, false , false )
259+ : makeGemmFragmentAB16x32CDNATransposed (k_pack)->Repeat (
260+ {1 , 1 }, false , false );
238261 auto warp_layout =
239- base_layout->Repeat ({block_k / 16 , warp_m / 16 }, false , true );
262+ base_layout->Repeat ({block_k / mfma_k , warp_m / 16 }, false , true );
240263 auto block_layout = warp_layout->Repeat ({1 , block_m / warp_m}, true , true )
241264 ->Replicate (block_n / warp_n);
242265 return block_layout;
243266 } else {
244267 auto base_layout =
245- makeGemmFragmentAB16x16CDNA ()->Repeat ({1 , 1 }, false , false );
268+ element_size == 16
269+ ? makeGemmFragmentAB16x16CDNA (k_pack)->Repeat ({1 , 1 }, false , false )
270+ : makeGemmFragmentAB16x32CDNA (k_pack)->Repeat ({1 , 1 }, false , false );
246271 auto warp_layout =
247- base_layout->Repeat ({warp_m / 16 , block_k / 16 }, false , false );
272+ base_layout->Repeat ({warp_m / 16 , block_k / mfma_k }, false , false );
248273 auto block_layout = warp_layout->Repeat ({block_m / warp_m, 1 }, true , true )
249274 ->Replicate (block_n / warp_n);
250275 return block_layout;
@@ -397,7 +422,7 @@ Layout makeMatrixCoreSwizzleLayout(int stride, int continuous, int element_size,
397422 const int numBanks = 32 ;
398423 const int bankBitWidth = 32 ;
399424 const int SIMDWidth = 16 ;
400- const int vecSize = 4 * kPack ;
425+ const int vecSize = ( 64 / element_size) * kPack ;
401426 const int innerDimLength = continuous;
402427 const int typeWidthInBit = element_size;
403428
@@ -616,12 +641,7 @@ Layout makeGemmABLayoutHopper(int mat_stride, int mat_continuous,
616641
617642Layout makeGemmABLayoutCDNA (int stride, int continuous, int element_size,
618643 int kPack ) {
619- int vector_size = 128 / element_size;
620- if (continuous % (vector_size * 4 ) == 0 )
621- return makeMatrixCoreSwizzleLayout (stride, continuous, element_size, kPack );
622- else {
623- return makeGemmABLayoutPadded (stride, continuous, element_size);
624- }
644+ return makeMatrixCoreSwizzleLayout (stride, continuous, element_size, kPack );
625645}
626646} // namespace tl
627647} // namespace tvm
0 commit comments