Commit 4b1622afb for llama.cpp
commit 4b1622afb7dccc819df968d1ef9bb5ebeeac5538
Author: Masashi Yoshimura <yoshimura.masashi.frbs@gmail.com>
Date: Thu Oct 1 22:09:07 2026 +0900
webgpu: add bfloat16 support for MUL_MAT/MUL_MAT_ID/GET_ROWS- #29358 (#29358)
diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
index 778a1b4bf..53e89f94e 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 dd806ab99..c5750ebbe 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 4a500e4ec..9efd080b9 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<f32>(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 487edb327..a3114a4be 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 44b6bb710..83f6fa172 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 864b4bd2c..841777df9 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<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]);
}