vulkan: skip unneeded MoE work in mul_mm coopmat1 path (#25483)

This commit is contained in:
Jiang, Fish
2026-09-17 14:53:07 +08:00
committed by GitHub
parent 817e5f83eb
commit 7490357f22
@@ -313,6 +313,9 @@ void main() {
// Workgroup has no work
if (ic * BN >= _ne1) return;
uint required_work_items = (_ne1 - ic * BN) * BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B;
uint required_warp_c = (_ne1 - ic * BN + WN - 1) / WN;
#endif
#ifdef MUL_MAT_ID
@@ -363,6 +366,9 @@ void main() {
[[unroll]] for (uint l = 0; l < BM; l += loadstride_a) {
load_a_to_shmem(pos_a, loadr_a, loadc_a + l, ir * BM + loadc_a + l, block, end_k);
}
#ifdef MUL_MAT_ID
if (gl_LocalInvocationID.x < required_work_items) {
#endif
[[unroll]] for (uint l = 0; l < BN; l += loadstride_b) {
#if !defined(MUL_MAT_ID)
load_b_to_shmem(pos_b, loadr_b, loadc_b + l, ic * BN + loadc_b + l, block, end_k);
@@ -370,6 +376,9 @@ void main() {
load_b_to_shmem(pos_b, loadr_b, loadc_b + l, ic, _ne1, block, end_k);
#endif
}
#ifdef MUL_MAT_ID
}
#endif
barrier();
@@ -377,6 +386,9 @@ void main() {
pos_b += BK / LOAD_VEC_B_EFF;
#ifdef COOPMAT
#ifdef MUL_MAT_ID
if (warp_c < required_warp_c) {
#endif
[[unroll]] for (uint i = 0; i < BK; i += TK) {
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
// Load from shared into cache
@@ -389,6 +401,9 @@ void main() {
}
}
}
#ifdef MUL_MAT_ID
}
#endif
#else
[[unroll]] for (uint i = 0; i < BK / BK_STEP; i++) {
// Load from shared into cache