use shmem arrays for LUTs

This commit is contained in:
Ruben Ortlam
2026-08-24 17:37:24 +02:00
parent 6a08f25c40
commit ad46e9d244
2 changed files with 24 additions and 8 deletions
@@ -101,6 +101,10 @@ shared float buf_b_s[BN * BK_STEP];
#define QUANT_OFFSET 16.0
#endif
#if defined(DATA_A_IQ4_NL) || defined(DATA_A_MXFP4)
shared int8_t cm1_kvalues[16];
#endif
#define LOAD_VEC_A (4 * QUANT_R)
#define LOAD_VEC_B 16
@@ -119,6 +123,18 @@ shared int32_t cm_layout_probe[TM * TN];
#include "mul_mmq_cm1_funcs.glsl"
void main() {
#if defined(DATA_A_IQ4_NL)
if (gl_LocalInvocationIndex < 16u) {
cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex];
}
barrier();
#elif defined(DATA_A_MXFP4)
if (gl_LocalInvocationIndex < 16u) {
cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
}
barrier();
#endif
const uint blocks_m = (p.M + BM - 1) / BM;
const uint ik = gl_WorkGroupID.x / blocks_m;
@@ -172,11 +172,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] =
pack32(i8vec4(kvalues_iq4nl_const[lo_idx.x], kvalues_iq4nl_const[lo_idx.y],
kvalues_iq4nl_const[lo_idx.z], kvalues_iq4nl_const[lo_idx.w]));
pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
pack32(i8vec4(kvalues_iq4nl_const[hi_idx.x], kvalues_iq4nl_const[hi_idx.y],
kvalues_iq4nl_const[hi_idx.z], kvalues_iq4nl_const[hi_idx.w]));
pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = float(blk.d);
@@ -204,11 +204,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr ] =
pack32(i8vec4(kvalues_mxfp4_const[lo_idx.x], kvalues_mxfp4_const[lo_idx.y],
kvalues_mxfp4_const[lo_idx.z], kvalues_mxfp4_const[lo_idx.w]));
pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
pack32(i8vec4(kvalues_mxfp4_const[hi_idx.x], kvalues_mxfp4_const[hi_idx.y],
kvalues_mxfp4_const[hi_idx.z], kvalues_mxfp4_const[hi_idx.w]));
pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
if (loadr == 0) {
buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5;