From 7df5fa8e62fa96bbfd7caf2eecb9252f84930c2c Mon Sep 17 00:00:00 2001 From: Masashi Yoshimura Date: Wed, 19 Aug 2026 16:29:33 +0900 Subject: [PATCH] ggml-webgpu: add mulmat with overlapping src0/src1 (e.g., for minimax-01) (llama/27321) --- .../ggml-webgpu/ggml-webgpu-shader-lib.hpp | 26 ++++- ggml/src/ggml-webgpu/ggml-webgpu.cpp | 63 +++++++----- .../wgsl-shaders/common_decls.tmpl | 14 ++- .../wgsl-shaders/mul_mat_decls.tmpl | 11 ++- .../wgsl-shaders/mul_mat_reg_tile.wgsl | 16 ++- .../wgsl-shaders/mul_mat_subgroup_matrix.wgsl | 14 ++- .../ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl | 18 +++- .../wgsl-shaders/mul_mat_vec_acc.tmpl | 97 ++++++++++--------- 8 files changed, 171 insertions(+), 88 deletions(-) diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 0604e1c2b..7a67ccf4f 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -954,10 +954,11 @@ struct ggml_webgpu_mul_mat_vec_pipeline_key { int vectorized; uint32_t num_cols; bool use_mmvq; + bool src_overlap; bool operator==(const ggml_webgpu_mul_mat_vec_pipeline_key & other) const { return src0_type == other.src0_type && src1_type == other.src1_type && vectorized == other.vectorized && - num_cols == other.num_cols && use_mmvq == other.use_mmvq; + num_cols == other.num_cols && use_mmvq == other.use_mmvq && src_overlap == other.src_overlap; } }; @@ -969,6 +970,7 @@ struct ggml_webgpu_mul_mat_vec_pipeline_key_hash { ggml_webgpu_hash_combine(seed, key.vectorized); ggml_webgpu_hash_combine(seed, key.num_cols); ggml_webgpu_hash_combine(seed, key.use_mmvq); + ggml_webgpu_hash_combine(seed, key.src_overlap); return seed; } }; @@ -977,6 +979,7 @@ struct ggml_webgpu_mul_mat_vec_shader_decisions { uint32_t wg_size; uint32_t outputs_per_wg; uint32_t vec_size; + bool src_overlap = false; }; struct ggml_webgpu_quantize_q8_pipeline_key { @@ -998,10 +1001,11 @@ struct ggml_webgpu_mul_mat_pipeline_key { ggml_type src1_type; int vectorized; int use_subgroup_matrix; + bool src_overlap; bool operator==(const ggml_webgpu_mul_mat_pipeline_key & other) const { return src0_type == other.src0_type && src1_type == other.src1_type && vectorized == other.vectorized && - use_subgroup_matrix == other.use_subgroup_matrix; + use_subgroup_matrix == other.use_subgroup_matrix && src_overlap == other.src_overlap; } }; @@ -1012,6 +1016,7 @@ struct ggml_webgpu_mul_mat_pipeline_key_hash { ggml_webgpu_hash_combine(seed, key.src1_type); ggml_webgpu_hash_combine(seed, key.vectorized); ggml_webgpu_hash_combine(seed, key.use_subgroup_matrix); + ggml_webgpu_hash_combine(seed, key.src_overlap); return seed; } }; @@ -1034,6 +1039,7 @@ struct ggml_webgpu_mul_mat_shader_decisions { uint32_t subgroup_matrix_n; uint32_t mul_mat_wg_size; + bool src_overlap = false; }; /** MUL_MAT_ID **/ @@ -1950,7 +1956,7 @@ class ggml_webgpu_shader_lib { return quantize_q8_pipelines[key]; } - webgpu_pipeline get_mul_mat_vec_pipeline(const ggml_webgpu_shader_lib_context & context) { + webgpu_pipeline get_mul_mat_vec_pipeline(const ggml_webgpu_shader_lib_context & context, bool src_overlap) { ggml_webgpu_mul_mat_vec_pipeline_key key = {}; key.src0_type = context.src0->type; key.src1_type = context.src1->type; @@ -1961,6 +1967,7 @@ class ggml_webgpu_shader_lib { key.num_cols = context.dst->ne[1]; key.use_mmvq = ggml_webgpu_can_use_mmvq(context.src0, context.src1, context.supports_dot_product, context.vendor); + key.src_overlap = src_overlap; auto it = mul_mat_vec_pipelines.find(key); if (it != mul_mat_vec_pipelines.end()) { @@ -2068,6 +2075,11 @@ class ggml_webgpu_shader_lib { defines.push_back("Q8_1_T"); } + if (key.src_overlap) { + defines.push_back("SRC_OVERLAP"); + variant += "_src_overlap"; + } + defines.push_back(std::string("WG_SIZE=") + std::to_string(wg_size)); defines.push_back(std::string("OUTPUTS_PER_WG=") + std::to_string(outputs_per_wg)); defines.push_back(context.supports_subgroups ? "USE_SUBGROUP_REDUCTION" : "USE_WORKGROUP_REDUCTION"); @@ -2089,7 +2101,7 @@ class ggml_webgpu_shader_lib { return mul_mat_vec_pipelines[key]; } - webgpu_pipeline get_mul_mat_fast_pipeline(const ggml_webgpu_shader_lib_context & context) { + webgpu_pipeline get_mul_mat_fast_pipeline(const ggml_webgpu_shader_lib_context & context, bool src_overlap) { ggml_webgpu_mul_mat_pipeline_key key = {}; key.src0_type = context.src0->type; key.src1_type = context.src1->type; @@ -2098,6 +2110,7 @@ class ggml_webgpu_shader_lib { 1 : 0; key.use_subgroup_matrix = context.supports_subgroup_matrix; + key.src_overlap = src_overlap; auto it = mul_mat_fast_pipelines.find(key); if (it != mul_mat_fast_pipelines.end()) { @@ -2216,6 +2229,11 @@ class ggml_webgpu_shader_lib { variant += "_vectorized"; } + if (key.src_overlap) { + defines.push_back("SRC_OVERLAP"); + variant += "_src_overlap"; + } + if (!key.use_subgroup_matrix) { defines.push_back("WORKGROUP_SIZE_M=" + std::to_string(WEBGPU_MUL_MAT_WG_SIZE_M) + "u"); defines.push_back("WORKGROUP_SIZE_N=" + std::to_string(WEBGPU_MUL_MAT_WG_SIZE_N) + "u"); diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 394aeeda2..b79f71b68 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -1628,48 +1628,65 @@ static webgpu_encoded_op ggml_webgpu_mul_mat(webgpu_context & ctx, // Get or create pipeline webgpu_pipeline pipeline; std::vector dispatches; + const bool src_overlap = ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src0, src1) && !use_mmvq; if (use_mat_vec) { if (use_mmvq) { ggml_webgpu_quantize_q8_dispatch(ctx, src0, src1, dst, dispatches); } - pipeline = ctx->shader_lib->get_mul_mat_vec_pipeline(shader_lib_ctx); + pipeline = ctx->shader_lib->get_mul_mat_vec_pipeline(shader_lib_ctx, src_overlap); } else { - pipeline = ctx->shader_lib->get_mul_mat_fast_pipeline(shader_lib_ctx); + pipeline = ctx->shader_lib->get_mul_mat_fast_pipeline(shader_lib_ctx, src_overlap); + } + + uint32_t offset_src0 = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src0) / ggml_type_size(src0->type)); + uint32_t offset_src1 = (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src1) / ggml_type_size(src1->type)); + size_t merged_offset = 0; + size_t merged_size = 0; + if (src_overlap) { + const ggml_webgpu_merged_binding_range merged_range = + ggml_webgpu_tensor_merged_binding_range(ctx, { src0, src1 }); + merged_offset = merged_range.offset; + merged_size = merged_range.size; + offset_src0 = ggml_webgpu_tensor_merged_element_offset(src0, merged_range); + offset_src1 = ggml_webgpu_tensor_merged_element_offset(src1, merged_range); } // Build params - std::vector params = { - (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src0) / ggml_type_size(src0->type)), - (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src1) / ggml_type_size(src1->type)), - (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, dst) / ggml_type_size(dst->type)), - (uint32_t) dst->ne[0], - (uint32_t) dst->ne[1], - (uint32_t) src0->ne[0], - (uint32_t) (src0->nb[1] / ggml_type_size(src0->type)), - (uint32_t) (src1->nb[1] / ggml_type_size(src1->type)), - (uint32_t) (src0->nb[2] / ggml_type_size(src0->type)), - (uint32_t) (src1->nb[2] / ggml_type_size(src1->type)), - (uint32_t) (src0->nb[3] / ggml_type_size(src0->type)), - (uint32_t) (src1->nb[3] / ggml_type_size(src1->type)), - (uint32_t) src0->ne[2], - (uint32_t) src0->ne[3], - (uint32_t) (src1->ne[2] / src0->ne[2]), - (uint32_t) (src1->ne[3] / src0->ne[3]) - }; + std::vector params = { offset_src0, + offset_src1, + (uint32_t) (ggml_webgpu_tensor_misalignment(ctx, dst) / ggml_type_size(dst->type)), + (uint32_t) dst->ne[0], + (uint32_t) dst->ne[1], + (uint32_t) src0->ne[0], + (uint32_t) (src0->nb[1] / ggml_type_size(src0->type)), + (uint32_t) (src1->nb[1] / ggml_type_size(src1->type)), + (uint32_t) (src0->nb[2] / ggml_type_size(src0->type)), + (uint32_t) (src1->nb[2] / ggml_type_size(src1->type)), + (uint32_t) (src0->nb[3] / ggml_type_size(src0->type)), + (uint32_t) (src1->nb[3] / ggml_type_size(src1->type)), + (uint32_t) src0->ne[2], + (uint32_t) src0->ne[3], + (uint32_t) (src1->ne[2] / src0->ne[2]), + (uint32_t) (src1->ne[3] / src0->ne[3]) }; // Build bind group entries std::vector entries = {}; - - entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 0, src0)); if (use_mmvq) { + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 0, src0)); auto & mmvq_qq8_entry = dispatches[0].bind_group_entries[1]; entries.push_back(ggml_webgpu_make_bind_group_entry(1, ggml_webgpu_tensor_buf(dst), mmvq_qq8_entry.offset, mmvq_qq8_entry.size)); + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst)); + } else if (src_overlap) { + entries.push_back( + ggml_webgpu_make_bind_group_entry(0, ggml_webgpu_tensor_buf(src0), merged_offset, merged_size)); + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst)); } else { + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 0, src0)); entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, src1)); + entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst)); } - entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst)); // Calculate workgroup dimensions uint32_t wg_x = 1; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl index b0cf2853e..4a500e4ec 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl @@ -1,3 +1,7 @@ +#ifndef SRC0 +#define SRC0 src0 +#endif + #ifdef BYTE_HELPERS fn get_byte(value: u32, index: u32) -> u32 { return (value >> (index * 8)) & 0xFF; @@ -46,7 +50,7 @@ fn load_f16_as_f32_at_src(byte_offset: u32) -> f32 { #ifdef DECLARE_BYTE_LOADERS_SRC0 fn load_u16_at_src0(byte_offset: u32) -> u32 { - let word = src0[byte_offset / 4u]; + let word = SRC0[byte_offset / 4u]; let shift = (byte_offset & 0x2u) * 8u; return (word >> shift) & 0xFFFFu; } @@ -55,14 +59,14 @@ fn load_u16_at_src0(byte_offset: u32) -> u32 { // Caller extracts the 16-bit half it needs via & 0xFFFFu or >> 16u. // this is used in k-quants for better performance fn load_u32_at_src0_aligned(byte_offset: u32) -> u32 { - return src0[(byte_offset & ~3u) / 4u]; + return SRC0[(byte_offset & ~3u) / 4u]; } fn load_u32_at_src0(byte_offset: u32) -> u32 { let word_idx = byte_offset / 4u; let shift = (byte_offset & 0x3u) * 8u; - let lo = src0[word_idx]; - let hi = src0[word_idx + 1u]; + let lo = SRC0[word_idx]; + let hi = SRC0[word_idx + 1u]; let shifted = (lo >> shift) | (hi << (32u - shift)); return select(shifted, lo, shift == 0u); } @@ -73,7 +77,7 @@ fn load_f16_at_src0(byte_offset: u32) -> f16 { } fn load_f16_as_f32_at_src0(byte_offset: u32) -> f32 { - let word = src0[byte_offset / 4u]; + let word = SRC0[byte_offset / 4u]; let shift = (byte_offset & 0x2u) * 8u; let d_bits = (word >> shift) & 0xFFFFu; return unpack2x16float(d_bits)[0]; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl index 13996ab51..44b6bb710 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl @@ -1,3 +1,10 @@ +#ifndef SRC0 +#define SRC0 src0 +#endif +#ifndef SRC1 +#define SRC1 src1 +#endif + #ifdef VEC #define VEC_SIZE 4 #define SHMEM_TYPE vec4 @@ -39,7 +46,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3 let src0_idx = batch_offset + global_m * params.stride_01 + global_k; let src0_val = select( // taking a slight performance hit to avoid oob SRC0_TYPE(0.0), - src0[src0_idx/VEC_SIZE], + SRC0[src0_idx/VEC_SIZE], global_m < params.m && global_k < params.k); store_shmem(SHMEM_TYPE(src0_val), elem_idx); } @@ -57,7 +64,7 @@ fn init_shmem_src1(thread_id: u32, batch_offset: u32, offset_n: u32, k_outer: u3 let src1_idx = batch_offset + global_n * params.stride_11 + global_k; let src1_val = select( SRC1_TYPE(0.0), - src1[src1_idx/VEC_SIZE], + SRC1[src1_idx/VEC_SIZE], global_n < params.n && global_k < params.k); store_shmem(SHMEM_TYPE(src1_val), TILE_SRC0_SHMEM + elem_idx); } diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl index 98bbdeb83..0e17fae16 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl @@ -1,8 +1,12 @@ enable f16; #define DECLARE_BYTE_LOADERS_SRC0 -#include "common_decls.tmpl" +#ifdef SRC_OVERLAP +#define SRC0 merged_src +#define SRC1 merged_src +#endif +#include "common_decls.tmpl" #include "mul_mat_decls.tmpl" #ifdef VEC @@ -36,11 +40,17 @@ struct MulMatParams { broadcast3: u32 }; +#ifdef SRC_OVERLAP +@group(0) @binding(0) var merged_src: array; +#define DST_BINDING 1 +#else @group(0) @binding(0) var src0: array; // M rows, K columns @group(0) @binding(1) var src1: array; // K rows, N columns (transposed) -@group(0) @binding(2) var dst: array; // M rows, N columns (transposed) +#define DST_BINDING 2 +#endif -@group(0) @binding(3) var params: MulMatParams; +@group(0) @binding(DST_BINDING) var dst: array; // M rows, N columns (transposed) +@group(0) @binding(DST_BINDING + 1) var params: MulMatParams; fn get_local_n(thread_id: u32) -> u32 { return thread_id / WORKGROUP_SIZE_M; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl index d86a72ce6..35998a9b0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl @@ -4,6 +4,10 @@ enable subgroups; enable chromium_experimental_subgroup_matrix; #define DECLARE_BYTE_LOADERS_SRC0 +#ifdef SRC_OVERLAP +#define SRC0 merged_src +#define SRC1 merged_src +#endif #include "common_decls.tmpl" #include "mul_mat_decls.tmpl" @@ -48,11 +52,17 @@ struct MulMatParams { }; // SRC0_TYPE and SRC1_TYPE are defined in mul_mat_decls, which is included +#ifdef SRC_OVERLAP +@group(0) @binding(0) var merged_src: array; +#define DST_BINDING 1 +#else @group(0) @binding(0) var src0: array; // M rows, K columns @group(0) @binding(1) var src1: array; // K rows, N columns (transposed) -@group(0) @binding(2) var dst: array; // M rows, N columns (transposed) +#define DST_BINDING 2 +#endif -@group(0) @binding(3) var params: MulMatParams; +@group(0) @binding(DST_BINDING) var dst: array; // M rows, N columns (transposed) +@group(0) @binding(DST_BINDING + 1) var params: MulMatParams; const WG_M_SG_TILE_SIZE = SUBGROUP_M * SUBGROUP_MATRIX_M * SUBGROUP_MATRIX_M_SIZE; const WG_N_SG_TILE_SIZE = SUBGROUP_N * SUBGROUP_MATRIX_N * SUBGROUP_MATRIX_N_SIZE; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl index ebdf09513..1781a6c79 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl @@ -7,6 +7,11 @@ enable f16; requires packed_4x8_integer_dot_product; #endif +#ifdef SRC_OVERLAP +#define SRC0 merged_src +#define SRC1 merged_src +#endif + #define DECLARE_BYTE_LOADERS_SRC0 #include "common_decls.tmpl" @@ -35,17 +40,22 @@ struct MulMatParams { broadcast3: u32 }; +#if defined(MMVQ) @group(0) @binding(0) var src0: array; - -#ifdef MMVQ @group(0) @binding(1) var src1q: array; +#define DST_BINDING 2 +#elif defined(SRC_OVERLAP) +@group(0) @binding(0) var merged_src: array; +#define DST_BINDING 1 #else +@group(0) @binding(0) var src0: array; @group(0) @binding(1) var src1: array; +#define DST_BINDING 2 #endif -@group(0) @binding(2) var dst: array; +@group(0) @binding(DST_BINDING) var dst: array; // "mul_mat_vec_acc.tmpl" requires params.k, params.m, params.stride_01 -@group(0) @binding(3) var params: MulMatParams; +@group(0) @binding(DST_BINDING + 1) var params: MulMatParams; // Flattened as [row][thread] to keep each row's reduction contiguous in memory. var partial_sums: array; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl index 8fd0d1907..864b4bd2c 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl @@ -1,3 +1,10 @@ +#ifndef SRC0 +#define SRC0 src0 +#endif +#ifndef SRC1 +#define SRC1 src1 +#endif + #ifdef U32_DEQUANT_HELPERS #define SRC0_TYPE u32 @@ -43,13 +50,13 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src for (var k = thread_id; k < k_vec; k += WG_SIZE) { var x_vals: array; for (var col = 0u;col < NUM_COLS;col += 1) { - x_vals[col] = src1[src1_idx_base_vec + col * (params.stride_11 / VEC_SIZE) + k]; + x_vals[col] = SRC1[src1_idx_base_vec + col * (params.stride_11 / VEC_SIZE) + 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; - let w = src0[src0_idx]; + let w = SRC0[src0_idx]; for (var col = 0u;col < NUM_COLS;col += 1) { acc[col][row] += inner_dot(w, x_vals[col]); } @@ -76,7 +83,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -116,8 +123,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -160,8 +167,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -205,8 +212,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -253,8 +260,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -302,7 +309,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -347,7 +354,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -409,10 +416,10 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 4u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + 32u + i]); - x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 64u + i]); - x_block[col][i + 12u] = f32(src1[x_base + col * params.stride_11 + 96u + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4u] = f32(SRC1[x_base + col * params.stride_11 + 32u + i]); + x_block[col][i + 8u] = f32(SRC1[x_base + col * params.stride_11 + 64u + i]); + x_block[col][i + 12u] = f32(SRC1[x_base + col * params.stride_11 + 96u + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -518,8 +525,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 8u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 32u + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 8u] = f32(SRC1[x_base + col * params.stride_11 + 32u + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -610,10 +617,10 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src for (var col = 0u; col < NUM_COLS;col += 1) { let col_base = x_base + col * params.stride_11; for (var i = 0u; i < 4u; i++) { - x_block[col][i] = f32(src1[col_base + i]); - x_block[col][i + 4u] = f32(src1[col_base + 32u + i]); - x_block[col][i + 8u] = f32(src1[col_base + 128u + i]); - x_block[col][i + 12u] = f32(src1[col_base + 160u + i]); + x_block[col][i] = f32(SRC1[col_base + i]); + x_block[col][i + 4u] = f32(SRC1[col_base + 32u + i]); + x_block[col][i + 8u] = f32(SRC1[col_base + 128u + i]); + x_block[col][i + 12u] = f32(SRC1[col_base + 160u + i]); } } @@ -713,10 +720,10 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src for (var col = 0u; col < NUM_COLS;col += 1) { let col_base = x_base + col * params.stride_11; for (var i = 0u; i < 4u; i++) { - x_block[col][i] = f32(src1[col_base + i]); - x_block[col][i + 4u] = f32(src1[col_base + 32u + i]); - x_block[col][i + 8u] = f32(src1[col_base + 128u + i]); - x_block[col][i + 12u] = f32(src1[col_base + 160u + i]); + x_block[col][i] = f32(SRC1[col_base + i]); + x_block[col][i + 4u] = f32(SRC1[col_base + 32u + i]); + x_block[col][i + 8u] = f32(SRC1[col_base + 128u + i]); + x_block[col][i + 12u] = f32(SRC1[col_base + 160u + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -823,10 +830,10 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src for (var col = 0u; col < NUM_COLS;col += 1) { let col_base = x_base + col * params.stride_11; for (var l = 0u; l < 4u; l++) { - x_block[col][l] = f32(src1[col_base + l]); - x_block[col][l + 4u] = f32(src1[col_base + 32u + l]); - x_block[col][l + 8u] = f32(src1[col_base + 64u + l]); - x_block[col][l + 12u] = f32(src1[col_base + 96u + l]); + x_block[col][l] = f32(SRC1[col_base + l]); + x_block[col][l + 4u] = f32(SRC1[col_base + 32u + l]); + x_block[col][l + 8u] = f32(SRC1[col_base + 64u + l]); + x_block[col][l + 12u] = f32(SRC1[col_base + 96u + l]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -899,7 +906,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -960,7 +967,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1039,7 +1046,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1101,7 +1108,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1168,7 +1175,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1234,7 +1241,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1302,7 +1309,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1367,8 +1374,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + i + 16u]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4u] = f32(SRC1[x_base + col * params.stride_11 + i + 16u]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1418,7 +1425,7 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < 16u; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1476,8 +1483,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) { @@ -1521,8 +1528,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src var x_block: array, NUM_COLS>; for (var col = 0u; col < NUM_COLS;col += 1) { for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) { - x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]); - x_block[col][i + 8] = f32(src1[x_base + col * params.stride_11 + i + 8]); + x_block[col][i] = f32(SRC1[x_base + col * params.stride_11 + i]); + x_block[col][i + 8] = f32(SRC1[x_base + col * params.stride_11 + i + 8]); } } for (var row = 0u; row < OUTPUTS_PER_WG; row++) {