mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-10-11 16:55:39 +02:00
webgpu: add f16 support to fill/set_rows (llama/29897)
This commit is contained in:
@@ -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";
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user