mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 10:57:33 -05:00
This commit is contained in:
@@ -1607,6 +1607,13 @@ class ggml_webgpu_shader_lib {
|
||||
defines.push_back("BLOCK_SIZE=1u");
|
||||
variant += "_i32";
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
defines.push_back("BF16");
|
||||
defines.push_back("SRC_TYPE=u32");
|
||||
defines.push_back("DST_TYPE=f32");
|
||||
defines.push_back("BLOCK_SIZE=1u");
|
||||
variant += "_bf16";
|
||||
break;
|
||||
default:
|
||||
{
|
||||
std::string type_upper = type_str;
|
||||
@@ -1992,13 +1999,21 @@ class ggml_webgpu_shader_lib {
|
||||
case GGML_TYPE_F32:
|
||||
defines.push_back("SRC0_INNER_TYPE=f32");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_F32");
|
||||
variant += "_f32";
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
defines.push_back("SRC0_INNER_TYPE=f16");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_F16");
|
||||
variant += "_f16";
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
defines.push_back("SRC0_INNER_TYPE=u32");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_BF16");
|
||||
variant += "_bf16";
|
||||
break;
|
||||
default:
|
||||
{
|
||||
// Quantized types: use helpers but accumulate in f16
|
||||
@@ -2149,7 +2164,7 @@ class ggml_webgpu_shader_lib {
|
||||
switch (context.src0->type) {
|
||||
case GGML_TYPE_F32:
|
||||
defines.push_back("SRC0_INNER_TYPE=f32");
|
||||
defines.push_back("FLOAT");
|
||||
defines.push_back("TYPE_F32");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
@@ -2157,12 +2172,20 @@ class ggml_webgpu_shader_lib {
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
defines.push_back("SRC0_INNER_TYPE=f16");
|
||||
defines.push_back("FLOAT");
|
||||
defines.push_back("TYPE_F16");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
variant += "_f16";
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
defines.push_back("SRC0_INNER_TYPE=u32");
|
||||
defines.push_back("TYPE_BF16");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
variant += "_bf16";
|
||||
break;
|
||||
default:
|
||||
{
|
||||
std::string type_upper = src0_name;
|
||||
@@ -2333,14 +2356,23 @@ class ggml_webgpu_shader_lib {
|
||||
defines.push_back("SRC0_INNER_TYPE=f32");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
defines.push_back("TYPE_F32");
|
||||
variant += "_f32";
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
defines.push_back("SRC0_INNER_TYPE=f16");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
defines.push_back("TYPE_F16");
|
||||
variant += "_f16";
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
defines.push_back("SRC0_INNER_TYPE=u32");
|
||||
defines.push_back("INIT_SRC0_SHMEM_FLOAT");
|
||||
defines.push_back("INIT_SRC1_SHMEM_FLOAT");
|
||||
defines.push_back("TYPE_BF16");
|
||||
variant += "_bf16";
|
||||
break;
|
||||
default:
|
||||
{
|
||||
std::string type_upper = src0_name;
|
||||
@@ -2453,13 +2485,21 @@ class ggml_webgpu_shader_lib {
|
||||
case GGML_TYPE_F32:
|
||||
defines.push_back("SRC0_INNER_TYPE=f32");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_F32");
|
||||
variant += "_f32";
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
defines.push_back("SRC0_INNER_TYPE=f16");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_F16");
|
||||
variant += "_f16";
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
defines.push_back("SRC0_INNER_TYPE=u32");
|
||||
defines.push_back("MUL_ACC_FLOAT");
|
||||
defines.push_back("TYPE_BF16");
|
||||
variant += "_bf16";
|
||||
break;
|
||||
default:
|
||||
{
|
||||
// Quantized types: use helpers but accumulate in f16
|
||||
|
||||
@@ -4432,7 +4432,8 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
src0->type == GGML_TYPE_F32 && (src1->type == GGML_TYPE_I64 || src1->type == GGML_TYPE_I32));
|
||||
break;
|
||||
case GGML_OP_GET_ROWS:
|
||||
if (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || ggml_webgpu_supported_qtype(src0->type)) {
|
||||
if (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16 ||
|
||||
ggml_webgpu_supported_qtype(src0->type)) {
|
||||
supports_op = (op->type == GGML_TYPE_F32);
|
||||
} else if (src0->type == GGML_TYPE_I32) {
|
||||
supports_op = op->type == GGML_TYPE_I32;
|
||||
@@ -4448,6 +4449,7 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
switch (src0->type) {
|
||||
case GGML_TYPE_F32:
|
||||
case GGML_TYPE_F16:
|
||||
case GGML_TYPE_BF16:
|
||||
case GGML_TYPE_Q1_0:
|
||||
case GGML_TYPE_Q4_0:
|
||||
case GGML_TYPE_Q4_1:
|
||||
@@ -4489,6 +4491,7 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
switch (src0->type) {
|
||||
case GGML_TYPE_F32:
|
||||
case GGML_TYPE_F16:
|
||||
case GGML_TYPE_BF16:
|
||||
case GGML_TYPE_Q1_0:
|
||||
case GGML_TYPE_Q4_0:
|
||||
case GGML_TYPE_Q4_1:
|
||||
|
||||
@@ -124,7 +124,12 @@ fn load_v_u32_at(byte_offset: u32) -> u32 {
|
||||
|
||||
#endif // U32_DEQUANT_HELPERS
|
||||
|
||||
|
||||
// bf16 helpers
|
||||
#if defined(TYPE_BF16) || defined(BF16)
|
||||
fn bf16_word_to_f32(word: u32, odd: u32) -> f32 {
|
||||
return bitcast<f32>(select(word << 16u, word & 0xFFFF0000u, odd == 1u));
|
||||
}
|
||||
#endif // TYPE_BF16 || BF16
|
||||
|
||||
#ifdef Q4_1_T
|
||||
struct q4_1 {
|
||||
|
||||
@@ -27,6 +27,12 @@ fn copy_elements(src_base: u32, dst_base: u32, offset: u32) {
|
||||
}
|
||||
#endif
|
||||
|
||||
#ifdef BF16
|
||||
fn copy_elements(src_base: u32, dst_base: u32, offset: u32) {
|
||||
dst[dst_base + offset] = bf16_word_to_f32(src[(src_base + offset) / 2u], (src_base + offset) & 1u);
|
||||
}
|
||||
#endif
|
||||
|
||||
#ifdef Q1_0
|
||||
fn copy_elements(src_base: u32, dst_base: u32, offset: u32) {
|
||||
let block_byte_base = (src_base + offset) * 18;
|
||||
|
||||
@@ -45,8 +45,14 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
|
||||
let global_k = k_outer + tile_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
|
||||
#if defined(TYPE_F16) || defined(TYPE_F32)
|
||||
SRC0_TYPE(0.0),
|
||||
SRC0[src0_idx/VEC_SIZE],
|
||||
#endif
|
||||
#ifdef TYPE_BF16
|
||||
f32(0.0),
|
||||
bf16_word_to_f32(SRC0[src0_idx / 2u], src0_idx & 1u),
|
||||
#endif
|
||||
global_m < params.m && global_k < params.k);
|
||||
store_shmem(SHMEM_TYPE(src0_val), elem_idx);
|
||||
}
|
||||
|
||||
@@ -26,17 +26,23 @@ fn sbyte_of(v: u32, b: u32) -> i32 {
|
||||
fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
|
||||
return f32(dot(SRC1_TYPE(src0_val), src1_val));
|
||||
}
|
||||
#endif
|
||||
#endif // VEC
|
||||
|
||||
#ifdef SCALAR
|
||||
#define VEC_SIZE 1u
|
||||
#define SRC0_TYPE SRC0_INNER_TYPE
|
||||
#define SRC1_TYPE SRC1_INNER_TYPE
|
||||
|
||||
#ifdef TYPE_BF16
|
||||
fn inner_dot(src0_val: f32, src1_val: SRC1_TYPE) -> f32 {
|
||||
return src0_val * f32(src1_val);
|
||||
}
|
||||
#else
|
||||
fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
|
||||
return f32(src0_val) * f32(src1_val);
|
||||
}
|
||||
#endif
|
||||
#endif // SCALAR
|
||||
|
||||
#ifdef MUL_ACC_FLOAT
|
||||
fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
|
||||
@@ -56,7 +62,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
|
||||
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;
|
||||
#if defined(TYPE_F16) || defined(TYPE_F32)
|
||||
let w = SRC0[src0_idx];
|
||||
#endif
|
||||
#ifdef TYPE_BF16
|
||||
let w = bf16_word_to_f32(SRC0[src0_idx / 2u], src0_idx & 1u);
|
||||
#endif
|
||||
for (var col = 0u;col < NUM_COLS;col += 1) {
|
||||
acc[col][row] += inner_dot(w, x_vals[col]);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user