move quant-specific prefetch function out of main file

This commit is contained in:
Ruben Ortlam
2026-08-24 15:12:28 +02:00
parent 43d831641c
commit 8fc44e56be
2 changed files with 78 additions and 170 deletions
@@ -245,89 +245,7 @@ void main() {
ivec4 pre_b_qs[B_LOADS * BK_STEP];
float16_t pre_b_d [B_LOADS * BK_STEP];
// Prefetch: global memory → registers
#define PREFETCH_BLOCK(blk) \
[[unroll]] for (uint li = 0; li < A_LOADS; li++) { \
const uint buf_ib = loadc_a + li * loadstride_a; \
if (buf_ib < BM) { \
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
pre_a_qs[li * BK_STEP + ks] = \
pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \
data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \
pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \
} \
} \
} \
[[unroll]] for (uint li = 0; li < B_LOADS; li++) { \
const uint buf_ib = loadc_b + li * loadstride_b; \
if (buf_ib < BN) { \
B_IB_CALC \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \
const uint ib_outer = ib_k / 4; \
const uint ib_inner = ib_k % 4; \
pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \
pre_b_d[li * BK_STEP + ks] = data_b[ib_outer].ds[ib_inner].x; \
} \
} \
}
#ifdef MUL_MAT_ID
#define B_IB_CALC \
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;
#else
#define B_IB_CALC \
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK;
#endif
// Store: registers → shared memory (with quant-specific unpacking)
#define STORE_BLOCK_TO_LDS(blk) \
[[unroll]] for (uint li = 0; li < A_LOADS; li++) { \
const uint buf_ib = loadc_a + li * loadstride_a; \
if (buf_ib < BM) { \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const uint idx = li * BK_STEP + ks; \
STORE_A_QS(buf_ib, ks, idx) \
if (loadr_a == 0) { \
buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \
} \
} \
} \
} \
[[unroll]] for (uint li = 0; li < B_LOADS; li++) { \
const uint buf_ib = loadc_b + li * loadstride_b; \
if (buf_ib < BN) { \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const bool in_bounds = (blk) + ks * BK < end_k; \
const uint idx = li * BK_STEP + ks; \
const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \
const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \
buf_b_qs[base ] = v.x; \
buf_b_qs[base + 1] = v.y; \
buf_b_qs[base + 2] = v.z; \
buf_b_qs[base + 3] = v.w; \
if (loadr_b == 0) { \
buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \
} \
} \
} \
}
#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1)
#define STORE_A_QS(buf_ib, ks, idx) \
uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \
uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \
lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \
hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4;
#elif defined(DATA_A_Q8_0)
#define STORE_A_QS(buf_ib, ks, idx) \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx];
#endif
#include "mul_mmq_cm1_funcs.glsl"
// Prefetch first block
if (start_k < end_k) {
@@ -1,94 +1,84 @@
#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require
#include "types.glsl"
// Each iqs value maps to a 32-bit integer
// Quant-specific A-side unpacking: registers → shared memory
#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, 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]));
#else // DATA_A_Q4_1
const uint32_t vui = data_a_packed32[ib].qs[iqs];
#define STORE_A_QS(buf_ib, ks, idx) \
uint32_t lo4 = pre_a_qs[idx] & 0x0F0F0F0F; \
uint32_t hi4 = (pre_a_qs[idx] >> 4) & 0x0F0F0F0F; \
lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080; \
hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080; \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a ] = lo4; \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a + 4] = hi4;
#elif defined(DATA_A_Q8_0)
#define STORE_A_QS(buf_ib, ks, idx) \
buf_a_qs[(buf_ib) * QPITCH + (ks) * (BK / 4) + loadr_a] = pre_a_qs[idx];
#endif
uint32_t lo4 = vui & 0x0F0F0F0F;
uint32_t hi4 = (vui >> 4) & 0x0F0F0F0F;
// subtract 8 from each byte
lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080;
hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080;
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[ks * BM + buf_ib] = float(data_a_packed16[ib].d);
#else // DATA_A_Q4_1
#endif
}
}
#ifdef MUL_MAT_ID
#define B_IB_CALC \
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;
#else
#define B_IB_CALC \
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK;
#endif
#if defined(DATA_A_Q5_0) || defined(DATA_A_Q5_1)
// 2-byte loads for Q5_0 blocks (22 bytes)
// 4-byte loads for Q5_1 blocks (24 bytes)
}
#endif
#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, 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 * QPITCH + ks * (BK / 4) + iqs] = vui;
if (iqs == 0) {
buf_a_d[ks * BM + buf_ib] = float(data_a_packed16[ib].d);
}
}
#endif
#if defined(DATA_A_MXFP4)
// 1-byte loads for mxfp4 blocks (17 bytes)
#endif
// For k-quants, ib and iqs still assume 32-wide blocks, but k-quants are 256-wide
// iqs still refers to a 32-bit integer, meaning 0..7 for 32-wide quants
#if defined(DATA_A_Q2_K)
// 4-byte loads for Q2_K blocks (84 bytes)
#endif
#if defined(DATA_A_Q3_K)
// 2-byte loads for Q3_K blocks (110 bytes)
#endif
#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
// 4-byte loads for Q4_K blocks (144 bytes) and Q5_K blocks (176 bytes)
#endif
#if defined(DATA_A_Q6_K)
// 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, const uint ks) {
const uint ib_outer = ib / 4;
const uint ib_inner = ib % 4;
if (iqs == 0) {
buf_b_d[ks * BN + buf_ib] = float(data_b[ib_outer].ds[ib_inner].x);
// Prefetch: global memory → registers
#define PREFETCH_BLOCK(blk) \
[[unroll]] for (uint li = 0; li < A_LOADS; li++) { \
const uint buf_ib = loadc_a + li * loadstride_a; \
if (buf_ib < BM) { \
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK; \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
pre_a_qs[li * BK_STEP + ks] = \
pack32(u16vec2(data_a_packed16[ib + ks].qs[loadr_a * 2], \
data_a_packed16[ib + ks].qs[loadr_a * 2 + 1])); \
pre_a_d[li * BK_STEP + ks] = data_a_packed16[ib + ks].d; \
} \
} \
} \
[[unroll]] for (uint li = 0; li < B_LOADS; li++) { \
const uint buf_ib = loadc_b + li * loadstride_b; \
if (buf_ib < BN) { \
B_IB_CALC \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const uint ib_k = ((blk) + ks * BK < end_k) ? (ib + ks) : ib; \
const uint ib_outer = ib_k / 4; \
const uint ib_inner = ib_k % 4; \
pre_b_qs[li * BK_STEP + ks] = data_b[ib_outer].qs[ib_inner * 2 + loadr_b]; \
pre_b_d[li * BK_STEP + ks] = 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 * 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;
}
// Store: registers → shared memory (with quant-specific unpacking)
#define STORE_BLOCK_TO_LDS(blk) \
[[unroll]] for (uint li = 0; li < A_LOADS; li++) { \
const uint buf_ib = loadc_a + li * loadstride_a; \
if (buf_ib < BM) { \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const uint idx = li * BK_STEP + ks; \
STORE_A_QS(buf_ib, ks, idx) \
if (loadr_a == 0) { \
buf_a_d[ks * BM + buf_ib] = float(pre_a_d[idx]); \
} \
} \
} \
} \
[[unroll]] for (uint li = 0; li < B_LOADS; li++) { \
const uint buf_ib = loadc_b + li * loadstride_b; \
if (buf_ib < BN) { \
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) { \
const bool in_bounds = (blk) + ks * BK < end_k; \
const uint idx = li * BK_STEP + ks; \
const ivec4 v = in_bounds ? pre_b_qs[idx] : ivec4(0); \
const uint base = buf_ib * QPITCH + ks * (BK / 4) + loadr_b * 4; \
buf_b_qs[base ] = v.x; \
buf_b_qs[base + 1] = v.y; \
buf_b_qs[base + 2] = v.z; \
buf_b_qs[base + 3] = v.w; \
if (loadr_b == 0) { \
buf_b_d[ks * BN + buf_ib] = in_bounds ? float(pre_b_d[idx]) : 0.0f; \
} \
} \
} \
}