mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-06 04:51:03 +02:00
use shmem arrays for LUTs
This commit is contained in:
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user