opencl: add int8 dp4 dense and MoE prefill optimization for Adreno GPUs (#25537)

* opencl: add int8 dp4 dense and moe GEMM

* opencl: refactor

---------

Co-authored-by: Li He <lih@qti.qualcomm.com>
This commit is contained in:
Hongqiang Wang
2026-07-10 23:05:58 -07:00
committed by GitHub
co-authored by Li He
parent 4f37f51972
commit 1d1d9a9ed7
21 changed files with 4791 additions and 137 deletions
+16
View File
@@ -114,7 +114,9 @@ set(GGML_OPENCL_KERNELS
mul_mv_id_mxfp4_f32
mul_mv_id_mxfp4_f32_flat
gemm_moe_q4_0_f32_ns
gemm_moe_q4_0_q8_1_dp4a
gemv_moe_q4_0_f32_ns
gemm_moe_q8_0_f32_ns
gemm_moe_q4_1_f32_ns
gemv_moe_q4_1_f32_ns
gemm_moe_q5_0_f32_ns
@@ -122,6 +124,18 @@ set(GGML_OPENCL_KERNELS
gemm_moe_q5_1_f32_ns
gemv_moe_q5_1_f32_ns
gemm_moe_q4_k_f32_ns
gemm_moe_q4_k_q8_1_dp4a
gemm_moe_q6_k_q8_1_dp4a
gemm_moe_q8_1_dp4a
moe_reorder_quant_a_q8_1
gemm_noshuffle_q4_k_q8_1_dp4a
gemm_noshuffle_q5_k_q8_1_dp4a
gemm_noshuffle_q6_k_q8_1_dp4a
gemm_noshuffle_q8_0_q8_1_dp4a
gemm_noshuffle_q5_0_q8_1_dp4a
gemm_noshuffle_iq4_nl_q8_1_dp4a
gemm_noshuffle_q4_0_q8_1_dp4a
quant_a_q8_1
gemv_moe_q4_k_f32_ns
gemm_moe_q5_k_f32_ns
gemv_moe_q5_k_f32_ns
@@ -130,8 +144,10 @@ set(GGML_OPENCL_KERNELS
gemm_moe_mxfp4_f32
gemv_moe_mxfp4_f32
gemm_moe_mxfp4_f32_ns
gemm_moe_mxfp4_q8_1_dp4a
gemv_moe_mxfp4_f32_ns
moe_reorder_b
moe_combine
moe_sort_by_expert
mul_mm_f32_f32_l4_lm
mul_mm_f16_f32_l4_lm
File diff suppressed because it is too large Load Diff
+118
View File
@@ -2372,3 +2372,121 @@ kernel void kernel_restore_block_iq4_nl_noshuffle(
b->qs[2*i + 1] = convert_uchar(((x0 & mask_F0) >> 4) | (x1 & mask_F0));
}
}
// ---------------------------------------------------------------------------
// kernel_moe_expand_scale_q8_0
//
// Expand the q8_0 per-32-block scale d (one half/block, [expert][row][block]) into
// the UNIFORM scale[16] format the generic dp4a MoE GEMM (kernel_gemm_moe_q8_1_dp4a,
// MOE_QT=80) consumes: 16 f16 per 256-superblock (per-16-element segment), where the
// two segments of each 32-block share the block's d. q8_0 is symmetric -> no min
// buffer (the GEMM runs with has_min=0). The int8 weight codes are reused verbatim
// from the existing flat q8_0 weight buffer (extra0_q8_0->q), so only the scale is
// rebuilt here. One work-item per (row, superblock, expert).
// ---------------------------------------------------------------------------
kernel void kernel_moe_expand_scale_q8_0(
global const half * src_d, // [expert][row][block], one scale per 32-block
global half * dst_scale, // [expert][row][block][2] (FLAT per-32-block)
int ne00,
int ne01
) {
int row = get_global_id(0);
int blk = get_global_id(1); // 32-block index along K
int e = get_global_id(2);
if (row >= ne01) { return; }
long nb = ne00 / 32; // 32-blocks per row (K only needs % 32 == 0)
half d = src_d[((long)e*ne01 + row)*nb + blk];
long b = (((long)e*ne01 + row)*nb + blk) * 2;
dst_scale[b + 0] = d;
dst_scale[b + 1] = d;
}
// ---------------------------------------------------------------------------
// kernel_moe_expand_scale_q5_0
//
// q5_0 = symmetric, value = d*(code-16), code = nibble | (hi<<4) in 0..31. The
// generic dp4a MoE GEMM keeps the unsigned code and centers via the min term:
// scale*dp4a(code,a) - min*sum(a), scale = d, min = d*16.
// Reads the existing q5_0 d ([expert][block][row], one half/32-block, from the
// trans4 convert) and writes the FLAT per-32-block uniform scale[2]/min[1] in
// [expert][row][block] order (a transpose). One work-item per (row, block, expert).
// ---------------------------------------------------------------------------
kernel void kernel_moe_expand_scale_q5_0(
global const half * src_d, // [expert][block][row]
global half * dst_scale, // [expert][row][block][2]
global half * dst_min, // [expert][row][block]
int ne00,
int ne01
) {
int row = get_global_id(0);
int blk = get_global_id(1);
int e = get_global_id(2);
if (row >= ne01) { return; }
long nb = ne00 / 32;
half d = src_d[(long)e*nb*ne01 + (long)blk*ne01 + row]; // [expert][block][row]
long sb = (((long)e*ne01 + row)*nb + blk) * 2;
long mb = ((long)e*ne01 + row)*nb + blk;
dst_scale[sb + 0] = d;
dst_scale[sb + 1] = d;
dst_min[mb] = (half)((float)d * 16.0f);
}
// ---------------------------------------------------------------------------
// kernel_moe_expand_scale_q5_K
//
// q5_K value = d*sv*code + (-dm*mn), with the 6-bit packed per-sub-block scale sv
// and min mn (8 sub-blocks of 32 per 256-superblock, decoded by get_scale_min_k4
// from the 12-byte s[]). The generic dp4a MoE GEMM (kernel_gemm_moe_q8_1_dp4a,
// MOE_QT=5) keeps the unsigned 5-bit code and applies scale/min via the uniform
// per-32-block buffers:
// acc += sc0*a_d*raw1 + sc1*a_d*raw2 - mn_u*a_s,
// sc0 = sc1 = d*sv (both per-16 segments of a 32-block share the sub-block scale),
// mn_u = dm*mn (positive; the GEMM subtracts it -> the -dm*mn min term).
// q5_K's q_img (low nibbles) + qh (hi-bit plane) are already in the layout the GEMM
// reads (same trans4_ns convert that feeds gemm_moe_q5_k_f32_ns), so only the scale
// is rebuilt here.
//
// One work-item per (row, superblock, expert); each emits 8 sub-blocks.
// ---------------------------------------------------------------------------
kernel void kernel_moe_expand_scale_q5_K(
global const uchar * src_s, // [expert][row][superblock][12]
global const half * src_d, // [expert][superblock][row]
global const half * src_dm, // [expert][superblock][row]
global half * dst_scale, // [expert][row][32block][2]
global half * dst_min, // [expert][row][32block]
int ne00,
int ne01
) {
int row = get_global_id(0);
int sb = get_global_id(1); // superblock index along K
int e = get_global_id(2);
if (row >= ne01) { return; }
long nsb = ne00 / 256; // superblocks per row
long nblk32 = ne00 / 32; // 32-blocks per row
float d = (float)src_d [((long)e*nsb + sb)*ne01 + row];
float dm = (float)src_dm[((long)e*nsb + sb)*ne01 + row];
__global const uchar * sc = src_s + ((long)e*ne01 + row)*nsb*12 + (long)sb*12;
for (int j = 0; j < 8; ++j) {
uchar sv, mn;
// get_scale_min_k4 (6-bit packed scale/min for sub-block j of 8)
if (j < 4) {
sv = sc[j] & 63;
mn = sc[j+4] & 63;
} else {
sv = (sc[j+4] & 0x0F) | ((sc[j-4] & 0xC0) >> 2);
mn = ((sc[j+4] >> 4) & 0x0F) | ((sc[j] & 0xC0) >> 2);
}
long sub = (long)sb*8 + j;
long sbase = (((long)e*ne01 + row)*nblk32 + sub) * 2;
half s_val = (half)(d * (float)sv);
dst_scale[sbase + 0] = s_val;
dst_scale[sbase + 1] = s_val;
dst_min[((long)e*ne01 + row)*nblk32 + sub] = (half)(dm * (float)mn);
}
}
@@ -0,0 +1,186 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_M 64
#define TILESIZE_N 32
// 2*mxfp4_value as signed int8, packed 4 codes per uint. Divergent nibble
// lookups read a __constant *uint* array + shift, never a byte array
// (byte-indexed __constant loads serialize on Adreno and are far slower).
// idx 0-3: 0, 1, 2, 3 = 0x03020100
// idx 4-7: 4, 6, 8, 12 = 0x0C080604
// idx 8-11: 0, -1, -2, -3 = 0xFDFEFF00 (-1=0xFF,-2=0xFE,-3=0xFD)
// idx 12-15:-4, -6, -8,-12 = 0xF4F8FAFC (-4=0xFC,-6=0xFA,-8=0xF8,-12=0xF4)
__constant uint mxfp4_i8x4[4] = {
0x03020100u, 0x0C080604u, 0xFDFEFF00u, 0xF4F8FAFCu
};
inline uint mxfp4_code(uint n) {
return (mxfp4_i8x4[n >> 2] >> ((n & 3u) * 8u)) & 0xFFu;
}
// 4 nibbles in the low 16 bits of u -> 4 codebook int8, packed for dp4a.
inline uint mxfp4_pack(ushort u) {
return mxfp4_code((uint)( u & 0xF))
| (mxfp4_code((uint)((u >> 4) & 0xF)) << 8)
| (mxfp4_code((uint)((u >> 8) & 0xF)) << 16)
| (mxfp4_code((uint)((u >> 12) & 0xF)) << 24);
}
static inline float e8m0_to_fp32(uchar x) {
int bits;
bits = (x == 0) ? 0x00400000 : ((uint) x << 23);
return as_float(bits);
}
// One token's dp4a dot (8 uints = 32 K elems) + mxfp4 block-scale epilogue.
// blk_scale already carries the 0.5 factor (== 0.5 * 2^e).
#define MOE_MXFP4_DP4A_T(t) do { \
int raw = 0; \
raw = dot_acc_sat_4x8packed_ss_int(qw[0], sh_qa[t][0], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[1], sh_qa[t][1], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[2], sh_qa[t][2], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[3], sh_qa[t][3], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[4], sh_qa[t][4], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[5], sh_qa[t][5], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[6], sh_qa[t][6], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[7], sh_qa[t][7], raw); \
acc[t] += blk_scale * (float)sh_d[t] * (float)raw; \
} while (0)
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_mxfp4_q8_1_dp4a(
__read_only image1d_buffer_t src0_q, // mxfp4 codes (transposed, packed nibbles)
__global uchar * src0_e, // e8m0 per-32-block scale
__global uint * src1_qa, // q8_1 activations: int8 quants (as uint, 4/elem)
__global half * src1_da, // q8_1 per-block scale [tok_slot * ne00/32]
__global uint * src2, // post-router (orig out positions)
__global ushort * src2_emap, // tile -> expert id
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01,
int is_ragged // 1: compute only real tokens per tile
) {
const uint block_id_m = get_global_id(1); // m_tile
const uint block_id_n = get_global_id(2); // n_tile
if (block_id_n >= total_tiles[0]) {
return;
}
const uint lid = get_local_id(0); // 0..63, == this WI's output row in the M-tile
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * TILESIZE_M;
const uint col = block_id_n * TILESIZE_N;
const uint num_blocks = ne00 >> 5; // blocks-of-32 per token
const uint row_idx = row + lid;
const uint ne00_u = ne00 >> 2; // ne00 in uint (int8x4) units
__local uint sh_qa[TILESIZE_N][8]; // 32 tokens x 8 uints (32 int8) = 1 KiB
__local half sh_d[TILESIZE_N];
// Real token count for this tile.
// Real tokens are packed contiguously at the tile start; padded slots hold
// 0xFFFFFFFF (only the last tile of each expert is partial). is_ragged skips
// the dp4a/staging/scatter for padded slots; is_ragged==0 forces n_real=32.
__local uint sh_src2[TILESIZE_N];
__local int sh_nreal;
if (lid < TILESIZE_N) {
sh_src2[lid] = src2[col + lid];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 0) {
int nr = TILESIZE_N;
if (is_ragged) {
nr = 0;
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) {
if (sh_src2[t] != 0xFFFFFFFFu) ++nr;
}
}
sh_nreal = nr;
}
barrier(CLK_LOCAL_MEM_FENCE);
const int n_real = sh_nreal;
float acc[TILESIZE_N];
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) acc[t] = 0.0f;
for (uint step = 0; step < ne00; step += 32) {
const uint sub = step >> 5; // 32-block index along K
// e8m0 block scale for this WI's row, this 32-block (folded x0.5)
const uint e_offset = row_idx + sub * ne01 + expert_id * num_blocks * ne01;
const float blk_scale = 0.5f * e8m0_to_fp32(src0_e[e_offset]);
// repack this WI's 32 weight nibbles into 8 dp4a uints
const uint qoff0 = row + ((ne01 * step) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint qoff1 = row + ((ne01 * (step + 16)) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint r0 = read_imageui(src0_q, qoff0 + lid).x;
const uint r1 = read_imageui(src0_q, qoff0 + lid + ne01).x;
const uint r2 = read_imageui(src0_q, qoff1 + lid).x;
const uint r3 = read_imageui(src0_q, qoff1 + lid + ne01).x;
uint qw[8];
qw[0] = mxfp4_pack((ushort)(r0)); qw[1] = mxfp4_pack((ushort)(r0 >> 16));
qw[2] = mxfp4_pack((ushort)(r1)); qw[3] = mxfp4_pack((ushort)(r1 >> 16));
qw[4] = mxfp4_pack((ushort)(r2)); qw[5] = mxfp4_pack((ushort)(r2 >> 16));
qw[6] = mxfp4_pack((ushort)(r3)); qw[7] = mxfp4_pack((ushort)(r3 >> 16));
// cooperatively stage the n_real-token x 32-K int8 activations
const uint stage_lim = (uint)n_real * 8;
for (uint idx = lid; idx < stage_lim; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
sh_qa[t][u] = src1_qa[(col + t) * ne00_u + (step >> 2) + u];
}
if (lid < (uint)n_real) {
sh_d[lid] = src1_da[(col + lid) * num_blocks + sub];
}
barrier(CLK_LOCAL_MEM_FENCE);
// Full tiles keep the fully-unrolled 32-wide loop; partial tiles run only n_real
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) { MOE_MXFP4_DP4A_T(t); }
} else {
#pragma unroll 4
for (int t = 0; t < n_real; ++t) { MOE_MXFP4_DP4A_T(t); }
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (row_idx >= ne01) {
return;
}
// scatter results to original output rows (reuse sh_src2 from the top)
__local uint out_idx[TILESIZE_N];
if (lid < TILESIZE_N) {
uint idx = sh_src2[lid];
if (idx == 0xFFFFFFFF) {
idx = sh_src2[0];
}
out_idx[lid] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
const uint m_offset = row + lid;
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 1; t < TILESIZE_N; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, acc[0]);
} else {
for (int t = 0; t < n_real; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
}
}
@@ -0,0 +1,165 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_M 64
#define TILESIZE_N 32
// Expand the 4 nibbles held in the low 16 bits of `u` into 4 bytes (one nibble
// per byte, value 0..15), packed for the int8 dp4a. The -8 zero-point is applied
// in the epilogue via the activation sum term (cheaper than biasing every byte).
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// One token's dp4a dot (8 uints = 32 K elems) + q4_0 scale/zero-point epilogue.
#define MOE_Q40_DP4A_T(t) do { \
int raw = 0; \
raw = dot_acc_sat_4x8packed_ss_int(qw[0], sh_qa[t][0], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[1], sh_qa[t][1], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[2], sh_qa[t][2], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[3], sh_qa[t][3], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[4], sh_qa[t][4], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[5], sh_qa[t][5], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[6], sh_qa[t][6], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[7], sh_qa[t][7], raw); \
acc[t] += d_val * ((float)sh_d[t] * (float)raw - 8.0f * (float)sh_s[t]); \
} while (0)
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_q4_0_q8_1_dp4a(
__read_only image1d_buffer_t src0_q, // q4_0 weights (transposed, packed nibbles)
__global half * src0_d, // per-32-block scale
__global uint * src1_qa, // q8_1 activations: int8 quants (as uint, 4/elem)
__global half * src1_da, // q8_1 per-block scale [tok_slot * ne00/32]
__global half * src1_sa, // q8_1 per-block sum*d [tok_slot * ne00/32]
__global uint * src2, // post-router (orig out positions)
__global ushort * src2_emap,// tile -> expert id
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01,
int is_ragged // 1: compute only real tokens per tile
) {
const uint block_id_m = get_global_id(1); // m_tile
const uint block_id_n = get_global_id(2); // n_tile
if (block_id_n >= total_tiles[0]) {
return;
}
const uint lid = get_local_id(0); // 0..63, == this WI's output row in the M-tile
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * TILESIZE_M;
const uint col = block_id_n * TILESIZE_N;
const uint num_blocks = ne00 >> 5; // blocks-of-32 per token
const uint row_idx = row + lid;
const uint ne00_u = ne00 >> 2; // ne00 in uint (int8x4) units
__local uint sh_qa[TILESIZE_N][8]; // 32 tokens x 8 uints (32 int8) = 1 KiB
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
// Real-token count for this tile
__local uint sh_src2[TILESIZE_N];
__local int sh_nreal;
if (lid < TILESIZE_N) {
sh_src2[lid] = src2[col + lid];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 0) {
int nr = TILESIZE_N;
if (is_ragged) {
nr = 0;
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) {
if (sh_src2[t] != 0xFFFFFFFFu) ++nr;
}
}
sh_nreal = nr;
}
barrier(CLK_LOCAL_MEM_FENCE);
const int n_real = sh_nreal;
float acc[TILESIZE_N];
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) acc[t] = 0.0f;
for (uint step = 0; step < ne00; step += 32) {
const uint sub = step >> 5; // 32-block index along K
// per-32-block scale for this WI's row
const uint d_offset = row_idx + sub * ne01 + expert_id * num_blocks * ne01;
const float d_val = (float)src0_d[d_offset];
// repack this WI's 32 weight nibbles into 8 dp4a uints
const uint qoff0 = row + ((ne01 * step) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint qoff1 = row + ((ne01 * (step + 16)) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint r0 = read_imageui(src0_q, qoff0 + lid).x;
const uint r1 = read_imageui(src0_q, qoff0 + lid + ne01).x;
const uint r2 = read_imageui(src0_q, qoff1 + lid).x;
const uint r3 = read_imageui(src0_q, qoff1 + lid + ne01).x;
uint qw[8];
qw[0] = EXP4(r0); qw[1] = EXP4(r0 >> 16);
qw[2] = EXP4(r1); qw[3] = EXP4(r1 >> 16);
qw[4] = EXP4(r2); qw[5] = EXP4(r2 >> 16);
qw[6] = EXP4(r3); qw[7] = EXP4(r3 >> 16);
// cooperatively stage the n_real-token x 32-K int8 activations
const uint stage_lim = (uint)n_real * 8;
for (uint idx = lid; idx < stage_lim; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
sh_qa[t][u] = src1_qa[(col + t) * ne00_u + (step >> 2) + u];
}
if (lid < (uint)n_real) {
sh_d[lid] = src1_da[(col + lid) * num_blocks + sub];
sh_s[lid] = src1_sa[(col + lid) * num_blocks + sub];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) { MOE_Q40_DP4A_T(t); }
} else {
#pragma unroll 4
for (int t = 0; t < n_real; ++t) { MOE_Q40_DP4A_T(t); }
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (row_idx >= ne01) {
return;
}
// scatter results to original output rows (reuse sh_src2 from the top)
__local uint out_idx[TILESIZE_N];
if (lid < TILESIZE_N) {
uint idx = sh_src2[lid];
if (idx == 0xFFFFFFFF) {
idx = sh_src2[0];
}
out_idx[lid] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
const uint m_offset = row + lid;
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 1; t < TILESIZE_N; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, acc[0]);
} else {
for (int t = 0; t < n_real; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
}
}
@@ -0,0 +1,202 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
// q4_K subblock (32 elems): w_i = scale*q_i - minv, q_i in [0,15], scale =
// d_super*sv6, minv = dmin_super*mn6. With activation block (a_d, a_s, qa[32]):
// Sum_i w_i * a_i = scale * a_d * dp4a(q, qa) - minv * a_s
// where a_s = a_d * Sum(qa) (the q8_1 "s" field)
#define TILESIZE_M 64
#define TILESIZE_N 32
#define QK_K 256
#define K_SCALE_SIZE 12
inline void get_scale_min_k4(
int j,
global const uchar * q,
uchar * d,
uchar * m
) {
if (j < 4) {
*d = q[j] & 63;
*m = q[j+4] & 63;
} else {
*d = (q[j+4] & 0x0F) | ((q[j-4] & 0xC0) >> 2);
*m = ((q[j+4] >> 4) & 0x0F) | ((q[j] & 0xC0) >> 2);
}
}
// Expand the 4 nibbles held in the low 16 bits of `u` into 4 bytes (one nibble
// per byte, value 0..15), packed for the int8 dp4a.
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// One token's dp4a dot (8 uints = 32 K elems) + q4_K scale/min epilogue into acc[t].
#define MOE_Q4K_DP4A_T(t) do { \
int raw = 0; \
raw = dot_acc_sat_4x8packed_ss_int(qw[0], sh_qa[t][0], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[1], sh_qa[t][1], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[2], sh_qa[t][2], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[3], sh_qa[t][3], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[4], sh_qa[t][4], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[5], sh_qa[t][5], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[6], sh_qa[t][6], raw); \
raw = dot_acc_sat_4x8packed_ss_int(qw[7], sh_qa[t][7], raw); \
acc[t] += scale * (float)sh_d[t] * (float)raw - minv * (float)sh_s[t]; \
} while (0)
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_q4_k_q8_1_dp4a(
__read_only image1d_buffer_t src0_q, // q4_K weights (transposed, packed nibbles)
__global half * src0_d, // per-superblock scale
__global half * src0_dm, // per-superblock min
__global uchar * src0_s, // 6-bit scale/min codes
__global uint * src1_qa, // q8_1 activations: int8 quants (as uint, 4/elem)
__global half * src1_da, // q8_1 per-block scale [tok_slot * ne00/32]
__global half * src1_sa, // q8_1 per-block sum*d [tok_slot * ne00/32]
__global uint * src2, // post-router (orig out positions)
__global ushort * src2_emap,// tile -> expert id
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01,
int is_ragged // 1: compute only real tokens per tile
) {
const uint block_id_m = get_global_id(1); // m_tile
const uint block_id_n = get_global_id(2); // n_tile
if (block_id_n >= total_tiles[0]) {
return;
}
const uint lid = get_local_id(0); // 0..63, == this WI's output row in the M-tile
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * TILESIZE_M;
const uint col = block_id_n * TILESIZE_N;
const uint num_superblocks = ne00 / QK_K;
const uint scales_per_row = num_superblocks * K_SCALE_SIZE;
const uint row_idx = row + lid;
const uint ne00_u = ne00 >> 2; // ne00 in uint (int8x4) units
const uint ne00_b = ne00 >> 5; // blocks-of-32 per token
__local uint sh_qa[TILESIZE_N][8]; // 32 tokens x 8 uints (32 int8) = 1 KiB
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
// Real token count for this tile
__local uint sh_src2[TILESIZE_N];
__local int sh_nreal;
if (lid < TILESIZE_N) {
sh_src2[lid] = src2[col + lid];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 0) {
int nr = TILESIZE_N;
if (is_ragged) {
nr = 0;
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) {
if (sh_src2[t] != 0xFFFFFFFFu) ++nr;
}
}
sh_nreal = nr;
}
barrier(CLK_LOCAL_MEM_FENCE);
const int n_real = sh_nreal;
float acc[TILESIZE_N];
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) acc[t] = 0.0f;
for (uint step = 0; step < ne00; step += 32) {
const uint sub = step >> 5; // subblock index along K
const uint sb = sub >> 3; // superblock index
const uint j = sub & 7; // subblock within superblock
// --- weight scale / min for this WI's row, this subblock ---
const uint d_offset = row + sb * ne01 + expert_id * num_superblocks * ne01 + lid;
const float d_val = (float)src0_d[d_offset];
const float dm_val = (float)src0_dm[d_offset];
global const uchar * sc = src0_s + (expert_id * ne01 + row_idx) * scales_per_row + sb * K_SCALE_SIZE;
uchar sv, mn;
get_scale_min_k4(j, sc, &sv, &mn);
const float scale = d_val * (float)sv;
const float minv = dm_val * (float)mn;
// --- repack this WI's 32 weight nibbles into 8 dp4a uints ---
const uint qoff0 = row + ((ne01 * step) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint qoff1 = row + ((ne01 * (step + 16)) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint r0 = read_imageui(src0_q, qoff0 + lid).x;
const uint r1 = read_imageui(src0_q, qoff0 + lid + ne01).x;
const uint r2 = read_imageui(src0_q, qoff1 + lid).x;
const uint r3 = read_imageui(src0_q, qoff1 + lid + ne01).x;
uint qw[8];
qw[0] = EXP4(r0); qw[1] = EXP4(r0 >> 16);
qw[2] = EXP4(r1); qw[3] = EXP4(r1 >> 16);
qw[4] = EXP4(r2); qw[5] = EXP4(r2 >> 16);
qw[6] = EXP4(r3); qw[7] = EXP4(r3 >> 16);
// --- cooperatively stage the n_real-token x 32-K int8 activations to LDS ---
const uint stage_lim = (uint)n_real * 8;
for (uint idx = lid; idx < stage_lim; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
sh_qa[t][u] = src1_qa[(col + t) * ne00_u + (step >> 2) + u];
}
if (lid < (uint)n_real) {
sh_d[lid] = src1_da[(col + lid) * ne00_b + sub];
sh_s[lid] = src1_sa[(col + lid) * ne00_b + sub];
}
barrier(CLK_LOCAL_MEM_FENCE);
// dp4a - each real token sum over 8 uints (32 K), then scale/min
// Full tiles keep the fully-unrolled 32-wide loop;
// partial tiles run only n_real (saves the padded-slot dp4a + staging).
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) { MOE_Q4K_DP4A_T(t); }
} else {
#pragma unroll 4
for (int t = 0; t < n_real; ++t) { MOE_Q4K_DP4A_T(t); }
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (row_idx >= ne01) {
return;
}
// scatter results to original output rows
__local uint out_idx[TILESIZE_N];
if (lid < TILESIZE_N) {
uint idx = sh_src2[lid];
if (idx == 0xFFFFFFFF) {
idx = sh_src2[0];
}
out_idx[lid] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
const uint m_offset = row + lid;
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 1; t < TILESIZE_N; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, acc[0]);
} else {
for (int t = 0; t < n_real; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
}
}
@@ -0,0 +1,196 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_N 32
#define QK_K 256
// 4 nibbles in the low 16 bits of `u` -> 4 bytes (value 0..15, in bits 0-3).
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 4 2-bit highs in byte `b` (8 bits) -> 4 bytes, value 0..3 in bits 4-5
// (pre-multiplied by 16 so it ORs with the EXP4 nibble to form q6 in 0..63).
#define EXP2(b) ( (((uint)((b) & 0x03u)) << 4) | \
(((uint)((b) & 0x0Cu)) << 10) | \
(((uint)((b) & 0x30u)) << 16) | \
(((uint)((b) & 0xC0u)) << 22) )
// q6 (0..63, bits 0-5 of each byte) -> (q6-32) as a signed int8 per byte.
// Flipping bit5 subtracts 32 in 6-bit two's complement; then replicate bit5
// into bits 6-7 to sign-extend to int8. Per-byte, no inter-byte carry.
inline uint SIGN6(uint q6p) {
uint x = q6p ^ 0x20202020u;
uint s = x & 0x20202020u;
return x | (s << 1) | (s << 2);
}
inline int dp4a_q6(uint qw0, uint qw1, uint qw2, uint qw3,
uint a0, uint a1, uint a2, uint a3) {
int raw = 0;
raw = dot_acc_sat_4x8packed_ss_int(qw0, a0, raw);
raw = dot_acc_sat_4x8packed_ss_int(qw1, a1, raw);
raw = dot_acc_sat_4x8packed_ss_int(qw2, a2, raw);
raw = dot_acc_sat_4x8packed_ss_int(qw3, a3, raw);
return raw;
}
// One token's q6_K dp4a dot (two halves, per-16 scales) + epilogue into acc[t].
#define MOE_Q6K_DP4A_T(t) do { \
const int raw1 = dp4a_q6(qw[0], qw[1], qw[2], qw[3], sh_qa[t][0], sh_qa[t][1], sh_qa[t][2], sh_qa[t][3]); \
const int raw2 = dp4a_q6(qw[4], qw[5], qw[6], qw[7], sh_qa[t][4], sh_qa[t][5], sh_qa[t][6], sh_qa[t][7]); \
const float a_d = (float)sh_d[t]; \
acc[t] += scale0 * a_d * (float)raw1 + scale1 * a_d * (float)raw2; \
} while (0)
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_q6_k_q8_1_dp4a(
__read_only image1d_buffer_t src0_ql, // q6_K low nibbles (image, q4_K-style layout)
__global uint * src0_qh, // q6_K high 2-bit (16 elems/uint)
__global char * src0_s, // int8 scales (one per 16 elems)
__global half * src0_d, // per-superblock scale
__global uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem)
__global half * src1_da, // q8_1 per-block scale [tok_slot * ne00/32]
__global uint * src2, // post-router (orig out positions)
__global ushort * src2_emap, // tile -> expert id
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01,
int is_ragged // 1: compute only real tokens per tile
) {
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
if (block_id_n >= total_tiles[0]) {
return;
}
const uint lid = get_local_id(0); // 0..63 -> row within M-tile
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * 64;
const uint col = block_id_n * TILESIZE_N;
const uint num_superblocks = ne00 / QK_K;
const uint scales_per_row = num_superblocks * 16;
const uint row_idx = row + lid;
const uint ne00_u = ne00 >> 2;
const uint ne00_b = ne00 >> 5;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
// Real token count for this tile
__local uint sh_src2[TILESIZE_N];
__local int sh_nreal;
if (lid < TILESIZE_N) {
sh_src2[lid] = src2[col + lid];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 0) {
int nr = TILESIZE_N;
if (is_ragged) {
nr = 0;
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) {
if (sh_src2[t] != 0xFFFFFFFFu) ++nr;
}
}
sh_nreal = nr;
}
barrier(CLK_LOCAL_MEM_FENCE);
const int n_real = sh_nreal;
float acc[TILESIZE_N];
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) acc[t] = 0.0f;
for (uint step = 0; step < ne00; step += 32) {
const uint sub = step >> 5;
const uint sb = sub >> 3;
const uint j = sub & 7;
const float d_val = (float)src0_d[row + sb * ne01 + expert_id * num_superblocks * ne01 + lid];
global const char * sc = src0_s + (expert_id * ne01 + row_idx) * scales_per_row + sb * 16;
const float scale0 = d_val * (float)sc[j * 2];
const float scale1 = d_val * (float)sc[j * 2 + 1];
// high bits: one uint covers 16 elems; first/second 16 of this 32-block
const uint qh_base = row + (sub * 2) * ne01 + expert_id * (num_superblocks * 16) * ne01 + lid;
const uint qh1 = src0_qh[qh_base];
const uint qh2 = src0_qh[qh_base + ne01];
// low nibbles: same image layout as q4_K (8 ushorts over the 32 K)
const uint qoff0 = row + ((ne01 * step) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint qoff1 = row + ((ne01 * (step + 16)) >> 3) + ((expert_id * ne00 * ne01) >> 3);
const uint r0 = read_imageui(src0_ql, qoff0 + lid).x;
const uint r1 = read_imageui(src0_ql, qoff0 + lid + ne01).x;
const uint r2 = read_imageui(src0_ql, qoff1 + lid).x;
const uint r3 = read_imageui(src0_ql, qoff1 + lid + ne01).x;
uint qw[8];
qw[0] = SIGN6(EXP4(r0) | EXP2((qh1) & 0xFFu));
qw[1] = SIGN6(EXP4(r0 >> 16) | EXP2((qh1 >> 8) & 0xFFu));
qw[2] = SIGN6(EXP4(r1) | EXP2((qh1 >> 16) & 0xFFu));
qw[3] = SIGN6(EXP4(r1 >> 16) | EXP2((qh1 >> 24) & 0xFFu));
qw[4] = SIGN6(EXP4(r2) | EXP2((qh2) & 0xFFu));
qw[5] = SIGN6(EXP4(r2 >> 16) | EXP2((qh2 >> 8) & 0xFFu));
qw[6] = SIGN6(EXP4(r3) | EXP2((qh2 >> 16) & 0xFFu));
qw[7] = SIGN6(EXP4(r3 >> 16) | EXP2((qh2 >> 24) & 0xFFu));
const uint stage_lim = (uint)n_real * 8;
for (uint idx = lid; idx < stage_lim; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
sh_qa[t][u] = src1_qa[(col + t) * ne00_u + (step >> 2) + u];
}
if (lid < (uint)n_real) {
sh_d[lid] = src1_da[(col + lid) * ne00_b + sub];
}
barrier(CLK_LOCAL_MEM_FENCE);
// Full tiles keep the fully-unrolled 32-wide loop; partial tiles run n_real.
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) { MOE_Q6K_DP4A_T(t); }
} else {
#pragma unroll 4
for (int t = 0; t < n_real; ++t) { MOE_Q6K_DP4A_T(t); }
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (row_idx >= ne01) {
return;
}
__local uint out_idx[TILESIZE_N];
if (lid < TILESIZE_N) {
uint idx = sh_src2[lid];
if (idx == 0xFFFFFFFF) {
idx = sh_src2[0];
}
out_idx[lid] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
const uint m_offset = row + lid;
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 1; t < TILESIZE_N; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, acc[0]);
} else {
for (int t = 0; t < n_real; ++t) {
write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
}
}
@@ -0,0 +1,221 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#pragma OPENCL EXTENSION cl_qcom_subgroup_uniform_load: enable
#pragma OPENCL EXTENSION cl_qcom_subgroup_constant_load: enable
#pragma OPENCL EXTENSION cl_qcom_extra_vector_types : enable
#define TILESIZE_K 16
#define TILESIZE_M 64
#define TILESIZE_N 32
// q8_0: 16 signed int8 weights (one uint4 = 16 chars) -> half16, scaled.
#define dequantize_q8_0(q4, a_f16, scale) \
a_f16 = convert_half16(as_char16(q4)) * scale;
#define dotx16_reduce8(a_reg, b_lm, c_reg, lm_offset) \
acc.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
acc.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
acc.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
acc.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
acc.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
acc.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
acc.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
acc.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
acc.s8 = dot(a_reg.s0123, b_lm[lm_offset + 8]); \
acc.s9 = dot(a_reg.s0123, b_lm[lm_offset + 9]); \
acc.sa = dot(a_reg.s0123, b_lm[lm_offset + 10]); \
acc.sb = dot(a_reg.s0123, b_lm[lm_offset + 11]); \
acc.sc = dot(a_reg.s0123, b_lm[lm_offset + 12]); \
acc.sd = dot(a_reg.s0123, b_lm[lm_offset + 13]); \
acc.se = dot(a_reg.s0123, b_lm[lm_offset + 14]); \
acc.sf = dot(a_reg.s0123, b_lm[lm_offset + 15]); \
acc.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
acc.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
acc.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
acc.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
acc.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
acc.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
acc.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
acc.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
acc.s8 += dot(a_reg.s4567, b_lm[lm_offset + 40]); \
acc.s9 += dot(a_reg.s4567, b_lm[lm_offset + 41]); \
acc.sa += dot(a_reg.s4567, b_lm[lm_offset + 42]); \
acc.sb += dot(a_reg.s4567, b_lm[lm_offset + 43]); \
acc.sc += dot(a_reg.s4567, b_lm[lm_offset + 44]); \
acc.sd += dot(a_reg.s4567, b_lm[lm_offset + 45]); \
acc.se += dot(a_reg.s4567, b_lm[lm_offset + 46]); \
acc.sf += dot(a_reg.s4567, b_lm[lm_offset + 47]); \
c_reg.lo += convert_float8(acc.lo); \
c_reg.hi += convert_float8(acc.hi); \
acc.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
acc.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
acc.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
acc.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
acc.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
acc.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
acc.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
acc.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
acc.s8 = dot(a_reg.s89ab, b_lm[lm_offset + 72]); \
acc.s9 = dot(a_reg.s89ab, b_lm[lm_offset + 73]); \
acc.sa = dot(a_reg.s89ab, b_lm[lm_offset + 74]); \
acc.sb = dot(a_reg.s89ab, b_lm[lm_offset + 75]); \
acc.sc = dot(a_reg.s89ab, b_lm[lm_offset + 76]); \
acc.sd = dot(a_reg.s89ab, b_lm[lm_offset + 77]); \
acc.se = dot(a_reg.s89ab, b_lm[lm_offset + 78]); \
acc.sf = dot(a_reg.s89ab, b_lm[lm_offset + 79]); \
acc.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
acc.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
acc.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
acc.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
acc.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
acc.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
acc.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
acc.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
acc.s8 += dot(a_reg.scdef, b_lm[lm_offset + 104]); \
acc.s9 += dot(a_reg.scdef, b_lm[lm_offset + 105]); \
acc.sa += dot(a_reg.scdef, b_lm[lm_offset + 106]); \
acc.sb += dot(a_reg.scdef, b_lm[lm_offset + 107]); \
acc.sc += dot(a_reg.scdef, b_lm[lm_offset + 108]); \
acc.sd += dot(a_reg.scdef, b_lm[lm_offset + 109]); \
acc.se += dot(a_reg.scdef, b_lm[lm_offset + 110]); \
acc.sf += dot(a_reg.scdef, b_lm[lm_offset + 111]); \
c_reg.lo += convert_float8(acc.lo); \
c_reg.hi += convert_float8(acc.hi); \
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_q8_0_f32_ns(
__global char * src0_q, // flat q8_0 quants [n_expert*ne01*ne00]
__global half * src0_d, // flat q8_0 scales [n_expert*ne01*nb]
__read_only image1d_buffer_t src1, // reordered activations (f32)
__global uint * src2, // post-router out indices
__global ushort * src2_emap,// expert per tile
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01
) {
uint block_id_m = get_global_id(1); // m_tile
uint block_id_n = get_global_id(2); // n_tile
if (block_id_n >= total_tiles[0]) {
return;
}
__private half16 reg_a;
__private float32 reg_c = (float32)(0);
__local half4 shared_b[128];
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * TILESIZE_M;
const uint col = block_id_n * TILESIZE_N;
const uint nb = ne00 >> 5; // blocks per row (ne00/32)
const uint w_row = expert_id * ne01 + row + get_local_id(0); // this lane's output row
__global char * w_q = src0_q + (ulong)w_row * ne00; // char base for the row
__global half * w_d = src0_d + (ulong)w_row * nb; // scale base for the row
uint sub_block_id_m = get_local_id(0);
uint2 b_global_offset;
b_global_offset.x = ((sub_block_id_m & 3) << 2) + (sub_block_id_m >> 2) * ne00;
b_global_offset.y = b_global_offset.x + (16 * ne00);
uint2 b_local_offset;
b_local_offset.x = (sub_block_id_m & 3) * 32 + (sub_block_id_m >> 2);
b_local_offset.y = b_local_offset.x + 16;
// Loop along K axis, 32 elements per iteration, split into 2 sub-blocks.
for (uint step = 0; step < ne00; step += TILESIZE_K * 2) {
half s = w_d[step >> 5]; // one q8_0 scale per 32-element block
// First sub-block: 16 weights (16 chars = one uint4) at K=step
uint4 q8x16 = *((__global uint4 *)(w_q + step));
uint b_sub_offset = col * ne00 + step;
float8 bx8_f32;
bx8_f32.lo = read_imagef(src1, (b_sub_offset + b_global_offset.x) / 4);
bx8_f32.hi = read_imagef(src1, (b_sub_offset + b_global_offset.y) / 4);
half8 bx8_f16 = convert_half8(bx8_f32);
shared_b[b_local_offset.x] = bx8_f16.lo;
shared_b[b_local_offset.y] = bx8_f16.hi;
dequantize_q8_0(q8x16, reg_a, s);
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
half16 acc;
dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
// Second sub-block: next 16 weights at K=step+16
uint half_step = step + TILESIZE_K;
q8x16 = *((__global uint4 *)(w_q + half_step));
b_sub_offset = col * ne00 + half_step;
bx8_f32.lo = read_imagef(src1, (b_sub_offset + b_global_offset.x) / 4);
bx8_f32.hi = read_imagef(src1, (b_sub_offset + b_global_offset.y) / 4);
bx8_f16 = convert_half8(bx8_f32);
shared_b[b_local_offset.x] = bx8_f16.lo;
shared_b[b_local_offset.y] = bx8_f16.hi;
dequantize_q8_0(q8x16, reg_a, s);
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
}
if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
return;
}
__local uint out_idx[TILESIZE_N];
if (get_local_id(0) < TILESIZE_N) {
uint idx = src2[block_id_n * TILESIZE_N + get_local_id(0)];
if (idx == 0xFFFFFFFF) {
idx = src2[block_id_n * TILESIZE_N + 0];
}
out_idx[get_local_id(0)] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
uint m_offset = row + get_local_id(0);
write_imagef(dst, out_idx[1] + m_offset, (reg_c.s1));
write_imagef(dst, out_idx[2] + m_offset, (reg_c.s2));
write_imagef(dst, out_idx[3] + m_offset, (reg_c.s3));
write_imagef(dst, out_idx[4] + m_offset, (reg_c.s4));
write_imagef(dst, out_idx[5] + m_offset, (reg_c.s5));
write_imagef(dst, out_idx[6] + m_offset, (reg_c.s6));
write_imagef(dst, out_idx[7] + m_offset, (reg_c.s7));
write_imagef(dst, out_idx[8] + m_offset, (reg_c.s8));
write_imagef(dst, out_idx[9] + m_offset, (reg_c.s9));
write_imagef(dst, out_idx[10] + m_offset, (reg_c.sa));
write_imagef(dst, out_idx[11] + m_offset, (reg_c.sb));
write_imagef(dst, out_idx[12] + m_offset, (reg_c.sc));
write_imagef(dst, out_idx[13] + m_offset, (reg_c.sd));
write_imagef(dst, out_idx[14] + m_offset, (reg_c.se));
write_imagef(dst, out_idx[15] + m_offset, (reg_c.sf));
write_imagef(dst, out_idx[16] + m_offset, (reg_c.sg));
write_imagef(dst, out_idx[17] + m_offset, (reg_c.sh));
write_imagef(dst, out_idx[18] + m_offset, (reg_c.si));
write_imagef(dst, out_idx[19] + m_offset, (reg_c.sj));
write_imagef(dst, out_idx[20] + m_offset, (reg_c.sk));
write_imagef(dst, out_idx[21] + m_offset, (reg_c.sl));
write_imagef(dst, out_idx[22] + m_offset, (reg_c.sm));
write_imagef(dst, out_idx[23] + m_offset, (reg_c.sn));
write_imagef(dst, out_idx[24] + m_offset, (reg_c.so));
write_imagef(dst, out_idx[25] + m_offset, (reg_c.sp));
write_imagef(dst, out_idx[26] + m_offset, (reg_c.sq));
write_imagef(dst, out_idx[27] + m_offset, (reg_c.sr));
write_imagef(dst, out_idx[28] + m_offset, (reg_c.ss));
write_imagef(dst, out_idx[29] + m_offset, (reg_c.st));
write_imagef(dst, out_idx[30] + m_offset, (reg_c.su));
write_imagef(dst, out_idx[31] + m_offset, (reg_c.sv));
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, (reg_c.s0));
}
@@ -0,0 +1,221 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
// Generic int8 dp4a MoE GEMM, specialized versions also exist
// MOE_QT:
// 4 (q4_K)/41(q4_1)/40(q4_0) NIBBLE image low nibbles -> EXP4
// 5 (q5_K)/51(q5_1)/50(q5_0) NIBBLE+HI image nibbles + qh high-bit plane
// 6 (q6_K) Q6 image nibbles + qh 2-bit -> SIGN6((nibble|hi2))
// 80(q8_0)/82(mxfp4) INT8 global int8 codes (mxfp4: convert applies kvalues LUT)
#define TILESIZE_M 64
#define TILESIZE_N 32
#define QK_K 256
#ifndef MOE_QT
#define MOE_QT 4
#endif
// 4 nibbles in low 16 bits of u -> 4 bytes (value 0..15)
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 4 2-bit highs in byte b -> 4 bytes, bits 4-5 (q6_K)
#define EXP2(b) ( (((uint)((b) & 0x03u)) << 4) | \
(((uint)((b) & 0x0Cu)) << 10) | \
(((uint)((b) & 0x30u)) << 16) | \
(((uint)((b) & 0xC0u)) << 22) )
// q6 (0..63) -> (q6-32) signed int8/byte (no inter-byte carry)
inline uint SIGN6(uint q6p){ uint x=q6p^0x20202020u; uint s=x&0x20202020u; return x|(s<<1)|(s<<2); }
// 4 high bits (one per element, in bits 0..3 of h) -> bit4 of each of 4 bytes (5-bit hi)
#define EXP1(h) ( (((uint)((h) & 0x1u)) << 4) | \
(((uint)((h) & 0x2u)) << 11) | \
(((uint)((h) & 0x4u)) << 18) | \
(((uint)((h) & 0x8u)) << 25) )
// per-type weight params + per-32-step unpack into qw[8] (8 int8 uints)
#if MOE_QT == 4 || MOE_QT == 41 || MOE_QT == 40
#define WEIGHT_PARAMS __read_only image1d_buffer_t src0_q,
#define LOAD_QW(step, sub) \
uint qw[8]; { \
const uint qoff0 = row + ((ne01*(step))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint qoff1 = row + ((ne01*((step)+16))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint r0=read_imageui(src0_q,qoff0+lid).x, r1=read_imageui(src0_q,qoff0+lid+ne01).x; \
const uint r2=read_imageui(src0_q,qoff1+lid).x, r3=read_imageui(src0_q,qoff1+lid+ne01).x; \
qw[0]=EXP4(r0); qw[1]=EXP4(r0>>16); qw[2]=EXP4(r1); qw[3]=EXP4(r1>>16); \
qw[4]=EXP4(r2); qw[5]=EXP4(r2>>16); qw[6]=EXP4(r3); qw[7]=EXP4(r3>>16); }
#elif MOE_QT == 5 || MOE_QT == 51 || MOE_QT == 50
// low nibbles via image (q4_K layout) + high-bit plane src0_qh: 1 uint per 32-block
// (bit i = high bit of element i). qh laid out [expert][block][row] to match the
// existing q5_0 trans4 convert
#define WEIGHT_PARAMS __read_only image1d_buffer_t src0_q, __global uint * src0_qh,
#define LOAD_QW(step, sub) \
uint qw[8]; { \
const uint qoff0 = row + ((ne01*(step))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint qoff1 = row + ((ne01*((step)+16))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint r0=read_imageui(src0_q,qoff0+lid).x, r1=read_imageui(src0_q,qoff0+lid+ne01).x; \
const uint r2=read_imageui(src0_q,qoff1+lid).x, r3=read_imageui(src0_q,qoff1+lid+ne01).x; \
const uint h = src0_qh[row_idx + (sub)*ne01 + expert_id*(ne00>>5)*ne01]; \
qw[0]=EXP4(r0)|EXP1(h); qw[1]=EXP4(r0>>16)|EXP1(h>>4); \
qw[2]=EXP4(r1)|EXP1(h>>8); qw[3]=EXP4(r1>>16)|EXP1(h>>12); \
qw[4]=EXP4(r2)|EXP1(h>>16); qw[5]=EXP4(r2>>16)|EXP1(h>>20); \
qw[6]=EXP4(r3)|EXP1(h>>24); qw[7]=EXP4(r3>>16)|EXP1(h>>28); }
#elif MOE_QT == 6
#define WEIGHT_PARAMS __read_only image1d_buffer_t src0_ql, __global uint * src0_qh,
#define LOAD_QW(step, sub) \
uint qw[8]; { \
const uint qoff0 = row + ((ne01*(step))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint qoff1 = row + ((ne01*((step)+16))>>3) + ((expert_id*ne00*ne01)>>3); \
const uint r0=read_imageui(src0_ql,qoff0+lid).x, r1=read_imageui(src0_ql,qoff0+lid+ne01).x; \
const uint r2=read_imageui(src0_ql,qoff1+lid).x, r3=read_imageui(src0_ql,qoff1+lid+ne01).x; \
const uint qhb = row + ((sub)*2)*ne01 + expert_id*((ne00>>5)*2)*ne01 + lid; \
const uint qh1=src0_qh[qhb], qh2=src0_qh[qhb+ne01]; \
qw[0]=SIGN6(EXP4(r0)|EXP2(qh1&0xFFu)); qw[1]=SIGN6(EXP4(r0>>16)|EXP2((qh1>>8)&0xFFu)); \
qw[2]=SIGN6(EXP4(r1)|EXP2((qh1>>16)&0xFFu)); qw[3]=SIGN6(EXP4(r1>>16)|EXP2((qh1>>24)&0xFFu)); \
qw[4]=SIGN6(EXP4(r2)|EXP2(qh2&0xFFu)); qw[5]=SIGN6(EXP4(r2>>16)|EXP2((qh2>>8)&0xFFu)); \
qw[6]=SIGN6(EXP4(r3)|EXP2((qh2>>16)&0xFFu)); qw[7]=SIGN6(EXP4(r3>>16)|EXP2((qh2>>24)&0xFFu)); }
#elif MOE_QT == 80 || MOE_QT == 82
// 8-bit direct: int8 codes 8 uints / 32-block, [expert][row][8*sub]. mxfp4: the
// convert resolves kvalues_mxfp4[nibble] -> int8 and stores the e8m0_half scale.
#define WEIGHT_PARAMS __global uint * src0_q8,
#define LOAD_QW(step, sub) \
uint qw[8]; { \
const uint qb = (expert_id*ne01 + row_idx)*(ne00>>2) + (sub)*8; \
qw[0]=src0_q8[qb+0]; qw[1]=src0_q8[qb+1]; qw[2]=src0_q8[qb+2]; qw[3]=src0_q8[qb+3]; \
qw[4]=src0_q8[qb+4]; qw[5]=src0_q8[qb+5]; qw[6]=src0_q8[qb+6]; qw[7]=src0_q8[qb+7]; }
#else
#error "unknown MOE_QT"
#endif
inline int dp4a4(uint w0,uint w1,uint w2,uint w3,uint a0,uint a1,uint a2,uint a3){
int r=0; r=dot_acc_sat_4x8packed_ss_int(w0,a0,r); r=dot_acc_sat_4x8packed_ss_int(w1,a1,r);
r=dot_acc_sat_4x8packed_ss_int(w2,a2,r); r=dot_acc_sat_4x8packed_ss_int(w3,a3,r); return r; }
// One token's two-half dp4a + uniform scale/min epilogue into acc[t].
#define MOE_DP4A_T(t) do { \
const int raw1 = dp4a4(qw[0],qw[1],qw[2],qw[3], sh_qa[t][0],sh_qa[t][1],sh_qa[t][2],sh_qa[t][3]); \
const int raw2 = dp4a4(qw[4],qw[5],qw[6],qw[7], sh_qa[t][4],sh_qa[t][5],sh_qa[t][6],sh_qa[t][7]); \
const float a_d = (float)sh_d[t]; \
acc[t] += sc0*a_d*(float)raw1 + sc1*a_d*(float)raw2 - mn*(float)sh_s[t]; \
} while (0)
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_moe_q8_1_dp4a(
WEIGHT_PARAMS // per-type native weight buffer(s)
__global half * src0_scale,// uniform f16 16/superblock (per-16), [expert,row]
__global half * src0_min, // uniform f16 8/superblock (per-32), [expert,row]
__global uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem)
__global half * src1_da, // q8_1 per-block scale [tok_slot * ne00/32]
__global half * src1_sa, // q8_1 per-block sum*d [tok_slot * ne00/32]
__global uint * src2, // post-router (orig out positions)
__global ushort * src2_emap, // tile -> expert id
__write_only image1d_buffer_t dst,
__global int * total_tiles,
uint ne00,
uint ne01,
int is_ragged,
int has_min // 0 for symmetric types (q8_0/q6_K/q4_0/...): skip min read
) {
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
if (block_id_n >= total_tiles[0]) return;
const uint lid = get_local_id(0); // 0..63 -> output row within M-tile
const ushort expert_id = src2_emap[block_id_n];
const uint row = block_id_m * TILESIZE_M;
const uint col = block_id_n * TILESIZE_N;
const uint row_idx = row + lid;
// Scale/min are laid out FLAT per-32-block (2 per-16-segment scales + 1 min per
// 32-block), so K only needs to be a multiple of 32 works for the 32-block
// types (q8_0/q5_0/q4_0/...) as well as the K-quants (K%256==0, same bytes).
const uint nblk32 = ne00 / 32;
const uint sc_per_row = nblk32 * 2;
const uint mn_per_row = nblk32;
const uint ne00_u = ne00 >> 2;
const uint ne00_b = ne00 >> 5;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
__local uint sh_src2[TILESIZE_N];
__local int sh_nreal;
if (lid < TILESIZE_N) sh_src2[lid] = src2[col + lid];
barrier(CLK_LOCAL_MEM_FENCE);
if (lid == 0) {
int nr = TILESIZE_N;
if (is_ragged) { nr = 0;
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) if (sh_src2[t] != 0xFFFFFFFFu) ++nr; }
sh_nreal = nr;
}
barrier(CLK_LOCAL_MEM_FENCE);
const int n_real = sh_nreal;
float acc[TILESIZE_N];
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) acc[t] = 0.0f;
for (uint step = 0; step < ne00; step += 32) {
const uint sub = step >> 5; // 32-block index along K
// uniform pre-decoded scale (2 per-16-seg) + min (1) for this row, this 32-block
__global half * scl = src0_scale + (expert_id*ne01 + row_idx)*sc_per_row + sub*2;
const float sc0 = (float)scl[0];
const float sc1 = (float)scl[1];
float mn = 0.0f;
if (has_min) mn = (float)src0_min[(expert_id*ne01 + row_idx)*mn_per_row + sub];
LOAD_QW(step, sub)
const uint stage_lim = (uint)n_real * 8;
for (uint idx = lid; idx < stage_lim; idx += 64) {
const uint t = idx >> 3, u = idx & 7;
sh_qa[t][u] = src1_qa[(col + t) * ne00_u + (step >> 2) + u];
}
if (lid < (uint)n_real) {
sh_d[lid] = src1_da[(col + lid) * ne00_b + sub];
sh_s[lid] = src1_sa[(col + lid) * ne00_b + sub];
}
barrier(CLK_LOCAL_MEM_FENCE);
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 0; t < TILESIZE_N; ++t) { MOE_DP4A_T(t); }
} else {
#pragma unroll 4
for (int t = 0; t < n_real; ++t) { MOE_DP4A_T(t); }
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (row_idx >= ne01) return;
__local uint out_idx[TILESIZE_N];
if (lid < TILESIZE_N) {
uint idx = sh_src2[lid];
if (idx == 0xFFFFFFFF) idx = sh_src2[0];
out_idx[lid] = idx * ne01;
}
barrier(CLK_LOCAL_MEM_FENCE);
const uint m_offset = row + lid;
if (n_real == TILESIZE_N) {
#pragma unroll
for (int t = 1; t < TILESIZE_N; ++t) write_imagef(dst, out_idx[t] + m_offset, acc[t]);
barrier(CLK_GLOBAL_MEM_FENCE);
write_imagef(dst, out_idx[0] + m_offset, acc[0]);
} else {
for (int t = 0; t < n_real; ++t) write_imagef(dst, out_idx[t] + m_offset, acc[t]);
}
}
@@ -0,0 +1,143 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
// Weight layout, feature-major:
// src0_q[row + (k/4)*m] ushort = 4 nibbles (K = 4*grp .. +3)
// src0_d[row + (k/32)*m] half = per-32-block scale
#define TILESIZE_N 32
// IQ4_NL non-linear codebook as signed int8, packed 4 codes per uint.
// divergent nibble lookups read a small __constant uint array + shift,
// never a byte array because byte-indexed __constant loads serialize on Adreno and tank perf
// idx 0-3: -127,-104,-83,-65 = 0x81,0x98,0xAD,0xBF
// idx 4-7: -49,-35,-22,-10 = 0xCF,0xDD,0xEA,0xF6
// idx 8-11: 1, 13, 25, 38 = 0x01,0x0D,0x19,0x26
// idx 12-15: 53, 69, 89,113 = 0x35,0x45,0x59,0x71
__constant uint kvalues_iq4nl_i8x4[4] = {
0xBFAD9881u, 0xF6EADDCFu, 0x26190D01u, 0x71594535u
};
// nibble (0..15) -> its codebook byte in the low 8 bits.
inline uint iq4nl_code(uint n) {
return (kvalues_iq4nl_i8x4[n >> 2] >> ((n & 3u) * 8u)) & 0xFFu;
}
// 4 nibbles in low 16 bits of u -> 4 codebook int8, packed for dp4a.
inline uint iq4nl_pack(ushort u) {
return iq4nl_code((uint)( u & 0xF))
| (iq4nl_code((uint)((u >> 4) & 0xF)) << 8)
| (iq4nl_code((uint)((u >> 8) & 0xF)) << 16)
| (iq4nl_code((uint)((u >> 12) & 0xF)) << 24);
}
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_iq4_nl_q8_1_dp4a(
__global const ushort * src0_q, // IQ4_NL nibbles (4/ushort, feature-major)
__global const half * src0_d, // per-32-block scale, feature-major
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k // K (== ne00)
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
// 8 weight uints (32 codebook int8) for this row, this 32-block.
const uint qsbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = iq4nl_pack(src0_q[qsbase + 0 * m]);
qw.s1 = iq4nl_pack(src0_q[qsbase + 1 * m]);
qw.s2 = iq4nl_pack(src0_q[qsbase + 2 * m]);
qw.s3 = iq4nl_pack(src0_q[qsbase + 3 * m]);
qw.s4 = iq4nl_pack(src0_q[qsbase + 4 * m]);
qw.s5 = iq4nl_pack(src0_q[qsbase + 5 * m]);
qw.s6 = iq4nl_pack(src0_q[qsbase + 6 * m]);
qw.s7 = iq4nl_pack(src0_q[qsbase + 7 * m]);
// cooperatively stage the 32-token x 32-K int8 activations to lm
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += d_w * LD4(sh_d, b) * rf;
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
// dst is [token, feature] row-major (stride m): dst[col*m + row].
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,127 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_N 32
// Expand the 4 nibbles in the low 16 bits of u into 4 bytes (value 0..15),
// packed for the int8 dp4a. The -8 zero-point is applied via the sum term.
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q4_0_q8_1_dp4a(
__global const ushort * src0_q, // q4_0 nibbles (4/ushort, feature-major)
__global const half * src0_d, // per-32-block scale, feature-major
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global const half * src1_sa, // q8_1 per-block sum*d [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k // K (== ne00)
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
// 8 weight uints (32 nibbles) for this row, this 32-block. Feature-major:
// src0_q[row + (k/4 + u)*m], k/4 = step/4 (= step>>2). EXP4 -> dp4a int8.
const uint qsbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = EXP4(src0_q[qsbase + 0 * m]);
qw.s1 = EXP4(src0_q[qsbase + 1 * m]);
qw.s2 = EXP4(src0_q[qsbase + 2 * m]);
qw.s3 = EXP4(src0_q[qsbase + 3 * m]);
qw.s4 = EXP4(src0_q[qsbase + 4 * m]);
qw.s5 = EXP4(src0_q[qsbase + 5 * m]);
qw.s6 = EXP4(src0_q[qsbase + 6 * m]);
qw.s7 = EXP4(src0_q[qsbase + 7 * m]);
// cooperatively stage the 32-token x 32-K int8 activations to LDS
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
// q4_0: w = d*(q-8) -> d_w * (a_d * dp4a(q,qa) - 8 * a_s)
acc[g] += d_w * (LD4(sh_d, b) * rf - 8.0f * LD4(sh_s, b));
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
// dst is [token, feature] row-major (stride m): dst[col*m + row].
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,281 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#ifndef TILESIZE_N
#define TILESIZE_N 32
#endif
#define QK_K 256
#define K_SCALE_SIZE 12
inline void get_scale_min_k4(
int j,
global const uchar * q,
uchar * d,
uchar * m,
uchar mask_d6,
uchar mask_d4,
uchar mask_hi2
) {
if (j < 4) {
*d = q[j] & mask_d6;
*m = q[j+4] & mask_d6;
} else {
*d = (q[j+4] & mask_d4) | ((q[j-4] & mask_hi2) >> 2);
*m = ((q[j+4] >> 4) & mask_d4) | ((q[j] & mask_hi2) >> 2);
}
}
// Expand the 4 nibbles in the low 16 bits of `u` into 4 bytes (one nibble per
// byte, value 0..15), packed for the int8 dp4a.
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 32-K dp4a dot of one token's int8 activations (8 packed uints in lm) against the
// row's 8 packed weight uints. qw passed by value as a uint8 (register), not an array.
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a(
__global const ushort * src0_q, // q4_K weights (noshuffle, packed nibbles)
__global const uchar * src0_s, // 6-bit scale/min codes
__global const half * src0_d, // per-superblock scale
__global const half * src0_dm, // per-superblock min
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global const half * src1_sa, // q8_1 per-block sum*d [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k, // K (== ne00)
uchar mask_d6,
uchar mask_d4,
uchar mask_hi2
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint num_superblocks = (uint)k / QK_K;
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
// One float4 vector-register accumulator per group of 4 tokens (NGROUPS = TILESIZE_N/4).
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) { acc[g] = (float4)(0.0f); }
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const uint sb_idx = step / QK_K;
const uint sub_idx = sub & 7;
// weight scale/min for this WI's row, this subblock
const float dd = (float)src0_d [rrow + sb_idx * m];
const float dmm = (float)src0_dm[rrow + sb_idx * m];
global const uchar * sc = src0_s + rrow * num_superblocks * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE;
uchar sv, mn;
get_scale_min_k4(sub_idx, sc, &sv, &mn, mask_d6, mask_d4, mask_hi2);
const float scale = dd * (float)sv;
const float minv = dmm * (float)mn;
// repack this row's 32 weight nibbles into 8 dp4a uints. The packed q4_K
// layout stores one ushort = 4 consecutive-K nibbles for a row at
// src0_q[row + (K_group)*m], K_group = step/4 + u.
const uint wbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = EXP4(src0_q[wbase + 0 * m]);
qw.s1 = EXP4(src0_q[wbase + 1 * m]);
qw.s2 = EXP4(src0_q[wbase + 2 * m]);
qw.s3 = EXP4(src0_q[wbase + 3 * m]);
qw.s4 = EXP4(src0_q[wbase + 4 * m]);
qw.s5 = EXP4(src0_q[wbase + 5 * m]);
qw.s6 = EXP4(src0_q[wbase + 6 * m]);
qw.s7 = EXP4(src0_q[wbase + 7 * m]);
// cooperatively stage the 32-token x 32-K int8 activations to lm
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += scale * LD4(sh_d, b) * rf - minv * LD4(sh_s, b);
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
// dst is [token, feature] row-major (stride m): dst[col*m + row]. Scatter each
// lane with a per-token padding guard (dst is non-contiguous in token).
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a_wimg(
__read_only image1d_buffer_t src0_q_img, // q4_K weights as uint32 texels (2 ushorts/texel)
__global const uchar * src0_s, // 6-bit scale/min codes
__global const half * src0_d, // per-superblock scale
__global const half * src0_dm, // per-superblock min
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global const half * src1_sa, // q8_1 per-block sum*d [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k, // K (== ne00)
uchar mask_d6,
uchar mask_d4,
uchar mask_hi2
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
// Constant per WI: the ushort the row needs always sits in the same half of
// its uint32 texel (m even => index parity == rrow parity). Hoist the shift.
const uint sel = (rrow & 1u) * 16u;
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
const uint num_superblocks = (uint)k / QK_K;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const uint sb_idx = step / QK_K;
const uint sub_idx = sub & 7;
const float dd = (float)src0_d [rrow + sb_idx * m];
const float dmm = (float)src0_dm[rrow + sb_idx * m];
global const uchar * sc = src0_s + rrow * num_superblocks * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE;
uchar sv, mn;
get_scale_min_k4(sub_idx, sc, &sv, &mn, mask_d6, mask_d4, mask_hi2);
const float scale = dd * (float)sv;
const float minv = dmm * (float)mn;
const uint wbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = EXP4(read_imageui(src0_q_img, (int)((wbase + 0 * m) >> 1)).x >> sel);
qw.s1 = EXP4(read_imageui(src0_q_img, (int)((wbase + 1 * m) >> 1)).x >> sel);
qw.s2 = EXP4(read_imageui(src0_q_img, (int)((wbase + 2 * m) >> 1)).x >> sel);
qw.s3 = EXP4(read_imageui(src0_q_img, (int)((wbase + 3 * m) >> 1)).x >> sel);
qw.s4 = EXP4(read_imageui(src0_q_img, (int)((wbase + 4 * m) >> 1)).x >> sel);
qw.s5 = EXP4(read_imageui(src0_q_img, (int)((wbase + 5 * m) >> 1)).x >> sel);
qw.s6 = EXP4(read_imageui(src0_q_img, (int)((wbase + 6 * m) >> 1)).x >> sel);
qw.s7 = EXP4(read_imageui(src0_q_img, (int)((wbase + 7 * m) >> 1)).x >> sel);
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += scale * LD4(sh_d, b) * rf - minv * LD4(sh_s, b);
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,235 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
// Weight layout
// src0_qs[row + (k/4)*m] ushort = 4 low nibbles (K = 4*grp .. +3)
// src0_qh[row + (k/8)*m] uchar = 8 high bits (one per element)
// src0_d [row + (k/32)*m] half = per-32-block scale
#define TILESIZE_N 32
// 4 nibbles in low 16 bits of u -> 4 bytes (value 0..15)
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 4 high bits (one per element, in bits 0..3 of h) -> bit4 of each of 4 bytes
#define EXP1(h) ( (((uint)((h) & 0x1u)) << 4) | \
(((uint)((h) & 0x2u)) << 11) | \
(((uint)((h) & 0x4u)) << 18) | \
(((uint)((h) & 0x8u)) << 25) )
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q5_0_q8_1_dp4a(
__global const ushort * src0_qs, // q5_0 low nibbles (4/ushort, feature-major)
__global const uchar * src0_qh, // q5_0 high-bit plane (8/uchar, feature-major)
__global const half * src0_d, // per-32-block scale, feature-major
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global const half * src1_sa, // q8_1 per-block sum*d [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k // K (== ne00)
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
const float minv = d_w * 16.0f; // -16 centering -> subtract via q8_1 sum
// 8 weight uints (32 elements) for this row, this 32-block.
// nibbles: src0_qs[row + (step/4 + u)*m]; high bits: src0_qh[row + (step/8 + u/2)*m],
// 4-bit group selected by (u&1)*4.
const uint qsbase = rrow + (step >> 2) * (uint)m;
const uint qhbase = rrow + (step >> 3) * (uint)m;
uint8 qw;
#define QW(u) (EXP4(src0_qs[qsbase + (u) * m]) | \
EXP1((uint)(src0_qh[qhbase + ((u) >> 1) * m] >> (((u) & 1u) * 4u)) & 0xFu))
qw.s0 = QW(0); qw.s1 = QW(1); qw.s2 = QW(2); qw.s3 = QW(3);
qw.s4 = QW(4); qw.s5 = QW(5); qw.s6 = QW(6); qw.s7 = QW(7);
#undef QW
// cooperatively stage the 32-token x 32-K int8 activations to lm
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += d_w * LD4(sh_d, b) * rf - minv * LD4(sh_s, b);
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q5_0_q8_1_dp4a_wimg(
__read_only image1d_buffer_t src0_qs_img, // q5_0 low nibbles as uint32 texels (2 ushorts/texel)
__global const uchar * src0_qh,
__global const half * src0_d,
__global const uint * src1_qa,
__global const half * src1_da,
__global const half * src1_sa,
__global float * dst,
ulong offsetd,
int m,
int n_no_padding,
int k
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0);
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0;
const uint sel = (rrow & 1u) * 16u; // constant per WI: qs ushort half in its uint32 texel
const uint k_u = (uint)k >> 2;
const uint k_b = (uint)k >> 5;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
const float minv = d_w * 16.0f;
const uint qsbase = rrow + (step >> 2) * (uint)m; // ushort index
const uint qhbase = rrow + (step >> 3) * (uint)m;
uint8 qw;
// qs ushort via texture: uint32 texel = ushort_index>>1, half = sel.
#define QSU(u) ((read_imageui(src0_qs_img, (int)((qsbase + (u) * m) >> 1)).x >> sel) & 0xFFFFu)
#define QW(u) (EXP4(QSU(u)) | \
EXP1((uint)(src0_qh[qhbase + ((u) >> 1) * m] >> (((u) & 1u) * 4u)) & 0xFu))
qw.s0 = QW(0); qw.s1 = QW(1); qw.s2 = QW(2); qw.s3 = QW(3);
qw.s4 = QW(4); qw.s5 = QW(5); qw.s6 = QW(6); qw.s7 = QW(7);
#undef QW
#undef QSU
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += d_w * LD4(sh_d, b) * rf - minv * LD4(sh_s, b);
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,164 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_N 32
#define QK_K 256
#define K_SCALE_SIZE 12
inline void get_scale_min_k4(
int j,
global const uchar * q,
uchar * d,
uchar * m,
uchar mask_d6,
uchar mask_d4,
uchar mask_hi2
) {
if (j < 4) {
*d = q[j] & mask_d6;
*m = q[j+4] & mask_d6;
} else {
*d = (q[j+4] & mask_d4) | ((q[j-4] & mask_hi2) >> 2);
*m = ((q[j+4] >> 4) & mask_d4) | ((q[j] & mask_hi2) >> 2);
}
}
// 4 nibbles in the low 16 bits of `u` -> 4 bytes (value 0..15, bits 0-3).
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 4 high bits (one per element, in bits 0-3 of h) -> bit 4 of each of 4 bytes,
// so OR with EXP4 forms the 5-bit q5_K code 0..31.
#define EXP1(h) ( (((uint)((h) & 0x1u)) << 4) | \
(((uint)((h) & 0x2u)) << 11) | \
(((uint)((h) & 0x4u)) << 18) | \
(((uint)((h) & 0x8u)) << 25) )
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q5_k_q8_1_dp4a(
__global const ushort * src0_q, // q5_K low nibbles (transposed, ushort = 4 nibbles)
__global const uchar * src0_qh, // q5_K high bits (transposed, uchar = 8 elems/byte)
__global const uchar * src0_s, // 6-bit scale/min codes [row][superblock][12]
__global const half * src0_d, // per-superblock scale (transposed)
__global const half * src0_dm, // per-superblock min (transposed)
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global const half * src1_sa, // q8_1 per-block sum*d [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k, // K (== ne00)
uchar mask_d6,
uchar mask_d4,
uchar mask_hi2
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0;
const uint num_superblocks = (uint)k / QK_K;
const uint k_u = (uint)k >> 2;
const uint k_b = (uint)k >> 5;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
__local half sh_s[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const uint sb_idx = step / QK_K;
const uint sub_idx = sub & 7;
const float dd = (float)src0_d [rrow + sb_idx * m];
const float dmm = (float)src0_dm[rrow + sb_idx * m];
global const uchar * sc = src0_s + rrow * num_superblocks * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE;
uchar sv, mn;
get_scale_min_k4(sub_idx, sc, &sv, &mn, mask_d6, mask_d4, mask_hi2);
const float scale = dd * (float)sv;
const float minv = dmm * (float)mn;
// repack this row's 32 weights (nibble | high-bit) into 8 dp4a uints.
// ushort u -> 4 elements at K = step + u*4; its 4 high bits are nibble
// (u&1) of qh byte (step/8 + u/2).
const uint wbase = rrow + (step >> 2) * (uint)m;
const uint qhbase = rrow + (step >> 3) * (uint)m;
uint8 qw;
#define QWU(u) ( EXP4((uint)src0_q[wbase + (uint)(u) * m]) \
| EXP1( (uint)((src0_qh[qhbase + (uint)((u) >> 1) * m] >> (((u) & 1) * 4)) & 0x0Fu) ) )
qw.s0 = QWU(0); qw.s1 = QWU(1); qw.s2 = QWU(2); qw.s3 = QWU(3);
qw.s4 = QWU(4); qw.s5 = QWU(5); qw.s6 = QWU(6); qw.s7 = QWU(7);
#undef QWU
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
sh_s[lid] = (c < (uint)n_no_padding) ? src1_sa[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += scale * LD4(sh_d, b) * rf - minv * LD4(sh_s, b);
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,144 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
#define TILESIZE_N 32
#define QK_K 256
// 4 nibbles in the low 16 bits of `u` -> 4 bytes (value 0..15, in bits 0-3).
#define EXP4(u) ( ((uint)((u) & 0x000Fu)) | \
(((uint)((u) & 0x00F0u)) << 4) | \
(((uint)((u) & 0x0F00u)) << 8) | \
(((uint)((u) & 0xF000u)) << 12) )
// 4 2-bit highs in byte `b` -> 4 bytes, value 0..3 in bits 4-5 (pre-multiplied
// by 16 so it ORs with the EXP4 nibble to form q6 in 0..63).
#define EXP2(b) ( (((uint)((b) & 0x03u)) << 4) | \
(((uint)((b) & 0x0Cu)) << 10) | \
(((uint)((b) & 0x30u)) << 16) | \
(((uint)((b) & 0xC0u)) << 22) )
// q6 (0..63, bits 0-5 of each byte) -> (q6-32) as a signed int8 per byte.
inline uint SIGN6(uint q6p) {
uint x = q6p ^ 0x20202020u;
uint s = x & 0x20202020u;
return x | (s << 1) | (s << 2);
}
// 16-K dp4a dot: 4 packed weight uints against 4 packed int8 activation uints.
inline int dot4_q8a(uint w0, uint w1, uint w2, uint w3,
uint a0, uint a1, uint a2, uint a3) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(w0, a0, r);
r = dot_acc_sat_4x8packed_ss_int(w1, a1, r);
r = dot_acc_sat_4x8packed_ss_int(w2, a2, r);
r = dot_acc_sat_4x8packed_ss_int(w3, a3, r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q6_k_q8_1_dp4a(
__global const ushort * src0_ql, // q6_K low nibbles (noshuffle)
__global const uchar * src0_qh, // q6_K high 2-bit (uchar, 4 highs/elem)
__global const ushort * src0_s, // int8 scale codes (2 chars/ushort, per 16)
__global const half * src0_d, // per-superblock scale
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k // K (== ne00)
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5; // 32-block index along K
const uint sb_idx = step / QK_K; // superblock index
// q6_K superblock scale + the two int8 sub-scales spanning this 32-block
const float dd = (float)src0_d[rrow + sb_idx * m];
const char2 sc = as_char2(src0_s[rrow + sub * m]);
const float scale0 = dd * (float)sc.s0; // K step..step+15
const float scale1 = dd * (float)sc.s1; // K step+16..step+31
// repack this row's 32 weights into 8 dp4a uints (4 K each). ql ushort +
// qh uchar are co-located at src0_*[row + (step/4 + u)*m].
const uint wbase = rrow + (step >> 2) * (uint)m;
uint qw[8];
#pragma unroll
for (int u = 0; u < 8; ++u) {
const uint o = wbase + (uint)u * (uint)m;
qw[u] = SIGN6(EXP4((uint)src0_ql[o]) | EXP2((uint)src0_qh[o]));
}
// cooperatively stage the 32-token x 32-K int8 activations + scale
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
#define DOT_TOK(j) { \
__local const uint * a = sh_qa[b + (j)]; \
const int raw1 = dot4_q8a(qw[0], qw[1], qw[2], qw[3], a[0], a[1], a[2], a[3]); \
const int raw2 = dot4_q8a(qw[4], qw[5], qw[6], qw[7], a[4], a[5], a[6], a[7]); \
rf.s##j = scale0 * (float)raw1 + scale1 * (float)raw2; \
}
DOT_TOK(0); DOT_TOK(1); DOT_TOK(2); DOT_TOK(3);
#undef DOT_TOK
const float4 ad = (float4)((float)sh_d[b+0], (float)sh_d[b+1], (float)sh_d[b+2], (float)sh_d[b+3]);
acc[g] += ad * rf;
}
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
// dst is [token, feature] row-major (stride m): dst[col*m + row].
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -0,0 +1,212 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
#ifdef cl_khr_integer_dot_product
#pragma OPENCL EXTENSION cl_khr_integer_dot_product : enable
#endif
// ne1<=8 keeps the f16 / bin small-batch path.
#define TILESIZE_N 32
// 32-K dp4a dot of one token's int8 activations (8 packed uints in lm) against
// 8 packed weight uints. q8_0 weights are already dp4a-format signed int8.
inline int dot8_q8a(uint8 qw, __local const uint * a) {
int r = 0;
r = dot_acc_sat_4x8packed_ss_int(qw.s0, a[0], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s1, a[1], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s2, a[2], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s3, a[3], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s4, a[4], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s5, a[5], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s6, a[6], r);
r = dot_acc_sat_4x8packed_ss_int(qw.s7, a[7], r);
return r;
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q8_0_q8_1_dp4a(
__global const uint * src0_q, // q8_0 weights: signed int8, 4/uint, feature-major
__global const half * src0_d, // per-32-block scale, feature-major [row + (k/32)*m]
__global const uint * src1_qa, // q8_1 activations int8 (as uint, 4/elem) [N, K]
__global const half * src1_da, // q8_1 per-block scale [N, K/32]
__global float * dst,
ulong offsetd,
int m, // output features (rows)
int n_no_padding, // tokens (cols)
int k // K (== ne00)
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0); // 0..63 -> row within the M-tile
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
// 8 weight uints (32 int8) for this row, this 32-block. Feature-major:
// src0_q[row + (k/4 + u)*m], k/4 = step/4 (= step>>2).
const uint wbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = src0_q[wbase + 0 * m];
qw.s1 = src0_q[wbase + 1 * m];
qw.s2 = src0_q[wbase + 2 * m];
qw.s3 = src0_q[wbase + 3 * m];
qw.s4 = src0_q[wbase + 4 * m];
qw.s5 = src0_q[wbase + 5 * m];
qw.s6 = src0_q[wbase + 6 * m];
qw.s7 = src0_q[wbase + 7 * m];
// cooperatively stage the 32-token x 32-K int8 activations to LDS
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += d_w * LD4(sh_d, b) * rf;
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
// dst is [token, feature] row-major (stride m): dst[col*m + row].
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
__attribute__((qcom_wave_pair_mode(1)))
kernel void kernel_gemm_noshuffle_q8_0_q8_1_dp4a_wimg(
__read_only image1d_buffer_t src0_q_img, // q8_0 weights as uint32 texels (4 int8/texel)
__global const half * src0_d,
__global const uint * src1_qa,
__global const half * src1_da,
__global float * dst,
ulong offsetd,
int m,
int n_no_padding,
int k
) {
dst = (global float *)((global char *)dst + offsetd);
const uint lid = get_local_id(0);
const uint block_id_m = get_global_id(1);
const uint block_id_n = get_global_id(2);
const uint row = block_id_m * 64 + lid;
const uint col_base = block_id_n * TILESIZE_N;
const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0;
const uint k_u = (uint)k >> 2;
const uint k_b = (uint)k >> 5;
__local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N];
#define NGROUPS (TILESIZE_N / 4)
float4 acc[NGROUPS];
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) acc[g] = (float4)(0.0f);
for (uint step = 0; step < (uint)k; step += 32) {
const uint sub = step >> 5;
const float d_w = (float)src0_d[rrow + sub * (uint)m];
const uint wbase = rrow + (step >> 2) * (uint)m;
uint8 qw;
qw.s0 = read_imageui(src0_q_img, (int)(wbase + 0 * m)).x;
qw.s1 = read_imageui(src0_q_img, (int)(wbase + 1 * m)).x;
qw.s2 = read_imageui(src0_q_img, (int)(wbase + 2 * m)).x;
qw.s3 = read_imageui(src0_q_img, (int)(wbase + 3 * m)).x;
qw.s4 = read_imageui(src0_q_img, (int)(wbase + 4 * m)).x;
qw.s5 = read_imageui(src0_q_img, (int)(wbase + 5 * m)).x;
qw.s6 = read_imageui(src0_q_img, (int)(wbase + 6 * m)).x;
qw.s7 = read_imageui(src0_q_img, (int)(wbase + 7 * m)).x;
for (uint idx = lid; idx < TILESIZE_N * 8; idx += 64) {
const uint t = idx >> 3;
const uint u = idx & 7;
const uint c = col_base + t;
sh_qa[t][u] = (c < (uint)n_no_padding) ? src1_qa[c * k_u + (step >> 2) + u] : 0u;
}
if (lid < TILESIZE_N) {
const uint c = col_base + lid;
sh_d[lid] = (c < (uint)n_no_padding) ? src1_da[c * k_b + sub] : (half)0;
}
barrier(CLK_LOCAL_MEM_FENCE);
#define LD4(arr, b) ((float4)((float)arr[(b)+0], (float)arr[(b)+1], (float)arr[(b)+2], (float)arr[(b)+3]))
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const int b = g * 4;
float4 rf;
rf.s0 = (float)dot8_q8a(qw, sh_qa[b+0]); rf.s1 = (float)dot8_q8a(qw, sh_qa[b+1]);
rf.s2 = (float)dot8_q8a(qw, sh_qa[b+2]); rf.s3 = (float)dot8_q8a(qw, sh_qa[b+3]);
acc[g] += d_w * LD4(sh_d, b) * rf;
}
#undef LD4
barrier(CLK_LOCAL_MEM_FENCE);
}
if (!row_valid) {
return;
}
#pragma unroll
for (int g = 0; g < NGROUPS; ++g) {
const uint b = (uint)(g * 4);
const float4 a = acc[g];
const uint c0 = col_base + b;
if (c0 + 0 < (uint)n_no_padding) dst[(c0 + 0) * (uint)m + row] = a.s0;
if (c0 + 1 < (uint)n_no_padding) dst[(c0 + 1) * (uint)m + row] = a.s1;
if (c0 + 2 < (uint)n_no_padding) dst[(c0 + 2) * (uint)m + row] = a.s2;
if (c0 + 3 < (uint)n_no_padding) dst[(c0 + 3) * (uint)m + row] = a.s3;
}
#undef NGROUPS
}
@@ -163,3 +163,95 @@ __kernel void kernel_gemv_moe_mxfp4_f32_ns(
}
}
__attribute__((qcom_reqd_sub_group_size("half")))
__kernel void kernel_gemv_moe_mxfp4_f32_ns_wimg(
__read_only image1d_buffer_t src0_q,
__global uchar * src0_e,
__read_only image1d_buffer_t src1,
__global uint * src2,
__global float * dst,
ulong offsetd,
int ne00,
int ne01,
int ne11
) {
uint i01 = get_global_id(0);
uint i20 = get_global_id(2);
uint sgid = get_local_id(1);
uint slid = get_sub_group_local_id();
if (i01 >= ne01) {
return;
}
uint i11 = i20 % ne11;
uint expert_id = src2[i20];
uint expert_offset = expert_id * ne00 * ne01 / 32;
__private float sum = 0.0f;
for (uint ib00 = sgid; ib00 < (ne00 / QK_MXFP4); ib00 += N_SIMDGROUP) {
uint4 regQ;
uint block_offset = expert_offset * 4 + ib00 * ne01 * 4 + i01;
regQ.s0 = read_imageui(src0_q, (int)(block_offset)).x;
regQ.s1 = read_imageui(src0_q, (int)(block_offset + ne01)).x;
regQ.s2 = read_imageui(src0_q, (int)(block_offset + ne01 * 2)).x;
regQ.s3 = read_imageui(src0_q, (int)(block_offset + ne01 * 3)).x;
uint offset = i11 * ne00 / 4 + ib00 * 8;
half8 fp16x8 = mxfp4_to_fp16_packed8(as_ushort2(regQ.s0));
float4 shared_y4;
shared_y4 = read_imagef(src1, (offset + 0));
float4 acc = shared_y4 * convert_float4(fp16x8.lo);
shared_y4 = read_imagef(src1, (offset + 1));
acc += shared_y4 * convert_float4(fp16x8.hi);
fp16x8 = mxfp4_to_fp16_packed8(as_ushort2(regQ.s1));
shared_y4 = read_imagef(src1, (offset + 2));
acc += shared_y4 * convert_float4(fp16x8.lo);
shared_y4 = read_imagef(src1, (offset + 3));
acc += shared_y4 * convert_float4(fp16x8.hi);
fp16x8 = mxfp4_to_fp16_packed8(as_ushort2(regQ.s2));
shared_y4 = read_imagef(src1, (offset + 4));
acc += shared_y4 * convert_float4(fp16x8.lo);
shared_y4 = read_imagef(src1, (offset + 5));
acc += shared_y4 * convert_float4(fp16x8.hi);
fp16x8 = mxfp4_to_fp16_packed8(as_ushort2(regQ.s3));
shared_y4 = read_imagef(src1, (offset + 6));
acc += shared_y4 * convert_float4(fp16x8.lo);
shared_y4 = read_imagef(src1, (offset + 7));
acc += shared_y4 * convert_float4(fp16x8.hi);
uchar regE = src0_e[ib00 * ne01 + i01 + expert_offset];
sum += e8m0_to_fp32(regE) * ((acc.s0 + acc.s1) + (acc.s2 + acc.s3));
}
__local float reduceLM[SIMDGROUP_WIDTH * (N_SIMDGROUP - 1)];
if (sgid == 1) reduceLM[SIMDGROUP_WIDTH * 0 + slid] = sum;
if (sgid == 2) reduceLM[SIMDGROUP_WIDTH * 1 + slid] = sum;
if (sgid == 3) reduceLM[SIMDGROUP_WIDTH * 2 + slid] = sum;
barrier(CLK_LOCAL_MEM_FENCE);
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 0 + slid];
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 1 + slid];
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 2 + slid];
if (sgid == 0) {
dst = dst + (offsetd >> 2);
dst[i01 + i20 * ne01] = sum;
}
}
@@ -153,3 +153,114 @@ __kernel void kernel_gemv_moe_q4_k_f32_ns(
dst[i01 + i20 * ne01] = sum;
}
}
__attribute__((qcom_reqd_sub_group_size("half")))
__kernel void kernel_gemv_moe_q4_k_f32_ns_wimg(
__read_only image1d_buffer_t src0_q,
__global half * src0_d,
__global half * src0_dm,
__global uchar * src0_s,
__read_only image1d_buffer_t src1,
__global uint * src2,
__global float * dst,
ulong offsetd,
int ne00,
int ne01,
int ne11
) {
uint i01 = get_global_id(0);
uint i20 = get_global_id(2);
uint sgid = get_local_id(1);
uint slid = get_sub_group_local_id();
if (i01 >= ne01) {
return;
}
uint i11 = i20 % ne11;
uint expert_id = src2[i20];
int num_superblocks = ne00 / QK_K;
int num_subblocks = ne00 / 32;
int scales_per_row = num_superblocks * K_SCALE_SIZE;
uint expert_q_offset = expert_id * (ne00 / 8) * ne01;
uint expert_d_offset = expert_id * num_superblocks * ne01;
__private float sum = 0.0f;
for (uint ib = sgid; ib < num_subblocks; ib += N_SIMDGROUP) {
uint sb = ib / 8;
uint j = ib % 8;
half d_val = src0_d[expert_d_offset + sb * ne01 + i01];
half dm_val = src0_dm[expert_d_offset + sb * ne01 + i01];
global const uchar * sc = src0_s + (expert_id * ne01 + i01) * scales_per_row + sb * K_SCALE_SIZE;
uchar sv, mn;
get_scale_min_k4(j, sc, &sv, &mn);
float scale = (float)d_val * (float)sv;
float minv = (float)dm_val * (float)mn;
uint q_base = expert_q_offset + ib * ne01 * 4 + i01;
uint4 regQ;
regQ.s0 = read_imageui(src0_q, (int)(q_base)).x;
regQ.s1 = read_imageui(src0_q, (int)(q_base + ne01)).x;
regQ.s2 = read_imageui(src0_q, (int)(q_base + ne01 * 2)).x;
regQ.s3 = read_imageui(src0_q, (int)(q_base + ne01 * 3)).x;
uint y_offset = i11 * ne00 / 4 + ib * 8;
float8 fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s0), scale, minv);
float4 shared_y4;
shared_y4 = read_imagef(src1, (y_offset + 0));
float4 acc = shared_y4 * fp32x8.lo;
shared_y4 = read_imagef(src1, (y_offset + 1));
acc += shared_y4 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s1), scale, minv);
shared_y4 = read_imagef(src1, (y_offset + 2));
acc += shared_y4 * fp32x8.lo;
shared_y4 = read_imagef(src1, (y_offset + 3));
acc += shared_y4 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s2), scale, minv);
shared_y4 = read_imagef(src1, (y_offset + 4));
acc += shared_y4 * fp32x8.lo;
shared_y4 = read_imagef(src1, (y_offset + 5));
acc += shared_y4 * fp32x8.hi;
fp32x8 = q4_k_to_fp32_packed8(as_ushort2(regQ.s3), scale, minv);
shared_y4 = read_imagef(src1, (y_offset + 6));
acc += shared_y4 * fp32x8.lo;
shared_y4 = read_imagef(src1, (y_offset + 7));
acc += shared_y4 * fp32x8.hi;
sum += ((acc.s0 + acc.s1) + (acc.s2 + acc.s3));
}
__local float reduceLM[SIMDGROUP_WIDTH * (N_SIMDGROUP - 1)];
if (sgid == 1) reduceLM[SIMDGROUP_WIDTH * 0 + slid] = sum;
if (sgid == 2) reduceLM[SIMDGROUP_WIDTH * 1 + slid] = sum;
if (sgid == 3) reduceLM[SIMDGROUP_WIDTH * 2 + slid] = sum;
barrier(CLK_LOCAL_MEM_FENCE);
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 0 + slid];
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 1 + slid];
if (sgid == 0) sum += reduceLM[SIMDGROUP_WIDTH * 2 + slid];
if (sgid == 0) {
dst = dst + (offsetd >> 2);
dst[i01 + i20 * ne01] = sum;
}
}
@@ -0,0 +1,36 @@
// Fused MoE combine epilogue: replaces the router-weight MUL + the (n_expert_used-1)
// cross-expert ADD chain with ONE weighted-sum-across-experts pass.
// dst[row, tok] = sum_e experts[row, e, tok] * weights[0, e, tok]
// experts: [n_embd, n_expert_used, n_tokens] f32 (contiguous after down-proj GEMM)
// weights: [1, n_expert_used, n_tokens] f32
// dst: [n_embd, n_tokens] f32
// One read of experts + one write of dst (eliminates the intermediate weighted
// buffer and the k-1 elementwise add round-trips). Vectorized float4 over rows.
// strides e1/e2/w1/w2/d1 are in ELEMENTS (floats).
__kernel void kernel_moe_combine_f32(
__global const char * e_buf, ulong off_e,
__global const char * w_buf, ulong off_w,
__global char * d_buf, ulong off_d,
int n_embd4, // n_embd / 4
int k, // n_expert_used
int n_tokens,
uint e1, uint e2, // experts strides (elements): per-expert, per-token
uint w1, uint w2, // weights strides (elements)
uint d1) // dst per-token stride (elements)
{
const uint r4 = get_global_id(0);
const uint tok = get_global_id(1);
if (r4 >= (uint)n_embd4 || tok >= (uint)n_tokens) return;
__global const float * E = (__global const float *)(e_buf + off_e) + tok*e2 + r4*4u;
__global const float * W = (__global const float *)(w_buf + off_w) + tok*w2;
float4 acc = (float4)(0.0f);
for (int e = 0; e < k; ++e) {
acc = mad(vload4(0, E + (uint)e*e1), (float4)(W[(uint)e*w1]), acc);
}
__global float * D = (__global float *)(d_buf + off_d) + tok*d1 + r4*4u;
vstore4(acc, 0, D);
}
@@ -0,0 +1,64 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
// Fused MoE activation reorder + q8_1 quantization for the dp4a prefill GEMM.
// Combines kernel_moe_reorder_b (gather src1 rows per the post-router map) with
// the q8_1 quant pre-pass, so the f32 reordered-activation tile buffer is never
// materialised (saves a full write + read of [tok_slots * ne00] floats).
//
// One work-item per (token_slot, 32-block). Padding lanes (router 0xFFFFFFFF)
// emit d=0,s=0,qs=0 so they contribute nothing to the GEMM, exactly as the
// reorder zero-fill did. Output layout matches kernel_moe_quant_a_q8_1:
// qa[token_slot*K + blk*32 + i], da/sa[token_slot*(K/32) + blk].
__kernel void kernel_moe_reorder_quant_a_q8_1(
__global const float * src, // original activations (offset applied)
__global const uint * router, // post-router indices [tok_slots]
__global char * qa,
__global half * da,
__global half * sa,
__global const int * total_tiles,
uint K,
ushort map_ratio,
uint tile_size,
uint n_kblocks // K / 32
) {
const uint blk = get_global_id(0); // 32-block along K
const uint tok = get_global_id(1); // token slot (post_router_idx)
if (blk >= n_kblocks || tok >= (uint)total_tiles[0] * tile_size) {
return;
}
const uint out_base = tok * K + blk * 32;
const uint bidx = tok * n_kblocks + blk;
const uint router_idx = router[tok];
float v[32];
float amax = 0.0f;
if (router_idx == 0xFFFFFFFF) {
#pragma unroll
for (int i = 0; i < 32; ++i) v[i] = 0.0f;
} else {
const uint act_idx = router_idx / map_ratio;
const uint in_base = act_idx * K + blk * 32;
#pragma unroll
for (int i = 0; i < 32; ++i) {
v[i] = src[in_base + i];
amax = fmax(amax, fabs(v[i]));
}
}
const float d = amax / 127.0f;
const float id = (amax > 0.0f) ? (127.0f / amax) : 0.0f;
int sum = 0;
#pragma unroll
for (int i = 0; i < 32; ++i) {
const int q = (int)rint(v[i] * id);
qa[out_base + i] = (char)q;
sum += q;
}
da[bidx] = (half)d;
sa[bidx] = (half)(d * (float)sum);
}
@@ -0,0 +1,42 @@
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
// Quantize a contiguous [N, K] f32 activation buffer (token-major, K contiguous
// per token) into q8_1 blocks of 32: int8 quants + per-block scale d + per-block
// sum s (= d * Sum(qs)). Consumed by kernel_gemm_noshuffle_q4_k_q8_1_dp4a for the
// dp4a (int8) dense q4_K prefill GEMM. One work-item per 32-element block.
__kernel void kernel_quant_a_q8_1(
__global const float * src, // [N * K]
__global char * qa, // [N * K]
__global half * da, // [N * (K/32)]
__global half * sa, // [N * (K/32)]
int total_blocks // N * (K/32)
) {
const int blk = get_global_id(0);
if (blk >= total_blocks) {
return;
}
const int base = blk * 32;
float v[32];
float amax = 0.0f;
#pragma unroll
for (int i = 0; i < 32; ++i) {
v[i] = src[base + i];
amax = fmax(amax, fabs(v[i]));
}
const float d = amax / 127.0f;
const float id = (amax > 0.0f) ? (127.0f / amax) : 0.0f;
int sum = 0;
#pragma unroll
for (int i = 0; i < 32; ++i) {
const int q = (int)rint(v[i] * id);
qa[base + i] = (char)q;
sum += q;
}
da[blk] = (half)d;
sa[blk] = (half)(d * (float)sum);
}