ggml-webgpu: add mulmat with overlapping src0/src1 (e.g., for minimax-01) (llama/27321)
This commit is contained in:
parent
fa1e2bcaaf
commit
7df5fa8e62
|
|
@ -954,10 +954,11 @@ struct ggml_webgpu_mul_mat_vec_pipeline_key {
|
||||||
int vectorized;
|
int vectorized;
|
||||||
uint32_t num_cols;
|
uint32_t num_cols;
|
||||||
bool use_mmvq;
|
bool use_mmvq;
|
||||||
|
bool src_overlap;
|
||||||
|
|
||||||
bool operator==(const ggml_webgpu_mul_mat_vec_pipeline_key & other) const {
|
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 &&
|
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.vectorized);
|
||||||
ggml_webgpu_hash_combine(seed, key.num_cols);
|
ggml_webgpu_hash_combine(seed, key.num_cols);
|
||||||
ggml_webgpu_hash_combine(seed, key.use_mmvq);
|
ggml_webgpu_hash_combine(seed, key.use_mmvq);
|
||||||
|
ggml_webgpu_hash_combine(seed, key.src_overlap);
|
||||||
return seed;
|
return seed;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
@ -977,6 +979,7 @@ struct ggml_webgpu_mul_mat_vec_shader_decisions {
|
||||||
uint32_t wg_size;
|
uint32_t wg_size;
|
||||||
uint32_t outputs_per_wg;
|
uint32_t outputs_per_wg;
|
||||||
uint32_t vec_size;
|
uint32_t vec_size;
|
||||||
|
bool src_overlap = false;
|
||||||
};
|
};
|
||||||
|
|
||||||
struct ggml_webgpu_quantize_q8_pipeline_key {
|
struct ggml_webgpu_quantize_q8_pipeline_key {
|
||||||
|
|
@ -998,10 +1001,11 @@ struct ggml_webgpu_mul_mat_pipeline_key {
|
||||||
ggml_type src1_type;
|
ggml_type src1_type;
|
||||||
int vectorized;
|
int vectorized;
|
||||||
int use_subgroup_matrix;
|
int use_subgroup_matrix;
|
||||||
|
bool src_overlap;
|
||||||
|
|
||||||
bool operator==(const ggml_webgpu_mul_mat_pipeline_key & other) const {
|
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 &&
|
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.src1_type);
|
||||||
ggml_webgpu_hash_combine(seed, key.vectorized);
|
ggml_webgpu_hash_combine(seed, key.vectorized);
|
||||||
ggml_webgpu_hash_combine(seed, key.use_subgroup_matrix);
|
ggml_webgpu_hash_combine(seed, key.use_subgroup_matrix);
|
||||||
|
ggml_webgpu_hash_combine(seed, key.src_overlap);
|
||||||
return seed;
|
return seed;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
@ -1034,6 +1039,7 @@ struct ggml_webgpu_mul_mat_shader_decisions {
|
||||||
uint32_t subgroup_matrix_n;
|
uint32_t subgroup_matrix_n;
|
||||||
|
|
||||||
uint32_t mul_mat_wg_size;
|
uint32_t mul_mat_wg_size;
|
||||||
|
bool src_overlap = false;
|
||||||
};
|
};
|
||||||
|
|
||||||
/** MUL_MAT_ID **/
|
/** MUL_MAT_ID **/
|
||||||
|
|
@ -1950,7 +1956,7 @@ class ggml_webgpu_shader_lib {
|
||||||
return quantize_q8_pipelines[key];
|
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 = {};
|
ggml_webgpu_mul_mat_vec_pipeline_key key = {};
|
||||||
key.src0_type = context.src0->type;
|
key.src0_type = context.src0->type;
|
||||||
key.src1_type = context.src1->type;
|
key.src1_type = context.src1->type;
|
||||||
|
|
@ -1961,6 +1967,7 @@ class ggml_webgpu_shader_lib {
|
||||||
key.num_cols = context.dst->ne[1];
|
key.num_cols = context.dst->ne[1];
|
||||||
key.use_mmvq =
|
key.use_mmvq =
|
||||||
ggml_webgpu_can_use_mmvq(context.src0, context.src1, context.supports_dot_product, context.vendor);
|
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);
|
auto it = mul_mat_vec_pipelines.find(key);
|
||||||
if (it != mul_mat_vec_pipelines.end()) {
|
if (it != mul_mat_vec_pipelines.end()) {
|
||||||
|
|
@ -2068,6 +2075,11 @@ class ggml_webgpu_shader_lib {
|
||||||
defines.push_back("Q8_1_T");
|
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("WG_SIZE=") + std::to_string(wg_size));
|
||||||
defines.push_back(std::string("OUTPUTS_PER_WG=") + std::to_string(outputs_per_wg));
|
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");
|
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];
|
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 = {};
|
ggml_webgpu_mul_mat_pipeline_key key = {};
|
||||||
key.src0_type = context.src0->type;
|
key.src0_type = context.src0->type;
|
||||||
key.src1_type = context.src1->type;
|
key.src1_type = context.src1->type;
|
||||||
|
|
@ -2098,6 +2110,7 @@ class ggml_webgpu_shader_lib {
|
||||||
1 :
|
1 :
|
||||||
0;
|
0;
|
||||||
key.use_subgroup_matrix = context.supports_subgroup_matrix;
|
key.use_subgroup_matrix = context.supports_subgroup_matrix;
|
||||||
|
key.src_overlap = src_overlap;
|
||||||
|
|
||||||
auto it = mul_mat_fast_pipelines.find(key);
|
auto it = mul_mat_fast_pipelines.find(key);
|
||||||
if (it != mul_mat_fast_pipelines.end()) {
|
if (it != mul_mat_fast_pipelines.end()) {
|
||||||
|
|
@ -2216,6 +2229,11 @@ class ggml_webgpu_shader_lib {
|
||||||
variant += "_vectorized";
|
variant += "_vectorized";
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (key.src_overlap) {
|
||||||
|
defines.push_back("SRC_OVERLAP");
|
||||||
|
variant += "_src_overlap";
|
||||||
|
}
|
||||||
|
|
||||||
if (!key.use_subgroup_matrix) {
|
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_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");
|
defines.push_back("WORKGROUP_SIZE_N=" + std::to_string(WEBGPU_MUL_MAT_WG_SIZE_N) + "u");
|
||||||
|
|
|
||||||
|
|
@ -1628,48 +1628,65 @@ static webgpu_encoded_op ggml_webgpu_mul_mat(webgpu_context & ctx,
|
||||||
// Get or create pipeline
|
// Get or create pipeline
|
||||||
webgpu_pipeline pipeline;
|
webgpu_pipeline pipeline;
|
||||||
std::vector<webgpu_dispatch_desc> dispatches;
|
std::vector<webgpu_dispatch_desc> dispatches;
|
||||||
|
const bool src_overlap = ggml_webgpu_tensor_binding_overlap(ctx->global_ctx, src0, src1) && !use_mmvq;
|
||||||
|
|
||||||
if (use_mat_vec) {
|
if (use_mat_vec) {
|
||||||
if (use_mmvq) {
|
if (use_mmvq) {
|
||||||
ggml_webgpu_quantize_q8_dispatch(ctx, src0, src1, dst, dispatches);
|
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 {
|
} 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
|
// Build params
|
||||||
std::vector<uint32_t> params = {
|
std::vector<uint32_t> params = { offset_src0,
|
||||||
(uint32_t) (ggml_webgpu_tensor_misalignment(ctx, src0) / ggml_type_size(src0->type)),
|
offset_src1,
|
||||||
(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) (ggml_webgpu_tensor_misalignment(ctx, dst) / ggml_type_size(dst->type)),
|
(uint32_t) dst->ne[0],
|
||||||
(uint32_t) dst->ne[0],
|
(uint32_t) dst->ne[1],
|
||||||
(uint32_t) dst->ne[1],
|
(uint32_t) src0->ne[0],
|
||||||
(uint32_t) src0->ne[0],
|
(uint32_t) (src0->nb[1] / ggml_type_size(src0->type)),
|
||||||
(uint32_t) (src0->nb[1] / ggml_type_size(src0->type)),
|
(uint32_t) (src1->nb[1] / ggml_type_size(src1->type)),
|
||||||
(uint32_t) (src1->nb[1] / ggml_type_size(src1->type)),
|
(uint32_t) (src0->nb[2] / ggml_type_size(src0->type)),
|
||||||
(uint32_t) (src0->nb[2] / ggml_type_size(src0->type)),
|
(uint32_t) (src1->nb[2] / ggml_type_size(src1->type)),
|
||||||
(uint32_t) (src1->nb[2] / ggml_type_size(src1->type)),
|
(uint32_t) (src0->nb[3] / ggml_type_size(src0->type)),
|
||||||
(uint32_t) (src0->nb[3] / ggml_type_size(src0->type)),
|
(uint32_t) (src1->nb[3] / ggml_type_size(src1->type)),
|
||||||
(uint32_t) (src1->nb[3] / ggml_type_size(src1->type)),
|
(uint32_t) src0->ne[2],
|
||||||
(uint32_t) src0->ne[2],
|
(uint32_t) src0->ne[3],
|
||||||
(uint32_t) src0->ne[3],
|
(uint32_t) (src1->ne[2] / src0->ne[2]),
|
||||||
(uint32_t) (src1->ne[2] / src0->ne[2]),
|
(uint32_t) (src1->ne[3] / src0->ne[3]) };
|
||||||
(uint32_t) (src1->ne[3] / src0->ne[3])
|
|
||||||
};
|
|
||||||
|
|
||||||
// Build bind group entries
|
// Build bind group entries
|
||||||
std::vector<wgpu::BindGroupEntry> entries = {};
|
std::vector<wgpu::BindGroupEntry> entries = {};
|
||||||
|
|
||||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 0, src0));
|
|
||||||
if (use_mmvq) {
|
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];
|
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,
|
entries.push_back(ggml_webgpu_make_bind_group_entry(1, ggml_webgpu_tensor_buf(dst), mmvq_qq8_entry.offset,
|
||||||
mmvq_qq8_entry.size));
|
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 {
|
} 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, 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
|
// Calculate workgroup dimensions
|
||||||
uint32_t wg_x = 1;
|
uint32_t wg_x = 1;
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,7 @@
|
||||||
|
#ifndef SRC0
|
||||||
|
#define SRC0 src0
|
||||||
|
#endif
|
||||||
|
|
||||||
#ifdef BYTE_HELPERS
|
#ifdef BYTE_HELPERS
|
||||||
fn get_byte(value: u32, index: u32) -> u32 {
|
fn get_byte(value: u32, index: u32) -> u32 {
|
||||||
return (value >> (index * 8)) & 0xFF;
|
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
|
#ifdef DECLARE_BYTE_LOADERS_SRC0
|
||||||
fn load_u16_at_src0(byte_offset: u32) -> u32 {
|
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;
|
let shift = (byte_offset & 0x2u) * 8u;
|
||||||
return (word >> shift) & 0xFFFFu;
|
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.
|
// Caller extracts the 16-bit half it needs via & 0xFFFFu or >> 16u.
|
||||||
// this is used in k-quants for better performance
|
// this is used in k-quants for better performance
|
||||||
fn load_u32_at_src0_aligned(byte_offset: u32) -> u32 {
|
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 {
|
fn load_u32_at_src0(byte_offset: u32) -> u32 {
|
||||||
let word_idx = byte_offset / 4u;
|
let word_idx = byte_offset / 4u;
|
||||||
let shift = (byte_offset & 0x3u) * 8u;
|
let shift = (byte_offset & 0x3u) * 8u;
|
||||||
let lo = src0[word_idx];
|
let lo = SRC0[word_idx];
|
||||||
let hi = src0[word_idx + 1u];
|
let hi = SRC0[word_idx + 1u];
|
||||||
let shifted = (lo >> shift) | (hi << (32u - shift));
|
let shifted = (lo >> shift) | (hi << (32u - shift));
|
||||||
return select(shifted, lo, shift == 0u);
|
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 {
|
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 shift = (byte_offset & 0x2u) * 8u;
|
||||||
let d_bits = (word >> shift) & 0xFFFFu;
|
let d_bits = (word >> shift) & 0xFFFFu;
|
||||||
return unpack2x16float(d_bits)[0];
|
return unpack2x16float(d_bits)[0];
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,10 @@
|
||||||
|
#ifndef SRC0
|
||||||
|
#define SRC0 src0
|
||||||
|
#endif
|
||||||
|
#ifndef SRC1
|
||||||
|
#define SRC1 src1
|
||||||
|
#endif
|
||||||
|
|
||||||
#ifdef VEC
|
#ifdef VEC
|
||||||
#define VEC_SIZE 4
|
#define VEC_SIZE 4
|
||||||
#define SHMEM_TYPE vec4<f16>
|
#define SHMEM_TYPE vec4<f16>
|
||||||
|
|
@ -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_idx = batch_offset + global_m * params.stride_01 + global_k;
|
||||||
let src0_val = select( // taking a slight performance hit to avoid oob
|
let src0_val = select( // taking a slight performance hit to avoid oob
|
||||||
SRC0_TYPE(0.0),
|
SRC0_TYPE(0.0),
|
||||||
src0[src0_idx/VEC_SIZE],
|
SRC0[src0_idx/VEC_SIZE],
|
||||||
global_m < params.m && global_k < params.k);
|
global_m < params.m && global_k < params.k);
|
||||||
store_shmem(SHMEM_TYPE(src0_val), elem_idx);
|
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_idx = batch_offset + global_n * params.stride_11 + global_k;
|
||||||
let src1_val = select(
|
let src1_val = select(
|
||||||
SRC1_TYPE(0.0),
|
SRC1_TYPE(0.0),
|
||||||
src1[src1_idx/VEC_SIZE],
|
SRC1[src1_idx/VEC_SIZE],
|
||||||
global_n < params.n && global_k < params.k);
|
global_n < params.n && global_k < params.k);
|
||||||
store_shmem(SHMEM_TYPE(src1_val), TILE_SRC0_SHMEM + elem_idx);
|
store_shmem(SHMEM_TYPE(src1_val), TILE_SRC0_SHMEM + elem_idx);
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,12 @@
|
||||||
enable f16;
|
enable f16;
|
||||||
|
|
||||||
#define DECLARE_BYTE_LOADERS_SRC0
|
#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"
|
#include "mul_mat_decls.tmpl"
|
||||||
|
|
||||||
#ifdef VEC
|
#ifdef VEC
|
||||||
|
|
@ -36,11 +40,17 @@ struct MulMatParams {
|
||||||
broadcast3: u32
|
broadcast3: u32
|
||||||
};
|
};
|
||||||
|
|
||||||
|
#ifdef SRC_OVERLAP
|
||||||
|
@group(0) @binding(0) var<storage, read_write> merged_src: array<SRC0_TYPE>;
|
||||||
|
#define DST_BINDING 1
|
||||||
|
#else
|
||||||
@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>; // M rows, K columns
|
@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(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)
|
#define DST_BINDING 2
|
||||||
|
#endif
|
||||||
|
|
||||||
@group(0) @binding(3) var<uniform> params: MulMatParams;
|
@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<DST_TYPE>; // M rows, N columns (transposed)
|
||||||
|
@group(0) @binding(DST_BINDING + 1) var<uniform> params: MulMatParams;
|
||||||
|
|
||||||
fn get_local_n(thread_id: u32) -> u32 {
|
fn get_local_n(thread_id: u32) -> u32 {
|
||||||
return thread_id / WORKGROUP_SIZE_M;
|
return thread_id / WORKGROUP_SIZE_M;
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,10 @@ enable subgroups;
|
||||||
enable chromium_experimental_subgroup_matrix;
|
enable chromium_experimental_subgroup_matrix;
|
||||||
|
|
||||||
#define DECLARE_BYTE_LOADERS_SRC0
|
#define DECLARE_BYTE_LOADERS_SRC0
|
||||||
|
#ifdef SRC_OVERLAP
|
||||||
|
#define SRC0 merged_src
|
||||||
|
#define SRC1 merged_src
|
||||||
|
#endif
|
||||||
#include "common_decls.tmpl"
|
#include "common_decls.tmpl"
|
||||||
|
|
||||||
#include "mul_mat_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
|
// SRC0_TYPE and SRC1_TYPE are defined in mul_mat_decls, which is included
|
||||||
|
#ifdef SRC_OVERLAP
|
||||||
|
@group(0) @binding(0) var<storage, read_write> merged_src: array<SRC0_TYPE>;
|
||||||
|
#define DST_BINDING 1
|
||||||
|
#else
|
||||||
@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>; // M rows, K columns
|
@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(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)
|
#define DST_BINDING 2
|
||||||
|
#endif
|
||||||
|
|
||||||
@group(0) @binding(3) var<uniform> params: MulMatParams;
|
@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<DST_TYPE>; // M rows, N columns (transposed)
|
||||||
|
@group(0) @binding(DST_BINDING + 1) var<uniform> params: MulMatParams;
|
||||||
|
|
||||||
const WG_M_SG_TILE_SIZE = SUBGROUP_M * SUBGROUP_MATRIX_M * SUBGROUP_MATRIX_M_SIZE;
|
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;
|
const WG_N_SG_TILE_SIZE = SUBGROUP_N * SUBGROUP_MATRIX_N * SUBGROUP_MATRIX_N_SIZE;
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,11 @@ enable f16;
|
||||||
requires packed_4x8_integer_dot_product;
|
requires packed_4x8_integer_dot_product;
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
|
#ifdef SRC_OVERLAP
|
||||||
|
#define SRC0 merged_src
|
||||||
|
#define SRC1 merged_src
|
||||||
|
#endif
|
||||||
|
|
||||||
#define DECLARE_BYTE_LOADERS_SRC0
|
#define DECLARE_BYTE_LOADERS_SRC0
|
||||||
#include "common_decls.tmpl"
|
#include "common_decls.tmpl"
|
||||||
|
|
||||||
|
|
@ -35,17 +40,22 @@ struct MulMatParams {
|
||||||
broadcast3: u32
|
broadcast3: u32
|
||||||
};
|
};
|
||||||
|
|
||||||
|
#if defined(MMVQ)
|
||||||
@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>;
|
@group(0) @binding(0) var<storage, read_write> src0: array<SRC0_TYPE>;
|
||||||
|
|
||||||
#ifdef MMVQ
|
|
||||||
@group(0) @binding(1) var<storage, read_write> src1q: array<q8_1>;
|
@group(0) @binding(1) var<storage, read_write> src1q: array<q8_1>;
|
||||||
|
#define DST_BINDING 2
|
||||||
|
#elif defined(SRC_OVERLAP)
|
||||||
|
@group(0) @binding(0) var<storage, read_write> merged_src: array<SRC0_TYPE>;
|
||||||
|
#define DST_BINDING 1
|
||||||
#else
|
#else
|
||||||
|
@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(1) var<storage, read_write> src1: array<SRC1_TYPE>;
|
||||||
|
#define DST_BINDING 2
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
@group(0) @binding(2) var<storage, read_write> dst: array<f32>;
|
@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<f32>;
|
||||||
// "mul_mat_vec_acc.tmpl" requires params.k, params.m, params.stride_01
|
// "mul_mat_vec_acc.tmpl" requires params.k, params.m, params.stride_01
|
||||||
@group(0) @binding(3) var<uniform> params: MulMatParams;
|
@group(0) @binding(DST_BINDING + 1) var<uniform> params: MulMatParams;
|
||||||
|
|
||||||
// Flattened as [row][thread] to keep each row's reduction contiguous in memory.
|
// Flattened as [row][thread] to keep each row's reduction contiguous in memory.
|
||||||
var<workgroup> partial_sums: array<f32, OUTPUTS_PER_WG * WG_SIZE>;
|
var<workgroup> partial_sums: array<f32, OUTPUTS_PER_WG * WG_SIZE>;
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,10 @@
|
||||||
|
#ifndef SRC0
|
||||||
|
#define SRC0 src0
|
||||||
|
#endif
|
||||||
|
#ifndef SRC1
|
||||||
|
#define SRC1 src1
|
||||||
|
#endif
|
||||||
|
|
||||||
#ifdef U32_DEQUANT_HELPERS
|
#ifdef U32_DEQUANT_HELPERS
|
||||||
#define SRC0_TYPE u32
|
#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) {
|
for (var k = thread_id; k < k_vec; k += WG_SIZE) {
|
||||||
var x_vals: array<SRC1_TYPE, NUM_COLS>;
|
var x_vals: array<SRC1_TYPE, NUM_COLS>;
|
||||||
for (var col = 0u;col < NUM_COLS;col += 1) {
|
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++) {
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
||||||
let output_row = row_base + row;
|
let output_row = row_base + row;
|
||||||
if (output_row < params.m) {
|
if (output_row < params.m) {
|
||||||
let src0_idx = (src0_batch_offset + output_row * params.stride_01) / VEC_SIZE + k;
|
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) {
|
for (var col = 0u;col < NUM_COLS;col += 1) {
|
||||||
acc[col][row] += inner_dot(w, x_vals[col]);
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
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++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
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++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 4u; i++) {
|
for (var i = 0u; i < 4u; 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]);
|
||||||
x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + 32u + 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 + 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 + 12u] = f32(SRC1[x_base + col * params.stride_11 + 96u + i]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 8u; i++) {
|
for (var i = 0u; i < 8u; 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]);
|
||||||
x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 32u + 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++) {
|
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) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
let col_base = x_base + col * params.stride_11;
|
let col_base = x_base + col * params.stride_11;
|
||||||
for (var i = 0u; i < 4u; i++) {
|
for (var i = 0u; i < 4u; i++) {
|
||||||
x_block[col][i] = f32(src1[col_base + 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 + 4u] = f32(SRC1[col_base + 32u + i]);
|
||||||
x_block[col][i + 8u] = f32(src1[col_base + 128u + 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 + 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) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
let col_base = x_base + col * params.stride_11;
|
let col_base = x_base + col * params.stride_11;
|
||||||
for (var i = 0u; i < 4u; i++) {
|
for (var i = 0u; i < 4u; i++) {
|
||||||
x_block[col][i] = f32(src1[col_base + 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 + 4u] = f32(SRC1[col_base + 32u + i]);
|
||||||
x_block[col][i + 8u] = f32(src1[col_base + 128u + 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 + 12u] = f32(SRC1[col_base + 160u + i]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
let col_base = x_base + col * params.stride_11;
|
let col_base = x_base + col * params.stride_11;
|
||||||
for (var l = 0u; l < 4u; l++) {
|
for (var l = 0u; l < 4u; l++) {
|
||||||
x_block[col][l] = f32(src1[col_base + 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 + 4u] = f32(SRC1[col_base + 32u + l]);
|
||||||
x_block[col][l + 8u] = f32(src1[col_base + 64u + 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 + 12u] = f32(SRC1[col_base + 96u + l]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) {
|
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] = 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 + 4u] = f32(SRC1[x_base + col * params.stride_11 + i + 16u]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, 16>, NUM_COLS>;
|
var x_block: array<array<f32, 16>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < 16u; i++) {
|
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++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 4] = f32(SRC1[x_base + col * params.stride_11 + i + 16]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
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<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
|
||||||
for (var col = 0u; col < NUM_COLS;col += 1) {
|
for (var col = 0u; col < NUM_COLS;col += 1) {
|
||||||
for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
|
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] = 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 + 8] = f32(SRC1[x_base + col * params.stride_11 + i + 8]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue