}
};
+/* Add_Id */
+
+struct ggml_webgpu_add_id_pipeline_key {
+ bool inplace;
+
+ bool operator==(const ggml_webgpu_add_id_pipeline_key & other) const { return inplace == other.inplace; }
+};
+
+struct ggml_webgpu_add_id_pipeline_key_hash {
+ size_t operator()(const ggml_webgpu_add_id_pipeline_key & key) const {
+ size_t seed = 0;
+ ggml_webgpu_hash_combine(seed, key.inplace);
+ return seed;
+ }
+};
+
/** Unary **/
struct ggml_webgpu_unary_pipeline_key {
std::unordered_map<ggml_webgpu_pad_pipeline_key, webgpu_pipeline, ggml_webgpu_pad_pipeline_key_hash>
pad_pipelines; // circular/non-circular
std::unordered_map<ggml_webgpu_binary_pipeline_key, webgpu_pipeline, ggml_webgpu_binary_pipeline_key_hash>
- binary_pipelines; // type/op/inplace/overlap
+ binary_pipelines; // type/op/inplace/overlap/src_overlap
+ std::unordered_map<ggml_webgpu_add_id_pipeline_key, webgpu_pipeline, ggml_webgpu_add_id_pipeline_key_hash>
+ add_id_pipelines; // inplace
std::unordered_map<ggml_webgpu_concat_pipeline_key, webgpu_pipeline, ggml_webgpu_concat_pipeline_key_hash>
concat_pipelines; // type
std::unordered_map<ggml_webgpu_repeat_pipeline_key, webgpu_pipeline, ggml_webgpu_repeat_pipeline_key_hash>
case GGML_TYPE_IQ3_S:
case GGML_TYPE_IQ1_S:
case GGML_TYPE_IQ4_NL:
+ case GGML_TYPE_MXFP4:
{
// Quantized types using u32 buffers for portability.
defines.push_back("SRC_TYPE=u32");
defines.push_back(type_upper + "_SCALE_MIN");
defines.push_back(type_upper + "_TABLES");
defines.push_back(type_upper + "_GRID");
+ defines.push_back(type_upper + "_LUT");
variant += "_";
variant += type_str;
if (key.src_type == GGML_TYPE_Q1_0) {
defines.push_back("BLOCK_SIZE=128u");
} else if ((key.src_type >= GGML_TYPE_Q4_0 && key.src_type <= GGML_TYPE_Q8_1) ||
- key.src_type == GGML_TYPE_IQ4_NL) {
+ key.src_type == GGML_TYPE_IQ4_NL || key.src_type == GGML_TYPE_MXFP4) {
defines.push_back("BLOCK_SIZE=32u");
} else if (key.src_type >= GGML_TYPE_Q2_K) {
defines.push_back("BLOCK_SIZE=256u");
defines.push_back(type_upper + "_GRID");
defines.push_back(type_upper + "_TABLES");
break;
+ case GGML_TYPE_MXFP4:
+ defines.push_back(type_upper + "_LUT");
+ break;
default:
break;
}
defines.push_back(type_upper + "_GRID");
defines.push_back(type_upper + "_TABLES");
break;
+ case GGML_TYPE_MXFP4:
+ defines.push_back(type_upper + "_LUT");
+ break;
default:
break;
}
case GGML_TYPE_IQ3_S:
case GGML_TYPE_IQ1_S:
case GGML_TYPE_IQ4_NL:
+ case GGML_TYPE_MXFP4:
{
// Quantized types using u32 buffers for portability.
defines.push_back("SRC0_TYPE=u32");
defines.push_back(type_upper + "_GRID");
defines.push_back(type_upper + "_TABLES");
break;
+ case GGML_TYPE_MXFP4:
+ defines.push_back(type_upper + "_LUT");
+ break;
default:
break;
}
defines.push_back(type_upper + "_GRID");
defines.push_back(type_upper + "_TABLES");
break;
+ case GGML_TYPE_MXFP4:
+ defines.push_back(type_upper + "_LUT");
+ break;
default:
break;
}
return binary_pipelines[key];
}
+ webgpu_pipeline get_add_id_pipeline(const ggml_webgpu_shader_lib_context & context) {
+ ggml_webgpu_add_id_pipeline_key key = {};
+ key.inplace = ggml_webgpu_tensor_equal(context.src0, context.dst);
+
+ auto it = add_id_pipelines.find(key);
+ if (it != add_id_pipelines.end()) {
+ return it->second;
+ }
+
+ std::vector<std::string> defines;
+ std::string variant = "add_id";
+ const char * shader_src = wgsl_add_id;
+
+ if (key.inplace) {
+ defines.push_back("INPLACE");
+ variant += "_inplace";
+ }
+
+ defines.push_back(std::string("WG_SIZE=") + std::to_string(context.max_wg_size));
+
+ auto processed = preprocessor.preprocess(shader_src, defines);
+ auto pipeline_decisions = std::make_shared<ggml_webgpu_generic_shader_decisions>();
+ pipeline_decisions->wg_size = context.max_wg_size;
+ pipeline_decisions->inplace = key.inplace;
+
+ webgpu_pipeline pipeline = ggml_webgpu_create_pipeline(device, processed, variant);
+ pipeline.context = pipeline_decisions;
+ add_id_pipelines[key] = pipeline;
+ return pipeline;
+ }
+
webgpu_pipeline get_concat_pipeline(const ggml_webgpu_shader_lib_context & context) {
ggml_webgpu_concat_pipeline_key key = {};
key.type = context.dst->type;
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
override BLOCKS_K = TILE_K/BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
+const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+
+ let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k = 0u; k < 4u; k++) {
+
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte(q_packed, k);
let q_hi = (f16((q_byte >> 4) & 0xF) - 8.0) * d;
let q_lo = (f16(q_byte & 0xF) - 8.0) * d;
- shmem[shmem_idx + j * 2 + k] = q_lo;
- shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
}
}
}
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
override BLOCKS_K = TILE_K/BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
+const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+
+ let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k = 0u; k < 4u; k++) {
+
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte(q_packed, k);
let q_lo = f16(q_byte & 0xF) * d + m;
let q_hi = f16((q_byte >> 4) & 0xF) * d + m;
- shmem[shmem_idx + j * 2 + k] = q_lo;
- shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
}
}
}
#endif // INIT_SRC0_SHMEM_Q4_1
#ifdef INIT_SRC0_SHMEM_Q5_0
-// 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
const BLOCK_SIZE = 32u;
const BLOCK_SIZE_BYTES = 22u;
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
// tile_k is defined as 32u, so blocks_k ends up being 1 always
override BLOCKS_K = TILE_K / BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
+const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+ let q_byte_offset = block_byte_base + 6u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- let j_adjusted = j + (block_offset / 2u);
-
-
- for (var k = 0u; k < 4u; k++) {
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte(q_packed, k);
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
+ let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
+ let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi) - 16.0) * d;
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
+ let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
let q_lo = (f16((q_byte & 0xF) | qh_lo) - 16.0) * d;
-
- shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
- shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
}
}
}
#endif // INIT_SRC0_SHMEM_Q5_0
#ifdef INIT_SRC0_SHMEM_Q5_1
-// 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
const BLOCK_SIZE = 32u;
const BLOCK_SIZE_BYTES = 24u;
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
-// tile_k is defined as 32u, so blocks_k ends up being 1 always
override BLOCKS_K = TILE_K / BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
+const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+ let q_byte_offset = block_byte_base + 8u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- let j_adjusted = j + (block_offset / 2u);
-
-
- for (var k = 0u; k < 4u; k++) {
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte(q_packed, k);
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
- let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi)) * d + m;
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
- let q_lo = (f16((q_byte & 0xF) | qh_lo)) * d + m;
-
- shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
- shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
+ let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
+ let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
+ let q_hi = f16(((q_byte >> 4) & 0xF) | qh_hi) * d + m;
+ let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
+ let q_lo = f16((q_byte & 0xF) | qh_lo) * d + m;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
}
}
}
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
override BLOCKS_K = TILE_K/BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread
+const BYTES_PER_THREAD = 16u; // NQ(16) weights use 16 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+ let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k = 0u; k < 4u; k++) {
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte_i32(q_packed, k);
let q_val = f16(q_byte) * d;
- shmem[shmem_idx + j * 2 + k] = q_val;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
}
}
}
// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
override BLOCKS_K = TILE_K/BLOCK_SIZE;
const NQ = 16u;
-const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
-const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread, 2 threads per block
+const BYTES_PER_THREAD = 16u; // NQ(16) weights use 16 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
let blck_idx = i / BLOCK_SIZE;
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
let tile_m = blck_idx / BLOCKS_K;
let global_m = offset_m + tile_m;
let block_k = blck_idx % BLOCKS_K;
- let global_k = k_outer / BLOCK_SIZE + block_k;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
let d = 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);
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+ let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
let q_packed = load_u32_at_src0(q_byte_offset);
- for (var k = 0u; k < 4u; k++) {
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
let q_byte = get_byte_i32(q_packed, k);
let q_val = f16(q_byte) * d + m;
- shmem[shmem_idx + j * 2 + k] = q_val;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
}
}
}
}
}
#endif // INIT_SRC0_SHMEM_IQ3_S
+
+#ifdef INIT_SRC0_SHMEM_MXFP4
+const BLOCK_SIZE = 32u;
+const BLOCK_SIZE_BYTES = 17u;
+// the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
+override BLOCKS_K = TILE_K/BLOCK_SIZE;
+const NQ = 16u;
+const BYTES_PER_THREAD = 8u; // NQ(16) weights uses 8 bytes of q
+const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
+
+fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
+ for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
+ let blck_idx = i / BLOCK_SIZE;
+ let block_offset = (i % BLOCK_SIZE) / NQ;
+ let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
+
+ let tile_m = blck_idx / BLOCKS_K;
+ let global_m = offset_m + tile_m;
+ let block_k = blck_idx % BLOCKS_K;
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
+
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
+ let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
+ let eu8 = get_byte(load_u32_at_src0(block_byte_base), 0);
+ let e = ldexp(1.0, i32(eu8) - 128);
+
+ // store NQ(16) weights
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
+
+ let q_byte_offset = block_byte_base + 1u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
+ let q_packed = load_u32_at_src0(q_byte_offset);
+
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
+ let q_byte = get_byte(q_packed, k);
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * e;
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xF]) * e;
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = f16(q_lo);
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = f16(q_hi);
+ }
+ }
+ }
+ }
+}
+#endif // INIT_SRC0_SHMEM_MXFP4