webgpu: add f16 support to fill/set_rows (llama/29897)

This commit is contained in:
Masashi Yoshimura
2026-10-06 10:38:19 +03:00
committed by Georgi Gerganov
parent f2aa80ad26
commit 51bfbd6170
4 changed files with 36 additions and 15 deletions
@@ -179,20 +179,22 @@ struct ggml_webgpu_argsort_shader_lib_context {
/** Set Rows **/
struct ggml_webgpu_set_rows_pipeline_key {
int src0_type;
int dst_type;
int vec4;
int i64_idx;
int pair_blocks;
bool operator==(const ggml_webgpu_set_rows_pipeline_key & other) const {
return dst_type == other.dst_type && vec4 == other.vec4 && i64_idx == other.i64_idx &&
pair_blocks == other.pair_blocks;
return src0_type == other.src0_type && dst_type == other.dst_type && vec4 == other.vec4 &&
i64_idx == other.i64_idx && pair_blocks == other.pair_blocks;
}
};
struct ggml_webgpu_set_rows_pipeline_key_hash {
size_t operator()(const ggml_webgpu_set_rows_pipeline_key & key) const {
size_t seed = 0;
ggml_webgpu_hash_combine(seed, key.src0_type);
ggml_webgpu_hash_combine(seed, key.dst_type);
ggml_webgpu_hash_combine(seed, key.vec4);
ggml_webgpu_hash_combine(seed, key.i64_idx);
@@ -1387,9 +1389,10 @@ class ggml_webgpu_shader_lib {
webgpu_pipeline get_set_rows_pipeline(const ggml_webgpu_shader_lib_context & context) {
const bool quantized = ggml_is_quantized(context.dst->type);
ggml_webgpu_set_rows_pipeline_key key = {};
key.src0_type = context.src0->type;
key.dst_type = context.dst->type;
key.vec4 =
(context.dst->type == GGML_TYPE_F32 || context.dst->type == GGML_TYPE_F16) && context.src0->ne[0] % 4 == 0;
key.vec4 = (context.dst->type == GGML_TYPE_F32 || context.dst->type == GGML_TYPE_F16) &&
context.src0->type == GGML_TYPE_F32 && context.src0->ne[0] % 4 == 0;
key.i64_idx = context.src1->type == GGML_TYPE_I64;
key.pair_blocks = quantized && ((context.src0->ne[0] / ggml_blck_size(context.dst->type)) % 2 == 0);
@@ -1422,6 +1425,11 @@ class ggml_webgpu_shader_lib {
GGML_ABORT("Unsupported dst type for set_rows shader");
}
if (context.src0->type == GGML_TYPE_F16) {
defines.push_back("TYPE_F16");
variant += "_src0_f16";
}
if (key.vec4) {
defines.push_back("VEC4");
variant += "_vec4";
+3 -2
View File
@@ -4431,7 +4431,8 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
case GGML_OP_SET_ROWS:
supports_op = ((op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_Q8_0 ||
op->type == GGML_TYPE_Q4_0) &&
src0->type == GGML_TYPE_F32 && (src1->type == GGML_TYPE_I64 || src1->type == GGML_TYPE_I32));
(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16) &&
(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 || src0->type == GGML_TYPE_BF16 ||
@@ -4701,7 +4702,7 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
supports_op = (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16) && (src0->type == op->type);
break;
case GGML_OP_FILL:
supports_op = op->type == GGML_TYPE_F32 && src0->type == GGML_TYPE_F32;
supports_op = (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16) && (src0->type == op->type);
break;
case GGML_OP_LOG:
supports_op = (op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16) && (src0->type == op->type);
@@ -11,7 +11,11 @@ enable f16;
#define DST_TYPE vec4<DST_INNER_TYPE>
#define VEC_SIZE 4
#else
#ifdef TYPE_F16
#define SRC_TYPE f16
#else
#define SRC_TYPE f32
#endif
#define DST_TYPE DST_INNER_TYPE
#define VEC_SIZE 1
#endif
@@ -1,3 +1,5 @@
enable f16;
#ifdef DST_Q8_0
#define BLOCK_SIZE 32u
#define BLOCK_BYTES 34u
@@ -8,8 +10,14 @@
#define QS_WORDS 4u
#endif
#ifdef TYPE_F16
#define SRC_TYPE f16
#else
#define SRC_TYPE f32
#endif
@group(0) @binding(0)
var<storage, read_write> src: array<f32>;
var<storage, read_write> src: array<SRC_TYPE>;
@group(0) @binding(1)
var<storage, read_write> idx: array<u32>;
@@ -112,7 +120,7 @@ fn quantize_block_params(src_block: u32) -> vec2<f32> {
#ifdef DST_Q8_0
var amax = 0.0;
for (var j: u32 = 0u; j < BLOCK_SIZE; j++) {
amax = max(amax, abs(src[src_block + j]));
amax = max(amax, abs(f32(src[src_block + j])));
}
let d = amax / 127.0;
@@ -122,7 +130,7 @@ fn quantize_block_params(src_block: u32) -> vec2<f32> {
var amax = 0.0;
var max_val = 0.0;
for (var j: u32 = 0u; j < BLOCK_SIZE; j++) {
let v = src[src_block + j];
let v = f32(src[src_block + j]);
let av = abs(v);
if (amax < av) {
amax = av;
@@ -139,15 +147,15 @@ fn quantize_block_params(src_block: u32) -> vec2<f32> {
fn quantize_block_word(src_block: u32, j: u32, id: f32) -> u32 {
#ifdef DST_Q8_0
let base = src_block + j * 4u;
return (u32(i32(round(src[base + 0u] * id)) & 0xFF) << 0u) |
(u32(i32(round(src[base + 1u] * id)) & 0xFF) << 8u) |
(u32(i32(round(src[base + 2u] * id)) & 0xFF) << 16u) |
(u32(i32(round(src[base + 3u] * id)) & 0xFF) << 24u);
return (u32(i32(round(f32(src[base + 0u]) * id)) & 0xFF) << 0u) |
(u32(i32(round(f32(src[base + 1u]) * id)) & 0xFF) << 8u) |
(u32(i32(round(f32(src[base + 2u]) * id)) & 0xFF) << 16u) |
(u32(i32(round(f32(src[base + 3u]) * id)) & 0xFF) << 24u);
#elif defined(DST_Q4_0)
var packed_q = 0u;
for (var k: u32 = 0u; k < 4u; k++) {
let x0 = src[src_block + j * 4u + k] * id;
let x1 = src[src_block + 16u + j * 4u + k] * id;
let x0 = f32(src[src_block + j * 4u + k]) * id;
let x1 = f32(src[src_block + 16u + j * 4u + k]) * id;
let q0 = u32(clamp(i32(x0 + 8.5), 0, 15));
let q1 = u32(clamp(i32(x1 + 8.5), 0, 15));
packed_q |= (q0 & 0xFu) << (8u * k);