vulkan: Change mul_mat_id to pad K rather than N (#27925)

The N padding is needed for mul_mat, but not mul_mat_id. For mul_mat_id,
we indirect the row index through a shared memory lookup table which avoids
any OOB row coordinate. But that callback doesn't bounds check K, so we
actually need K padding instead.
This commit is contained in:
Jeff Bolz
2026-08-29 10:09:24 +03:00
committed by GitHub
parent d7bd3bfcad
commit 77f132cb1d
5 changed files with 68 additions and 52 deletions
@@ -88,7 +88,6 @@ layout (push_constant) uniform parameter
uint nei1;
uint nbi1;
uint ne11;
uint padded_N;
uint n_experts;
uint hoist_row_ids;
#else
@@ -56,6 +56,8 @@ layout (push_constant) uniform parameter
uint nei1;
uint nbi1;
uint ne11;
uint n_experts;
uint hoist_row_ids;
#else
uint base_work_group_z;
uint num_batches;
@@ -64,12 +66,8 @@ layout (push_constant) uniform parameter
uint ne12;
uint broadcast2;
uint broadcast3;
#endif
// N dimension for the B matrix can be >= p.N
uint padded_N;
#ifdef MUL_MAT_ID
uint n_experts;
uint hoist_row_ids;
#endif
} p;
@@ -332,7 +330,9 @@ void main() {
tensorLayoutNV<2> tensorLayoutA = createTensorLayoutNV(2);
tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutAClamp = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV);
tensorLayoutNV<2> tensorLayoutB = createTensorLayoutNV(2);
#ifndef MUL_MAT_ID
tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutBClamp = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV);
#endif
tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutD = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV);
#if QUANT_K > 1
@@ -345,12 +345,19 @@ void main() {
// Use end_k rather than p.K as the dimension because that's what
// we need to bound check against when using split_k.
// Bounds check B against padded_N, but bounds check D against N.
tensorLayoutA = setTensorLayoutDimensionNV(tensorLayoutA, p.M, end_k);
#ifdef MUL_MAT_ID
// MUL_MAT_ID pads each B row to stride_b so partial K tiles read zeros without clamping.
tensorLayoutB = setTensorLayoutDimensionNV(tensorLayoutB, BN, p.stride_b);
#else
// Bounds check B against padded_N, but bounds check D against N.
tensorLayoutB = setTensorLayoutDimensionNV(tensorLayoutB, p.padded_N, end_k);
#endif
tensorLayoutD = setTensorLayoutDimensionNV(tensorLayoutD, p.N, p.M);
tensorLayoutAClamp = setTensorLayoutDimensionNV(tensorLayoutAClamp, p.M, end_k);
#ifndef MUL_MAT_ID
tensorLayoutBClamp = setTensorLayoutDimensionNV(tensorLayoutBClamp, p.padded_N, end_k);
#endif
tensorLayoutD = setTensorLayoutStrideNV(tensorLayoutD, p.stride_d, 1);
@@ -527,7 +534,9 @@ void main() {
tensorLayoutB = setTensorLayoutStrideNV(tensorLayoutB, stride_b, 1);
#ifndef MUL_MAT_ID
tensorLayoutBClamp = setTensorLayoutStrideNV(tensorLayoutBClamp, stride_b, 1);
#endif
uint k_iters = (end_k - start_k + BK - 1) / BK;
@@ -56,7 +56,6 @@ layout (push_constant) uniform parameter
uint nei1;
uint nbi1;
uint ne11;
uint padded_N;
uint n_experts;
uint hoist_row_ids;
#else