add BK_STEP to shader, default to 2

This commit is contained in:
Ruben Ortlam
2026-08-24 09:33:15 +02:00
parent 33a4a5d637
commit 1a97c2c54d
2 changed files with 71 additions and 64 deletions
@@ -80,15 +80,16 @@ layout (constant_id = 9) const uint TK = 16;
layout (constant_id = 10) const uint WARP = 32;
#define BK 32
#define BK_STEP 2
const uint shmem_stride = (BK / 4) + 4;
const uint QPITCH = BK_STEP * (BK / 4) + 4;
// Shared memory cache
shared uint32_t buf_a_qs[BM * shmem_stride];
shared float16_t buf_a_d[BM];
shared uint32_t buf_a_qs[BM * QPITCH];
shared float16_t buf_a_d[BM * BK_STEP];
shared uint32_t buf_b_qs[BN * shmem_stride];
shared float16_t buf_b_d[BN];
shared uint32_t buf_b_qs[BN * QPITCH];
shared float16_t buf_b_d[BN * BK_STEP];
#define LOAD_VEC_A (4 * QUANT_R)
#define LOAD_VEC_B 16
@@ -228,58 +229,65 @@ void main() {
sums[i] = ACC_TYPE(0.0);
}
for (uint block = start_k; block < end_k; block += BK) {
[[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) {
const uint buf_ib = loadc_a + l;
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK;
const uint iqs = loadr_a;
block_a_to_shmem(buf_ib, ib, iqs);
}
[[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) {
const uint buf_ib = loadc_b + l;
for (uint block = start_k; block < end_k; block += BK * BK_STEP) {
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
const bool k_in_bounds = block + ks * BK < end_k;
[[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) {
const uint buf_ib = loadc_a + l;
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK + ks;
if (k_in_bounds) {
block_a_to_shmem(buf_ib, ib, loadr_a, ks);
} else if (loadr_a == 0) {
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(0.0);
}
}
[[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) {
const uint buf_ib = loadc_b + l;
#ifdef MUL_MAT_ID
const u16vec2 row_idx = row_ids[buf_ib];
const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK;
const u16vec2 row_idx = row_ids[buf_ib];
const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK + ks;
#else
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK;
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK + ks;
#endif
const uint iqs = loadr_b;
block_b_to_shmem(buf_ib, ib, iqs);
}
barrier();
pos_a_ib += 1;
pos_b_ib += 1;
[[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) {
cm_result[idx] = coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0);
}
// Calculate quants
[[unroll]] for (uint i = 0; i < BK; i += TK) {
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
coopMatLoad(cache_a, buf_a_qs, (warp_r * WM + cm_row * TM) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutRowMajor);
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
coopMatLoad(cache_b, buf_b_qs, (warp_c * WN + cm_col * TN) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutColumnMajor);
cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]);
if (k_in_bounds) {
block_b_to_shmem(buf_ib, ib, loadr_b, ks);
} else if (loadr_b == 0) {
buf_b_d[ks * BN + buf_ib] = FLOAT_TYPE(0.0);
}
}
}
// Apply scales directly from coopmat elements
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
const uint tile_idx = cm_col * cms_per_row + cm_row;
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
const ACC_TYPE da = ACC_TYPE(buf_a_d[warp_r * WM + cm_row * TM + elem_row[e]]);
const ACC_TYPE db = ACC_TYPE(buf_b_d[warp_c * WN + cm_col * TN + elem_col[e]]);
sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db;
barrier();
pos_a_ib += BK_STEP;
pos_b_ib += BK_STEP;
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
[[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) {
cm_result[idx] = coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0);
}
[[unroll]] for (uint i = 0; i < BK; i += TK) {
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
coopMatLoad(cache_a, buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutRowMajor);
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
coopMatLoad(cache_b, buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]);
}
}
}
// Apply scales directly from coopmat elements
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
const uint tile_idx = cm_col * cms_per_row + cm_row;
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
const ACC_TYPE da = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + cm_row * TM + elem_row[e]]);
const ACC_TYPE db = ACC_TYPE(buf_b_d[ks * BN + warp_c * WN + cm_col * TN + elem_col[e]]);
sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db;
}
}
}
}
@@ -9,7 +9,7 @@
#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1)
// 2-byte loads for Q4_0 blocks (18 bytes)
// 4-byte loads for Q4_1 blocks (20 bytes)
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
#ifdef DATA_A_Q4_0
const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2],
data_a_packed16[ib].qs[iqs * 2 + 1]));
@@ -24,12 +24,12 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080;
hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080;
buf_a_qs[buf_ib * shmem_stride + iqs ] = lo4;
buf_a_qs[buf_ib * shmem_stride + iqs + 4] = hi4;
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs ] = lo4;
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs + 4] = hi4;
if (iqs == 0) {
#ifdef DATA_A_Q4_0
buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
#else // DATA_A_Q4_1
#endif
}
@@ -44,14 +44,14 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
#if defined(DATA_A_Q8_0)
// 2-byte loads for Q8_0 blocks (34 bytes)
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2],
data_a_packed16[ib].qs[iqs * 2 + 1]));
buf_a_qs[buf_ib * shmem_stride + iqs] = vui;
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs] = vui;
if (iqs == 0) {
buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
}
}
#endif
@@ -78,18 +78,17 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
// 2-byte loads for Q6_K blocks (210 bytes)
#endif
void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
const uint ib_outer = ib / 4;
const uint ib_inner = ib % 4;
if (iqs == 0) {
// Divide by TK for matmul scale application
buf_b_d[buf_ib] = data_b[ib_outer].ds[ib_inner].x;
buf_b_d[ks * BN + buf_ib] = data_b[ib_outer].ds[ib_inner].x;
}
const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs];
buf_b_qs[buf_ib * shmem_stride + iqs * 4 ] = values.x;
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 1] = values.y;
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 2] = values.z;
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 3] = values.w;
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 ] = values.x;
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 1] = values.y;
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 2] = values.z;
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 3] = values.w;
}