+#ifdef USE_SUBGROUP_REDUCTION
+enable subgroups;
+#endif
enable f16;
#define DECLARE_BYTE_LOADERS_SRC0
#include "common_decls.tmpl"
+#ifdef U32_DEQUANT_HELPERS
+#define SRC0_TYPE u32
-#ifdef VEC
+fn byte_of(v: u32, b: u32) -> u32 {
+ return (v >> (b * 8u)) & 0xFFu;
+}
+
+fn sbyte_of(v: u32, b: u32) -> i32 {
+ let raw = i32((v >> (b * 8u)) & 0xFFu);
+ return select(raw, raw - 256, raw >= 128);
+}
+#endif
-#define VEC_SIZE 4
-#define DST_TYPE vec4<f32>
+#ifdef VEC
+#define VEC_SIZE 4u
#define SRC0_TYPE vec4<SRC0_INNER_TYPE>
#define SRC1_TYPE vec4<SRC1_INNER_TYPE>
fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
return f32(dot(SRC1_TYPE(src0_val), src1_val));
}
-
-fn store_val(group_base: u32) -> vec4<f32> {
- return vec4<f32>(partial_sums[group_base],
- partial_sums[group_base + THREADS_PER_OUTPUT],
- partial_sums[group_base + THREADS_PER_OUTPUT * 2],
- partial_sums[group_base + THREADS_PER_OUTPUT * 3]);
-}
#endif
#ifdef SCALAR
-
-#define VEC_SIZE 1
-#define DST_TYPE f32
+#define VEC_SIZE 1u
#define SRC0_TYPE SRC0_INNER_TYPE
#define SRC1_TYPE SRC1_INNER_TYPE
fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
return f32(src0_val) * f32(src1_val);
}
+#endif
+
+struct MulMatParams {
+ offset_src0: u32,
+ offset_src1: u32,
+ offset_dst: u32,
+ m: u32,
+ n: u32,
+ k: u32,
+ stride_01: u32,
+ stride_11: u32,
+ stride_02: u32,
+ stride_12: u32,
+ stride_03: u32,
+ stride_13: u32,
+ bs02: u32,
+ bs03: u32,
+ broadcast2: u32,
+ broadcast3: u32
+};
-fn store_val(group_base: u32) -> f32 {
- return partial_sums[group_base];
+@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>;
+@group(0) @binding(1) var<storage, read_write> src1: array<SRC1_TYPE>;
+@group(0) @binding(2) var<storage, read_write> dst: array<f32>;
+
+@group(0) @binding(3) var<uniform> params: MulMatParams;
+
+// Flattened as [row][thread] to keep each row's reduction contiguous in memory.
+var<workgroup> partial_sums: array<f32, OUTPUTS_PER_WG * WG_SIZE>;
+
+fn partial_index(row: u32, thread: u32) -> u32 {
+ return row * WG_SIZE + thread;
}
+
+@compute @workgroup_size(WG_SIZE)
+fn main(
+ @builtin(local_invocation_id) local_id: vec3<u32>,
+ @builtin(workgroup_id) wg_id: vec3<u32>,
+ @builtin(num_workgroups) num_wg: vec3<u32>
+#ifdef USE_SUBGROUP_REDUCTION
+ , @builtin(subgroup_id) subgroup_id: u32,
+ @builtin(subgroup_invocation_id) subgroup_invocation_id: u32,
+ @builtin(num_subgroups) num_subgroups: u32,
+ @builtin(subgroup_size) subgroup_size: u32
#endif
+) {
+ let thread_id = local_id.x;
+
+ let total_batches = params.bs02 * params.broadcast2 * params.bs03 * params.broadcast3;
+ let wg_linear = wg_id.y * num_wg.x + wg_id.x;
+ let output_groups = (params.m + OUTPUTS_PER_WG - 1u) / OUTPUTS_PER_WG;
+ let batch_idx = wg_linear / output_groups;
+ if (batch_idx >= total_batches) {
+ return;
+ }
+
+ let row_base = (wg_linear % output_groups) * OUTPUTS_PER_WG;
+
+ let dst2_stride = params.m * params.n;
+ let dst2_idx = batch_idx % (params.bs02 * params.broadcast2);
+ let dst3_stride = dst2_stride * params.bs02 * params.broadcast2;
+ let dst3_idx = batch_idx / (params.bs02 * params.broadcast2);
+ let src03_idx = dst3_idx / params.broadcast3;
+ let src13_idx = dst3_idx;
+ let src02_idx = dst2_idx / params.broadcast2;
+ let src12_idx = dst2_idx;
+
+ let src0_batch_offset = params.offset_src0 + src03_idx * params.stride_03 + src02_idx * params.stride_02;
+ let src1_idx_base = params.offset_src1 + src13_idx * params.stride_13 + src12_idx * params.stride_12;
+ let dst_idx_base = params.offset_dst + dst3_idx * dst3_stride + dst2_idx * dst2_stride + row_base;
+
+ var acc: array<f32, OUTPUTS_PER_WG>;
#ifdef MUL_ACC_FLOAT
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * VEC_SIZE; i < tile_size; i += THREADS_PER_OUTPUT * VEC_SIZE) {
- let a = src0[(idx_base + k_outer + i) / VEC_SIZE];
- let b = shared_vector[i / VEC_SIZE];
- local_sum += inner_dot(a, b);
+ let k_vec = params.k / VEC_SIZE;
+ let src1_idx_base_vec = src1_idx_base / VEC_SIZE;
+
+ // Each thread walks K, loads from the vector, and updates
+ // a small block of output rows held in registers.
+ for (var k = thread_id; k < k_vec; k += WG_SIZE) {
+ let x = src1[src1_idx_base_vec + k];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let src0_idx = (src0_batch_offset + output_row * params.stride_01) / VEC_SIZE + k;
+ acc[row] += inner_dot(src0[src0_idx], x);
+ }
+ }
}
- return local_sum;
-}
#endif
#ifdef MUL_ACC_Q4_0
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 18
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % 4;
+ for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4] = f32(src1[x_base + i + 16]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 18u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
- let q_byte_offset = block_byte_base + 2u + 2u * (block_offset + j);
- let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte(q_packed, k);
- let q_hi = (f32((q_byte >> 4) & 0xF) - 8.0) * d;
- let q_lo = (f32(q_byte & 0xF) - 8.0) * d;
- local_sum += q_lo * shared_vector[shmem_idx + j * 2 + k];
- local_sum += q_hi * shared_vector[shmem_idx + j * 2 + k + 16];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ var row_sum = 0.0;
+
+ let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_byte = get_byte(q_packed, byte_idx);
+ let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
+ let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
+ row_sum += q_lo * x_block[byte_idx];
+ row_sum += q_hi * x_block[byte_idx + 4u];
+ }
+ acc[row] += row_sum;
}
}
}
- return local_sum;
-}
#endif
#ifdef MUL_ACC_Q4_1
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 20
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4] = f32(src1[x_base + i + 16]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 20u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
- let m = f32(load_f16_at_src0(block_byte_base + 2u));
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
- let q_byte_offset = block_byte_base + 4u + 2u * (block_offset + j);
- let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte(q_packed, k);
- let q_hi = f32((q_byte >> 4) & 0xF) * d + m;
- let q_lo = f32(q_byte & 0xF) * d + m;
- local_sum += q_lo * shared_vector[shmem_idx + j * 2 + k];
- local_sum += q_hi * shared_vector[shmem_idx + j * 2 + k + 16];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
+ var row_sum = 0.0;
+
+ let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * thread_within_block);
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_byte = get_byte(q_packed, byte_idx);
+ let q_lo = f32(q_byte & 0xFu) * d + m;
+ let q_hi = f32((q_byte >> 4u) & 0xFu) * d + m;
+ row_sum += q_lo * x_block[byte_idx];
+ row_sum += q_hi * x_block[byte_idx + 4u];
+ }
+ acc[row] += row_sum;
}
}
}
- return local_sum;
-}
#endif
#ifdef MUL_ACC_Q5_0
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 22
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4] = f32(src1[x_base + i + 16]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 22u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
- let qh_packed = load_u32_at_src0(block_byte_base + 2u);
-
- for (var j = 0u; j < 2; j++) {
- let q_byte_offset = block_byte_base + 6u + 2u * (block_offset + j * 2u);
- let q_packed = load_u32_at_src0(q_byte_offset);
-
- let j_adjusted = j + (block_offset / 2u);
-
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte(q_packed, k);
-
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
- let q_hi = (f32(((q_byte >> 4) & 0xF) | qh_hi) - 16.0) * d;
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
- let q_lo = (f32((q_byte & 0xF) | qh_lo) - 16.0) * d;
-
- local_sum += q_lo * shared_vector[shmem_idx + j * 4 + k];
- local_sum += q_hi * shared_vector[shmem_idx + j * 4 + k + 16];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ let qh_packed = load_u32_at_src0(block_byte_base + 2u);
+ let q_packed = load_u32_at_src0(block_byte_base + 6u + 4u * thread_within_block);
+ let qh_shift = thread_within_block * 4u;
+ var row_sum = 0.0;
+
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_byte = get_byte(q_packed, byte_idx);
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
+ let q_lo = (f32((q_byte & 0xFu) | qh_lo) - 16.0) * d;
+ let q_hi = (f32(((q_byte >> 4u) & 0xFu) | qh_hi) - 16.0) * d;
+ row_sum += q_lo * x_block[byte_idx];
+ row_sum += q_hi * x_block[byte_idx + 4u];
+ }
+ acc[row] += row_sum;
}
-
}
}
- return local_sum;
-}
#endif
-
#ifdef MUL_ACC_Q5_1
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 24
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4] = f32(src1[x_base + i + 16]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 24u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
- let m = load_f16_at_src0(block_byte_base + 2u);
- let qh_packed = load_u32_at_src0(block_byte_base + 4u);
-
- for (var j = 0u; j < 2; j++) {
- let q_byte_offset = block_byte_base + 8u + 2u * (block_offset + j * 2u);
- let q_packed = load_u32_at_src0(q_byte_offset);
-
- let j_adjusted = j + (block_offset / 2u);
-
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte(q_packed, k);
-
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
- let q_hi = f32(((q_byte >> 4) & 0xF) | qh_hi) * d + f32(m);
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
- let q_lo = f32((q_byte & 0xF) | qh_lo) * d + f32(m);
-
- local_sum += q_lo * shared_vector[shmem_idx + j * 4 + k];
- local_sum += q_hi * shared_vector[shmem_idx + j * 4 + k + 16];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
+ let qh_packed = load_u32_at_src0(block_byte_base + 4u);
+ let q_packed = load_u32_at_src0(block_byte_base + 8u + 4u * thread_within_block);
+ let qh_shift = thread_within_block * 4u;
+ var row_sum = 0.0;
+
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_byte = get_byte(q_packed, byte_idx);
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
+ let q_lo = f32((q_byte & 0xFu) | qh_lo) * d + m;
+ let q_hi = f32(((q_byte >> 4u) & 0xFu) | qh_hi) * d + m;
+ row_sum += q_lo * x_block[byte_idx];
+ row_sum += q_hi * x_block[byte_idx + 4u];
+ }
+ acc[row] += row_sum;
}
-
}
}
- return local_sum;
-}
#endif
-
#ifdef MUL_ACC_Q8_0
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 34
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 34u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 2u;
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
-
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
- let q_byte_offset = block_byte_base + 2u + 2u * (block_offset + j);
- let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte_i32(q_packed, k);
- let q_val = f32(q_byte) * d;
- local_sum += q_val * shared_vector[shmem_idx + j * 2 + k];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ var row_sum = 0.0;
+
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
+ let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * (thread_within_block * 2u + packed_idx));
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_val = f32(get_byte_i32(q_packed, byte_idx)) * d;
+ row_sum += q_val * x_block[packed_idx * 4u + byte_idx];
+ }
+ }
+ acc[row] += row_sum;
}
}
}
- return local_sum;
-}
#endif
-
#ifdef MUL_ACC_Q8_1
+#define BLOCK_SIZE 32
+#define BLOCK_SIZE_BYTES 36
+#define THREADS_PER_BLOCK 4
+#define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
+
+ let num_blocks = params.k / BLOCK_SIZE;
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
+ var x_block: array<f32, ELEMS_PER_THREAD>;
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ }
-const BLOCK_SIZE = 32;
-const BLOCK_SIZE_BYTES = 36u;
-const NQ = 16u; // number of weights per thread
-const WEIGHTS_PER_F16 = 2u;
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
-
-fn mul_acc(tig:u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- var local_sum = 0.0;
- for (var i = tig * NQ; i < tile_size; i += THREADS_PER_OUTPUT * NQ) {
- let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let block_byte_base = (idx_base + k_outer / BLOCK_SIZE + blck_idx) * BLOCK_SIZE_BYTES;
- // each f16 contains offsets [block_offset, block_offset + 1] and [block_offset + 16, block_offset + 17]
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
- let d = f32(load_f16_at_src0(block_byte_base));
- let m = load_f16_at_src0(block_byte_base + 2u);
-
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
- let q_byte_offset = block_byte_base + 4u + 2u * (block_offset + j);
- let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k: u32 = 0; k < 4; k++) {
- let q_byte = get_byte_i32(q_packed, k);
- let q_val = f32(q_byte) * d + f32(m);
- local_sum += q_val * shared_vector[shmem_idx + j * 2 + k];
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+ let d = f32(load_f16_at_src0(block_byte_base));
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
+ var row_sum = 0.0;
+
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
+ let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * (thread_within_block * 2u + packed_idx));
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
+ let q_val = f32(get_byte_i32(q_packed, byte_idx)) * d + m;
+ row_sum += q_val * x_block[packed_idx * 4u + byte_idx];
+ }
+ }
+ acc[row] += row_sum;
}
}
}
- return local_sum;
-}
#endif
-#ifdef MUL_ACC_Q6_K
-
-const BLOCK_SIZE = 256u;
-const BLOCK_SIZE_BYTES = 210u;
-
-fn byte_of(v: u32, b: u32) -> u32 {
- return (v >> (b * 8u)) & 0xFFu;
-}
+#ifdef MUL_ACC_Q2_K
+#define BLOCK_SIZE 256
+#define BLOCK_SIZE_BYTES 84
+#define THREADS_PER_BLOCK 16
+
+ let tid = thread_id % THREADS_PER_BLOCK;
+ let block_group = thread_id / THREADS_PER_BLOCK;
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
+
+ let lane = tid / 2u;
+ let phase = tid % 2u;
+ let iq = lane / 4u;
+ let ir = lane % 4u;
+ let is = ir / 2u;
+
+ let y_offset = 128u * iq + 8u * ir + 4u * phase;
+ let sc0_byte = 8u * iq + is;
+ let sc2_byte = 8u * iq + is + 2u;
+ let sc4_byte = 8u * iq + is + 4u;
+ let sc6_byte = 8u * iq + is + 6u;
+ let qs_byte = 16u + (16u * iq + 4u * ir) * 2u + 4u * phase;
+
+ let num_blocks = params.k / BLOCK_SIZE;
+
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
+ var x_block: array<f32, 16>;
+ for (var i = 0u; i < 4u; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4u] = f32(src1[x_base + 32u + i]);
+ x_block[i + 8u] = f32(src1[x_base + 64u + i]);
+ x_block[i + 12u] = f32(src1[x_base + 96u + i]);
+ }
-fn sbyte_of(v: u32, b: u32) -> i32 {
- let raw = i32((v >> (b * 8u)) & 0xFFu);
- return select(raw, raw - 256, raw >= 128);
-}
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+
+ let dall = f32(load_f16_at_src0(block_byte_base + 80u));
+ let dmin = f32(load_f16_at_src0(block_byte_base + 82u)) * (1.0 / 16.0);
+
+ let sc0 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc0_byte), sc0_byte & 3u);
+ let sc2 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc2_byte), sc2_byte & 3u);
+ let sc4 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc4_byte), sc4_byte & 3u);
+ let sc6 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc6_byte), sc6_byte & 3u);
+
+ let q_u32 = load_u32_at_src0_aligned(block_byte_base + qs_byte);
+ let qs0 = q_u32 & 0xFFFFu;
+ let qs1 = q_u32 >> 16u;
+
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ var acc1 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ var acc2 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+
+ sumy[0] = x_block[0] + x_block[1] + x_block[2] + x_block[3];
+ sumy[1] = x_block[4] + x_block[5] + x_block[6] + x_block[7];
+ sumy[2] = x_block[8] + x_block[9] + x_block[10] + x_block[11];
+ sumy[3] = x_block[12] + x_block[13] + x_block[14] + x_block[15];
+
+ acc1[0] = x_block[0] * f32(qs0 & 0x0003u) + x_block[2] * f32(qs1 & 0x0003u);
+ acc2[0] = x_block[1] * f32(qs0 & 0x0300u) + x_block[3] * f32(qs1 & 0x0300u);
+ acc1[1] = x_block[4] * f32(qs0 & 0x000Cu) + x_block[6] * f32(qs1 & 0x000Cu);
+ acc2[1] = x_block[5] * f32(qs0 & 0x0C00u) + x_block[7] * f32(qs1 & 0x0C00u);
+ acc1[2] = x_block[8] * f32(qs0 & 0x0030u) + x_block[10] * f32(qs1 & 0x0030u);
+ acc2[2] = x_block[9] * f32(qs0 & 0x3000u) + x_block[11] * f32(qs1 & 0x3000u);
+ acc1[3] = x_block[12] * f32(qs0 & 0x00C0u) + x_block[14] * f32(qs1 & 0x00C0u);
+ acc2[3] = x_block[13] * f32(qs0 & 0xC000u) + x_block[15] * f32(qs1 & 0xC000u);
+
+ acc[row] += dall * ((acc1[0] + (1.0/256.0) * acc2[0]) * f32(sc0 & 0xFu) +
+ (acc1[1] + (1.0/256.0) * acc2[1]) * f32(sc2 & 0xFu) / 4.0 +
+ (acc1[2] + (1.0/256.0) * acc2[2]) * f32(sc4 & 0xFu) / 16.0 +
+ (acc1[3] + (1.0/256.0) * acc2[3]) * f32(sc6 & 0xFu) / 64.0)
+ - dmin * (sumy[0] * f32(sc0 & 0xF0u) + sumy[1] * f32(sc2 & 0xF0u) +
+ sumy[2] * f32(sc4 & 0xF0u) + sumy[3] * f32(sc6 & 0xF0u));
+ }
+ }
+ }
+#endif
-fn mul_acc(tig: u32, tile_size: u32, idx_base: u32, k_outer: u32) -> f32 {
- let tid = tig / 2u;
- let ix = tig % 2u;
- let ip = tid / 8u;
- let il = tid % 8u;
- let l0 = 4u * il;
- let is = 8u * ip + l0 / 16u;
- let y_offset = 128u * ip + l0;
- let q_offset_l = 64u * ip + l0;
- let q_offset_h = 32u * ip + l0;
+#ifdef MUL_ACC_Q3_K
+#define BLOCK_SIZE 256
+#define BLOCK_SIZE_BYTES 110
+#define THREADS_PER_BLOCK 16
- let nb = tile_size / BLOCK_SIZE;
- let k_block_start = k_outer / BLOCK_SIZE;
+ let tid = thread_id % THREADS_PER_BLOCK;
+ let block_group = thread_id / THREADS_PER_BLOCK;
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
- // Aligned scale byte position (is can be odd)
- let sc_base_byte = 192u + (is & ~3u);
- let sc_byte_pos = is & 3u;
+ let lane = tid / 2u;
+ let phase = tid % 2u;
+ let ip = lane / 4u;
+ let il = 2u * ((lane % 4u) / 2u);
+ let ir = lane % 2u;
+ let l0 = 8u * ir;
- var local_sum = 0.0;
+ let q_byte = 32u + 32u * ip + l0 + 16u * phase;
+ let h_byte = l0 + 16u * phase;
+ let y_offset = 128u * ip + 32u * il + l0 + 16u * phase;
- for (var i = ix; i < nb; i += 2u) {
- let bbase = (idx_base + k_block_start + i) * BLOCK_SIZE_BYTES;
+ let s_shift1 = 4u * ip;
+ let s_shift2 = s_shift1 + il;
- let d = f32(load_f16_at_src0(bbase + 208u));
+ let v1 = select(64.0, 4.0, il == 0u);
+ let v2 = 4.0 * v1;
+ let shift = 2u * il;
- let ql1_u32 = load_u32_at_src0(bbase + q_offset_l);
- let ql2_u32 = load_u32_at_src0(bbase + q_offset_l + 32u);
- let qh_u32 = load_u32_at_src0(bbase + 128u + q_offset_h);
- let sc_u32_0 = load_u32_at_src0(bbase + sc_base_byte);
- let sc_u32_1 = load_u32_at_src0(bbase + sc_base_byte + 4u);
+ var qm0: u32; var qm1: u32; var qm2: u32; var qm3: u32;
+ if (il == 0u) {
+ qm0 = 0x0003u; qm1 = 0x0300u; qm2 = 0x000Cu; qm3 = 0x0C00u;
+ } else {
+ qm0 = 0x0030u; qm1 = 0x3000u; qm2 = 0x00C0u; qm3 = 0xC000u;
+ }
- let sc0 = sbyte_of(sc_u32_0, sc_byte_pos);
- let sc2 = sbyte_of(sc_u32_0, sc_byte_pos + 2u);
- let sc4 = sbyte_of(sc_u32_1, sc_byte_pos);
- let sc6 = sbyte_of(sc_u32_1, sc_byte_pos + 2u);
+ let mm_idx = 2u * ip + il / 2u;
+ var hm0: u32; var hm1: u32; var hm2: u32; var hm3: u32;
+ switch (mm_idx) {
+ case 0u: { hm0=0x0001u; hm1=0x0100u; hm2=0x0002u; hm3=0x0200u; }
+ case 1u: { hm0=0x0004u; hm1=0x0400u; hm2=0x0008u; hm3=0x0800u; }
+ case 2u: { hm0=0x0010u; hm1=0x1000u; hm2=0x0020u; hm3=0x2000u; }
+ default: { hm0=0x0040u; hm1=0x4000u; hm2=0x0080u; hm3=0x8000u; }
+ }
- var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ let num_blocks = params.k / BLOCK_SIZE;
- for (var l = 0u; l < 4u; l++) {
- let y_base = i * BLOCK_SIZE + y_offset + l;
- let yl0 = f32(shared_vector[y_base]);
- let yl1 = f32(shared_vector[y_base + 32u]);
- let yl2 = f32(shared_vector[y_base + 64u]);
- let yl3 = f32(shared_vector[y_base + 96u]);
-
- let q1b = byte_of(ql1_u32, l);
- let q2b = byte_of(ql2_u32, l);
- let qhb = byte_of(qh_u32, l);
-
- let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
- let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
- let dq2 = f32(i32((q1b >> 4u) | ((qhb & 0x30u) )) - 32);
- let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
-
- sums[0] += yl0 * dq0;
- sums[1] += yl1 * dq1;
- sums[2] += yl2 * dq2;
- sums[3] += yl3 * dq3;
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
+ var x_block: array<f32, 16>;
+ for (var i = 0u; i < 8u; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 8u] = f32(src1[x_base + 32u + i]);
}
- local_sum += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
- sums[2] * f32(sc4) + sums[3] * f32(sc6));
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+
+ let d = f32(load_f16_at_src0(block_byte_base + 108u));
+ let a_base = 96u;
+ let a_il0 = load_u16_at_src0(block_byte_base + a_base + il * 2u);
+ let a_il1 = load_u16_at_src0(block_byte_base + a_base + (il + 1u) * 2u);
+ let a_4 = load_u16_at_src0(block_byte_base + a_base + 8u);
+ let a_5 = load_u16_at_src0(block_byte_base + a_base + 10u);
+
+ var scales32 = a_4 | (a_5 << 16u);
+ let aux32 = ((scales32 >> s_shift2) << 4u) & 0x30303030u;
+ scales32 = a_il0 | (a_il1 << 16u);
+ scales32 = ((scales32 >> s_shift1) & 0x0F0F0F0Fu) | aux32;
+
+ let scale0 = f32(i32(byte_of(scales32, phase + 0u)) - 32);
+ let scale1 = f32(i32(byte_of(scales32, phase + 2u)) - 32);
+
+ let q_u32_0 = load_u32_at_src0(block_byte_base + q_byte + 0u);
+ let q_u32_1 = load_u32_at_src0(block_byte_base + q_byte + 4u);
+ let h_u32_0 = load_u32_at_src0(block_byte_base + h_byte + 0u);
+ let h_u32_1 = load_u32_at_src0(block_byte_base + h_byte + 4u);
+
+ var s1 = 0.0; var s2 = 0.0; var s3 = 0.0;
+ var s4 = 0.0; var s5 = 0.0; var s6 = 0.0;
+
+ for (var l = 0u; l < 8u; l += 2u) {
+ let q_u32 = select(q_u32_0, q_u32_1, l >= 4u);
+ let qs = select(q_u32 & 0xFFFFu, q_u32 >> 16u, (l & 2u) != 0u);
+ let h_u32 = select(h_u32_0, h_u32_1, l >= 4u);
+ let hv = select(h_u32 & 0xFFFFu, h_u32 >> 16u, (l & 2u) != 0u);
+
+ s1 += x_block[l + 0u] * f32(qs & qm0);
+ s2 += x_block[l + 1u] * f32(qs & qm1);
+ s3 += select(0.0, x_block[l + 0u], (hv & hm0) == 0u) +
+ select(0.0, x_block[l + 1u], (hv & hm1) == 0u);
+ s4 += x_block[l + 8u] * f32(qs & qm2);
+ s5 += x_block[l + 9u] * f32(qs & qm3);
+ s6 += select(0.0, x_block[l + 8u], (hv & hm2) == 0u) +
+ select(0.0, x_block[l + 9u], (hv & hm3) == 0u);
+ }
+
+ let d1 = d * (s1 + (1.0/256.0) * s2 - s3 * v1);
+ let d2 = d * (s4 + (1.0/256.0) * s5 - s6 * v2);
+ acc[row] += (d1 * scale0 + 0.25 * d2 * scale1) / f32(1u << shift);
+ }
+ }
}
-
- return local_sum;
-}
#endif
-struct MulMatParams {
- offset_src0: u32,
- offset_src1: u32,
- offset_dst: u32,
- m: u32,
- n: u32,
- k: u32,
- stride_01: u32,
- stride_11: u32,
- stride_02: u32,
- stride_12: u32,
- stride_03: u32,
- stride_13: u32,
- bs02: u32,
- bs03: u32,
- broadcast2: u32,
- broadcast3: u32
-};
-
-// SRC0_TYPE and SRC1_TYPE are defined in mul_mat_decls, which is included
-@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>; // M rows, K columns
-@group(0) @binding(1) var<storage, read_write> src1: array<SRC1_TYPE>; // K rows, N columns (transposed)
-@group(0) @binding(2) var<storage, read_write> dst: array<DST_TYPE>; // M rows, N columns (transposed)
-
-@group(0) @binding(3) var<uniform> params: MulMatParams;
-
-const THREADS_PER_OUTPUT = WG_SIZE / OUTPUTS_PER_WG;
+#ifdef MUL_ACC_Q4_K
+#define BLOCK_SIZE 256
+#define BLOCK_SIZE_BYTES 144
+#define THREADS_PER_BLOCK 16
+
+ let tid = thread_id % THREADS_PER_BLOCK;
+ let block_group = thread_id / THREADS_PER_BLOCK;
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
+
+ let il = tid / 4u;
+ let ir = tid % 4u;
+ let im = il / 2u;
+ let in = il % 2u;
+ let l0 = 4u * (2u * ir + in);
+
+ let y_offset = 64u * im + l0;
+ let q_offset = 32u * im + l0;
+ let sc0_byte = 4u + im * 2u;
+ let sc2_byte = 4u + (im + 2u) * 2u;
+ let sc4_byte = 4u + (im + 4u) * 2u;
+
+ let num_blocks = params.k / BLOCK_SIZE;
+
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
+ var x_block: array<f32, 16>;
+ for (var i = 0u; i < 4u; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4u] = f32(src1[x_base + 32u + i]);
+ x_block[i + 8u] = f32(src1[x_base + 128u + i]);
+ x_block[i + 12u] = f32(src1[x_base + 160u + i]);
+ }
-// Shared memory for collaborative loading and reduction
-var<workgroup> shared_vector: array<SRC1_TYPE, TILE_K/VEC_SIZE>; // Cache vector tile
-var<workgroup> partial_sums: array<f32, WG_SIZE>; // For reduction
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+
+ let d = f32(load_f16_at_src0(block_byte_base + 0u));
+ let dmin = f32(load_f16_at_src0(block_byte_base + 2u));
+
+ let sc0_u32 = load_u32_at_src0_aligned(block_byte_base + sc0_byte);
+ let sc0 = select(sc0_u32 & 0xFFFFu, sc0_u32 >> 16u, (sc0_byte & 2u) != 0u);
+ let sc2_u32 = load_u32_at_src0_aligned(block_byte_base + sc2_byte);
+ let sc2 = select(sc2_u32 & 0xFFFFu, sc2_u32 >> 16u, (sc2_byte & 2u) != 0u);
+ let sc4_u32 = load_u32_at_src0_aligned(block_byte_base + sc4_byte);
+ let sc4 = select(sc4_u32 & 0xFFFFu, sc4_u32 >> 16u, (sc4_byte & 2u) != 0u);
+
+ let sc16_0 = sc0 & 0x3F3Fu;
+ let sc16_1 = sc2 & 0x3F3Fu;
+ let sc16_2 = (sc4 & 0x0F0Fu) | ((sc0 & 0xC0C0u) >> 2u);
+ let sc16_3 = ((sc4 >> 4u) & 0x0F0Fu) | ((sc2 & 0xC0C0u) >> 2u);
+
+ let scale0 = f32(sc16_0 & 0xFFu);
+ let scale1 = f32((sc16_0 >> 8u) & 0xFFu);
+ let min0 = f32(sc16_1 & 0xFFu);
+ let min1 = f32((sc16_1 >> 8u) & 0xFFu);
+ let scale2 = f32(sc16_2 & 0xFFu);
+ let scale3 = f32((sc16_2 >> 8u) & 0xFFu);
+ let min2 = f32(sc16_3 & 0xFFu);
+ let min3 = f32((sc16_3 >> 8u) & 0xFFu);
+
+ let q1_u32 = load_u32_at_src0_aligned(block_byte_base + 16u + q_offset);
+ let q2_u32 = load_u32_at_src0_aligned(block_byte_base + 80u + q_offset);
+
+ var dot = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ var sumx = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ for (var i = 0u; i < 4u; i++) {
+ let q1b = byte_of(q1_u32, i);
+ let q2b = byte_of(q2_u32, i);
+ dot[0] += x_block[i] * f32(q1b & 0x0Fu);
+ dot[1] += x_block[i + 4u] * f32(q1b >> 4u);
+ dot[2] += x_block[i + 8u] * f32(q2b & 0x0Fu);
+ dot[3] += x_block[i + 12u] * f32(q2b >> 4u);
+ sumx[0] += x_block[i];
+ sumx[1] += x_block[i + 4u];
+ sumx[2] += x_block[i + 8u];
+ sumx[3] += x_block[i + 12u];
+ }
+
+ acc[row] += d * (dot[0] * scale0 + dot[1] * scale1 + dot[2] * scale2 + dot[3] * scale3)
+ - dmin * (sumx[0] * min0 + sumx[1] * min1 + sumx[2] * min2 + sumx[3] * min3);
+ }
+ }
+ }
+#endif
-@compute @workgroup_size(WG_SIZE)
-fn main(
- @builtin(local_invocation_id) local_id: vec3<u32>,
- @builtin(workgroup_id) wg_id: vec3<u32>,
- @builtin(num_workgroups) num_wg: vec3<u32>) {
- let thread_id = local_id.x;
+#ifdef MUL_ACC_Q5_K
+#define BLOCK_SIZE 256
+#define BLOCK_SIZE_BYTES 176
+#define THREADS_PER_BLOCK 16
+
+ let tid = thread_id % THREADS_PER_BLOCK;
+ let block_group = thread_id / THREADS_PER_BLOCK;
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
+
+ let il = tid / 4u;
+ let ir = tid % 4u;
+ let im = il / 2u;
+ let in = il % 2u;
+ let l0 = 4u * (2u * ir + in);
+
+ let y_offset = 64u * im + l0;
+ let q_offset = 48u + 32u * im + l0;
+ let qh_offset = 16u + 8u * ir + 4u * in;
+ let sc0_byte = 4u + im * 2u;
+ let sc2_byte = 4u + (im + 2u) * 2u;
+ let sc4_byte = 4u + (im + 4u) * 2u;
+
+ let hm1 = 1u << (2u * im);
+ let hm2 = hm1 << 1u;
+ let hm3 = hm1 << 4u;
+ let hm4 = hm2 << 4u;
+
+ let num_blocks = params.k / BLOCK_SIZE;
+
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
+ var x_block: array<f32, 16>;
+ for (var i = 0u; i < 4u; i++) {
+ x_block[i] = f32(src1[x_base + i]);
+ x_block[i + 4u] = f32(src1[x_base + 32u + i]);
+ x_block[i + 8u] = f32(src1[x_base + 128u + i]);
+ x_block[i + 12u] = f32(src1[x_base + 160u + i]);
+ }
- // Handle batch dimensions
- let total_batches = params.bs02 * params.broadcast2 * params.bs03 * params.broadcast3;
- let wg_linear = wg_id.y * num_wg.x + wg_id.x;
- let output_groups = (params.m + OUTPUTS_PER_WG - 1u) / OUTPUTS_PER_WG;
- let batch_idx = wg_linear / output_groups;
- if (batch_idx >= total_batches) {
- return;
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+
+ let d = f32(load_f16_at_src0(block_byte_base + 0u));
+ let dmin = f32(load_f16_at_src0(block_byte_base + 2u));
+
+ let sc0_u32 = load_u32_at_src0_aligned(block_byte_base + sc0_byte);
+ let sc0 = select(sc0_u32 & 0xFFFFu, sc0_u32 >> 16u, (sc0_byte & 2u) != 0u);
+ let sc2_u32 = load_u32_at_src0_aligned(block_byte_base + sc2_byte);
+ let sc2 = select(sc2_u32 & 0xFFFFu, sc2_u32 >> 16u, (sc2_byte & 2u) != 0u);
+ let sc4_u32 = load_u32_at_src0_aligned(block_byte_base + sc4_byte);
+ let sc4 = select(sc4_u32 & 0xFFFFu, sc4_u32 >> 16u, (sc4_byte & 2u) != 0u);
+
+ let sc16_0 = sc0 & 0x3F3Fu;
+ let sc16_1 = sc2 & 0x3F3Fu;
+ let sc16_2 = (sc4 & 0x0F0Fu) | ((sc0 & 0xC0C0u) >> 2u);
+ let sc16_3 = ((sc4 >> 4u) & 0x0F0Fu) | ((sc2 & 0xC0C0u) >> 2u);
+
+ let f0 = f32(sc16_0 & 0xFFu);
+ let f1 = f32((sc16_0 >> 8u) & 0xFFu);
+ let m0 = f32(sc16_1 & 0xFFu);
+ let m1 = f32((sc16_1 >> 8u) & 0xFFu);
+ let f4 = f32(sc16_2 & 0xFFu);
+ let f5 = f32((sc16_2 >> 8u) & 0xFFu);
+ let m4 = f32(sc16_3 & 0xFFu);
+ let m5 = f32((sc16_3 >> 8u) & 0xFFu);
+
+ let q1_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset);
+ let q2_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset + 64u);
+ let qh_u32 = load_u32_at_src0_aligned(block_byte_base + qh_offset);
+
+ var vals = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+ for (var i = 0u; i < 4u; i++) {
+ let q1b = byte_of(q1_u32, i);
+ let q2b = byte_of(q2_u32, i);
+ let qhb = byte_of(qh_u32, i);
+
+ let yl0 = x_block[i];
+ let yl8 = x_block[i + 4u];
+ let yh0 = x_block[i + 8u];
+ let yh8 = x_block[i + 12u];
+
+ sumy[0] += yl0;
+ sumy[1] += yl8;
+ sumy[2] += yh0;
+ sumy[3] += yh8;
+
+ let q0 = f32((q1b & 0x0Fu) | select(0u, 0x10u, (qhb & hm1) != 0u));
+ let q1 = f32((q1b >> 4u) | select(0u, 0x10u, (qhb & hm2) != 0u));
+ let q2 = f32((q2b & 0x0Fu) | select(0u, 0x10u, (qhb & hm3) != 0u));
+ let q3 = f32((q2b >> 4u) | select(0u, 0x10u, (qhb & hm4) != 0u));
+
+ vals[0] += yl0 * q0;
+ vals[1] += yl8 * q1;
+ vals[2] += yh0 * q2;
+ vals[3] += yh8 * q3;
+ }
+
+ acc[row] += d * (f0 * vals[0] + f1 * vals[1] + f4 * vals[2] + f5 * vals[3])
+ - dmin * (sumy[0] * m0 + sumy[1] * m1 +
+ sumy[2] * m4 + sumy[3] * m5);
+ }
+ }
}
+#endif
- // Which of the outputs does this thread belong to?
- let thread_group = thread_id / THREADS_PER_OUTPUT;
- let thread_in_group = thread_id % THREADS_PER_OUTPUT;
+#ifdef MUL_ACC_Q6_K
+#define BLOCK_SIZE 256
+#define BLOCK_SIZE_BYTES 210
+#define THREADS_PER_BLOCK 16
- // Each workgroup computes OUTPUTS_PER_WG consecutive outputs
- let output_row = (wg_linear % output_groups) * OUTPUTS_PER_WG + thread_group;
+ let tid = thread_id % THREADS_PER_BLOCK;
+ let block_group = thread_id / THREADS_PER_BLOCK;
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
- let dst2_stride = params.m * params.n;
- let dst2_idx = batch_idx % (params.bs02 * params.broadcast2);
- let dst3_stride = dst2_stride * params.bs02 * params.broadcast2;
- let dst3_idx = batch_idx / (params.bs02 * params.broadcast2);
- let src03_idx = dst3_idx / params.broadcast3;
- let src13_idx = dst3_idx;
- let src02_idx = dst2_idx / params.broadcast2;
- let src12_idx = dst2_idx;
+ let ip = tid / 8u;
+ let il = tid % 8u;
+ let l0 = 4u * il;
+ let is = 8u * ip + l0 / 16u;
- let src0_idx_base = params.offset_src0 + src03_idx * params.stride_03 + src02_idx * params.stride_02 + output_row * params.stride_01;
- let src1_idx_base = params.offset_src1 + src13_idx * params.stride_13 + src12_idx * params.stride_12;
- let dst_idx = params.offset_dst + dst3_idx * dst3_stride + dst2_idx * dst2_stride + output_row;
+ let y_offset = 128u * ip + l0;
+ let q_offset_l = 64u * ip + l0;
+ let q_offset_h = 32u * ip + l0;
- var local_sum = 0.0;
+ let num_blocks = params.k / BLOCK_SIZE;
+ let sc_base_byte = 192u + (is & ~3u);
+ let sc_byte_pos = is & 3u;
+
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
+ var x_block: array<f32, 16>;
+ for (var l = 0u; l < 4u; l++) {
+ x_block[l] = f32(src1[x_base + l]);
+ x_block[l + 4u] = f32(src1[x_base + 32u + l]);
+ x_block[l + 8u] = f32(src1[x_base + 64u + l]);
+ x_block[l + 12u] = f32(src1[x_base + 96u + l]);
+ }
- // Each thread processes multiple K elements and accumulates
- for (var k_tile = 0u; k_tile < params.k; k_tile += TILE_K) {
- let tile_size = min(TILE_K, params.k - k_tile);
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let output_row = row_base + row;
+ if (output_row < params.m) {
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
+
+ let d = f32(load_f16_at_src0(block_byte_base + 208u));
+ let ql1_u32 = load_u32_at_src0(block_byte_base + q_offset_l);
+ let ql2_u32 = load_u32_at_src0(block_byte_base + q_offset_l + 32u);
+ let qh_u32 = load_u32_at_src0(block_byte_base + 128u + q_offset_h);
+ let sc_u32_0 = load_u32_at_src0(block_byte_base + sc_base_byte);
+ let sc_u32_1 = load_u32_at_src0(block_byte_base + sc_base_byte + 4u);
+
+ let sc0 = sbyte_of(sc_u32_0, sc_byte_pos);
+ let sc2 = sbyte_of(sc_u32_0, sc_byte_pos + 2u);
+ let sc4 = sbyte_of(sc_u32_1, sc_byte_pos);
+ let sc6 = sbyte_of(sc_u32_1, sc_byte_pos + 2u);
+
+ var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
+
+ for (var l = 0u; l < 4u; l++) {
+ let q1b = byte_of(ql1_u32, l);
+ let q2b = byte_of(ql2_u32, l);
+ let qhb = byte_of(qh_u32, l);
+
+ let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
+ let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
+ let dq2 = f32(i32((q1b >> 4u) | (qhb & 0x30u)) - 32);
+ let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
+
+ sums[0] += x_block[l] * dq0;
+ sums[1] += x_block[l + 4u] * dq1;
+ sums[2] += x_block[l + 8u] * dq2;
+ sums[3] += x_block[l + 12u] * dq3;
+ }
+
+ acc[row] += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
+ sums[2] * f32(sc4) + sums[3] * f32(sc6));
+ }
+ }
+ }
+#endif
- // Cooperatively load vector tile into shared memory (all threads)
- for (var i = thread_id * VEC_SIZE; i < tile_size; i += WG_SIZE * VEC_SIZE) {
- shared_vector[i / VEC_SIZE] = src1[(src1_idx_base + k_tile + i) / VEC_SIZE];
+#ifdef USE_SUBGROUP_REDUCTION
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ let subgroup_total = subgroupAdd(acc[row]);
+ if (subgroup_invocation_id == 0u) {
+ partial_sums[partial_index(row, subgroup_id)] = subgroup_total;
}
+ }
- workgroupBarrier();
+ workgroupBarrier();
- if (output_row < params.m) {
- local_sum += mul_acc(thread_in_group, tile_size, src0_idx_base, k_tile);
+ for (var row = subgroup_id; (row < OUTPUTS_PER_WG) && (row_base + row < params.m); row += num_subgroups) {
+ let output_row = row_base + row;
+ var row_acc = 0.0f;
+ for (var k = subgroup_invocation_id; k < num_subgroups; k += subgroup_size) {
+ row_acc += partial_sums[partial_index(row, k)];
}
+ let row_total = subgroupAdd(row_acc);
+ if (subgroup_invocation_id == 0) {
+ dst[dst_idx_base + row] = row_total;
+ }
+ }
+#endif
- workgroupBarrier();
+#ifdef USE_WORKGROUP_REDUCTION
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ partial_sums[partial_index(row, thread_id)] = acc[row];
}
- // Store partial sums and reduce within each partition
- partial_sums[thread_id] = local_sum;
workgroupBarrier();
- let group_base = thread_group * THREADS_PER_OUTPUT;
- let thread_base = group_base + thread_in_group;
- var offset: u32 = THREADS_PER_OUTPUT / 2;
- while (offset > 0) {
- if (thread_in_group < offset) {
- partial_sums[thread_base] += partial_sums[thread_base + offset];
+
+ var stride = WG_SIZE / 2u;
+
+ while (stride > 0) {
+ if (thread_id < stride) {
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
+ partial_sums[partial_index(row, thread_id)] += partial_sums[partial_index(row, thread_id + stride)];
+ }
}
- offset = offset / 2;
+
workgroupBarrier();
+ stride = stride / 2;
}
- // Store back to global memory
- if (output_row < params.m && thread_group % VEC_SIZE == 0 && thread_in_group == 0) {
- dst[dst_idx / VEC_SIZE] = store_val(group_base);
+ if (thread_id < OUTPUTS_PER_WG) {
+ let output_row = row_base + thread_id;
+ if (output_row < params.m) {
+ dst[dst_idx_base + thread_id] = partial_sums[partial_index(thread_id, 0)];
+ }
}
+#endif
}