mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-18 16:55:05 +02:00
vulkan: skip unneeded MoE work in mul_mm coopmat1 path (#25483)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user