@@ -369,35 +369,35 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
369369#endif
370370 for (var kv_block = subgroup_id ; kv_block < KV_BLOCKS ; kv_block += num_subgroups ) {
371371 let inter_offset = kv_block * SG_MAT_N ;
372- var acc : subgroup_matrix_result <f16 , SG_MAT_M , SG_MAT_N > = subgroupMatrixLoad <subgroup_matrix_result <f16 , SG_MAT_M , SG_MAT_N >>(& inter_shmem , inter_offset , false , KV_TILE );
372+ var acc : subgroup_matrix_result <f16 , SG_MAT_N , SG_MAT_M > = subgroupMatrixLoad <subgroup_matrix_result <f16 , SG_MAT_N , SG_MAT_M >>(& inter_shmem , inter_offset , false , KV_TILE );
373373
374- var q_cur = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K >>(& q_shmem , 0u , false , HEAD_DIM_QK );
374+ var q_cur = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M >>(& q_shmem , 0u , false , HEAD_DIM_QK );
375375
376376#ifdef KV_DIRECT
377- var k_cur = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& K , k_global_offset + 0u , true , params . stride_k1 );
377+ var k_cur = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& K , k_global_offset + 0u , true , params . stride_k1 );
378378#else
379- var k_cur = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& kv_shmem , k_block_offset + 0u , true , HEAD_DIM_QK );
379+ var k_cur = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& kv_shmem , k_block_offset + 0u , true , HEAD_DIM_QK );
380380#endif
381381
382382 var t : u32 = 1u ;
383383 for (; t + 1u < HEAD_DIM_QK / SG_MAT_K ; t += 2u ) {
384384 let h0 = t * SG_MAT_K ;
385- var q0 = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K >>(& q_shmem , h0 , false , HEAD_DIM_QK );
385+ var q0 = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M >>(& q_shmem , h0 , false , HEAD_DIM_QK );
386386#ifdef KV_DIRECT
387- var k0 = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& K , k_global_offset + h0 , true , params . stride_k1 );
387+ var k0 = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& K , k_global_offset + h0 , true , params . stride_k1 );
388388#else
389- var k0 = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& kv_shmem , k_block_offset + h0 , true , HEAD_DIM_QK );
389+ var k0 = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& kv_shmem , k_block_offset + h0 , true , HEAD_DIM_QK );
390390#endif
391391 acc = subgroupMatrixMultiplyAccumulate (q_cur , k_cur , acc );
392392 q_cur = q0 ;
393393 k_cur = k0 ;
394394
395395 let h1 = (t + 1u ) * SG_MAT_K ;
396- var q1g = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K >>(& q_shmem , h1 , false , HEAD_DIM_QK );
396+ var q1g = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M >>(& q_shmem , h1 , false , HEAD_DIM_QK );
397397#ifdef KV_DIRECT
398- var k1g = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& K , k_global_offset + h1 , true , params . stride_k1 );
398+ var k1g = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& K , k_global_offset + h1 , true , params . stride_k1 );
399399#else
400- var k1g = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& kv_shmem , k_block_offset + h1 , true , HEAD_DIM_QK );
400+ var k1g = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& kv_shmem , k_block_offset + h1 , true , HEAD_DIM_QK );
401401#endif
402402 acc = subgroupMatrixMultiplyAccumulate (q_cur , k_cur , acc );
403403 q_cur = q1g ;
@@ -407,11 +407,11 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
407407 // handle odd tail
408408 if (t < HEAD_DIM_QK / SG_MAT_K ) {
409409 let h = t * SG_MAT_K ;
410- var qn = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K >>(& q_shmem , h , false , HEAD_DIM_QK );
410+ var qn = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M >>(& q_shmem , h , false , HEAD_DIM_QK );
411411#ifdef KV_DIRECT
412- var kn = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& K , k_global_offset + h , true , params . stride_k1 );
412+ var kn = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& K , k_global_offset + h , true , params . stride_k1 );
413413#else
414- var kn = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(& kv_shmem , k_block_offset + h , true , HEAD_DIM_QK );
414+ var kn = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(& kv_shmem , k_block_offset + h , true , HEAD_DIM_QK );
415415#endif
416416 acc = subgroupMatrixMultiplyAccumulate (q_cur , k_cur , acc );
417417 q_cur = qn ;
@@ -566,15 +566,15 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
566566 head_dim_block < HEAD_DIM_V ;
567567 head_dim_block += num_subgroups * SG_MAT_N ) {
568568 // load O submatrix from shared memory
569- var o_sg_mat : subgroup_matrix_result <f16 , SG_MAT_M , SG_MAT_N > = subgroupMatrixLoad <subgroup_matrix_result <f16 , SG_MAT_M , SG_MAT_N >>(
569+ var o_sg_mat : subgroup_matrix_result <f16 , SG_MAT_N , SG_MAT_M > = subgroupMatrixLoad <subgroup_matrix_result <f16 , SG_MAT_N , SG_MAT_M >>(
570570 & o_shmem ,
571571 head_dim_block ,
572572 false ,
573573 HEAD_DIM_V
574574 );
575575 for (var kv_block = 0u ; kv_block < KV_BLOCKS ; kv_block ++ ) {
576576 let p_offset = kv_block * SG_MAT_N ;
577- var p_sg_mat : subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K > = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_M , SG_MAT_K >>(
577+ var p_sg_mat : subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M > = subgroupMatrixLoad <subgroup_matrix_left <f16 , SG_MAT_K , SG_MAT_M >>(
578578 & inter_shmem ,
579579 p_offset ,
580580 false ,
@@ -585,15 +585,15 @@ fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
585585#ifdef KV_DIRECT
586586 let v_block_row = kv_tile + kv_block * SG_MAT_N ;
587587 let v_global_offset = v_head_offset + v_block_row * params . stride_v1 + head_dim_block ;
588- var v_sg_mat : subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N > = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(
588+ var v_sg_mat : subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K > = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(
589589 & V ,
590590 v_global_offset ,
591591 false ,
592592 params . stride_v1
593593 );
594594#else
595595 let v_block_offset = kv_block * SG_MAT_N * HEAD_DIM_V ;
596- var v_sg_mat : subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N > = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_K , SG_MAT_N >>(
596+ var v_sg_mat : subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K > = subgroupMatrixLoad <subgroup_matrix_right <f16 , SG_MAT_N , SG_MAT_K >>(
597597 & kv_shmem ,
598598 v_block_offset + head_dim_block ,
599599 false ,
0 commit comments