diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 778a1b4bf7..53e89f94e1 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -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 diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index dd806ab99b..c5750ebbe9 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -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: diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl index 4a500e4ecd..9efd080b92 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl @@ -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(select(word << 16u, word & 0xFFFF0000u, odd == 1u)); +} +#endif // TYPE_BF16 || BF16 #ifdef Q4_1_T struct q4_1 { diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl index 487edb3275..a3114a4bea 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl @@ -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; 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 44b6bb710c..83f6fa1728 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl @@ -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); } 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 864b4bd2cd..841777df9c 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 @@ -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, 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]); }