merge shmem arrays

This commit is contained in:
Ruben Ortlam
2026-08-29 11:51:15 +02:00
parent 2ba1d225e1
commit 9ef5996da4
2 changed files with 26 additions and 27 deletions
@@ -95,13 +95,16 @@ const uint QPITCH = BK_STEP * (BK / 4) + 4;
// Shared memory cache
shared uint32_t buf_a_qs[BM * QPITCH];
#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
shared vec2 buf_a_dm[BM * BK_STEP]; // .x = d (== old buf_a_d), .y = m (== old buf_a_m)
#else
shared float buf_a_d[BM * BK_STEP * KSCALES];
#endif
shared uint32_t buf_b_qs[BN * QPITCH];
shared float buf_b_d[BN * BK_STEP];
#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
shared float buf_a_m[BM * BK_STEP];
shared float buf_b_s[BN * BK_STEP];
#endif
@@ -205,8 +208,8 @@ void main() {
const uint warp_c = warp_i / (BM / WM);
// Probe coopmat element layout to discover row/col mapping
uint elem_row[CM_ELEMS];
uint elem_col[CM_ELEMS];
uint8_t elem_row[CM_ELEMS];
uint8_t elem_col[CM_ELEMS];
for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) {
cm_layout_probe[i] = int32_t(i);
@@ -217,8 +220,8 @@ void main() {
coopMatLoad(probe, cm_layout_probe, 0, TN, gl_CooperativeMatrixLayoutRowMajor);
[[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) {
elem_row[e] = uint(probe[e]) / TN;
elem_col[e] = uint(probe[e]) % TN;
elem_row[e] = uint8_t(uint(probe[e]) / TN);
elem_col[e] = uint8_t(uint(probe[e]) % TN);
}
const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A);
@@ -333,7 +336,7 @@ void main() {
float scale_a[CM_ELEMS];
float nbias_a[CM_ELEMS];
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + elem_row[e]];
scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + uint(elem_row[e])];
if (USE_MAGIC_BIAS) {
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
}
@@ -343,7 +346,7 @@ void main() {
coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> cache_b;
coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
@@ -385,15 +388,15 @@ void main() {
float nbias_a[CM_ELEMS];
float ma[CM_ELEMS];
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + uint(elem_row[e])];
scale_a[e] = dm.x;
if (USE_MAGIC_BIAS) {
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
}
#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
ma[e] = float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]);
ma[e] = dm.y;
#else
ma[e] = float(QUANT_OFFSET) * scale_a[e]
+ float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]);
ma[e] = float(QUANT_OFFSET) * scale_a[e] + dm.y;
#endif
}
@@ -407,8 +410,8 @@ void main() {
}
const uint tile_idx = r * cms_per_col + c;
const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col[0]]);
const float sb = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + uint(elem_col[0])]);
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
if (USE_MAGIC_BIAS) {
const float t = fma(intBitsToFloat(int(acc[e])),
@@ -441,14 +444,14 @@ void main() {
float scale_b[cms_per_col];
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
}
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
float scale_a[CM_ELEMS];
float nbias_a[CM_ELEMS];
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + uint(elem_row[e])];
if (USE_MAGIC_BIAS) {
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
}
@@ -497,10 +500,10 @@ void main() {
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
const uint tile_idx = r * cms_per_col + c;
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
const uint col_i = dc + c * TN + elem_col[e];
const uint col_i = dc + c * TN + uint(elem_col[e]);
if (col_i >= _ne1) continue;
const uint row_g = dr + r * TM + elem_row[e];
const uint row_g = dr + r * TM + uint(elem_row[e]);
if (row_g >= p.M) continue;
const u16vec2 row_idx = row_ids[col_i - ic * BN];
@@ -516,8 +519,8 @@ void main() {
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
const uint tile_idx = r * cms_per_col + c;
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
const uint row_g = dr + r * TM + elem_row[e];
const uint col_g = dc + c * TN + elem_col[e];
const uint row_g = dr + r * TM + uint(elem_row[e]);
const uint col_g = dc + c * TN + uint(elem_col[e]);
if (row_g < p.M && col_g < p.N) {
data_d[offsets + col_g * p.stride_d + row_g] = D_TYPE(sums[tile_idx * CM_ELEMS + e]);
}
@@ -55,8 +55,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = float(blk.dm.x);
buf_a_m[ks * BM + buf_ib] = float(blk.dm.y);
buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
}
}
@@ -119,8 +118,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = float(blk.dm.x);
buf_a_m[ks * BM + buf_ib] = float(blk.dm.y);
buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
}
}
@@ -258,8 +256,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
}
vec2 dm = vec2(data_a_packed32[ib_k].dm);
float d_scaled = dm.x * float(sc_val);
buf_a_d[ks * BM + buf_ib] = d_scaled;
buf_a_m[ks * BM + buf_ib] = 8.0 * d_scaled - (dm.y * float(mn_val));
buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 8.0 * d_scaled - (dm.y * float(mn_val)));
}
}
@@ -316,8 +313,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
}
vec2 dm = vec2(data_a_packed32[ib_k].dm);
float d_scaled = dm.x * float(sc_val);
buf_a_d[ks * BM + buf_ib] = d_scaled;
buf_a_m[ks * BM + buf_ib] = 16.0 * d_scaled - (dm.y * float(mn_val));
buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 16.0 * d_scaled - (dm.y * float(mn_val)));
}
}