auto push_type_defines = [&](const char * prefix, ggml_type type) {
std::string s_prefix = prefix;
if (type == GGML_TYPE_F32) {
- defines.push_back(s_prefix + "_F32");
+ defines.push_back(s_prefix + "=f32");
} else if (type == GGML_TYPE_F16) {
- defines.push_back(s_prefix + "_F16");
+ defines.push_back(s_prefix + "=f16");
} else {
GGML_ABORT("Unsupported type for CONV_2D shader");
}
};
- push_type_defines("WEIGHT", key.weight_type);
- push_type_defines("INPUT", key.input_type);
- push_type_defines("OUTPUT", key.output_type);
+ push_type_defines("WEIGHT_TYPE", key.weight_type);
+ push_type_defines("INPUT_TYPE", key.input_type);
+ push_type_defines("OUTPUT_TYPE", key.output_type);
defines.push_back(std::string("WG_SIZE=") + std::to_string(context.max_wg_size));
auto push_type_defines = [&](const char * prefix, ggml_type type) {
std::string s_prefix = prefix;
if (type == GGML_TYPE_F32) {
- defines.push_back(s_prefix + "_F32");
+ defines.push_back(s_prefix + "=f32");
} else if (type == GGML_TYPE_F16) {
- defines.push_back(s_prefix + "_F16");
+ defines.push_back(s_prefix + "=f16");
} else {
- GGML_ABORT("Unsupported type for CONV_2D_DW shader");
+ GGML_ABORT("Unsupported type for CONV_2D shader");
}
};
- push_type_defines("WEIGHT", key.weight_type);
- push_type_defines("INPUT", key.input_type);
- push_type_defines("OUTPUT", key.output_type);
+ push_type_defines("WEIGHT_TYPE", key.weight_type);
+ push_type_defines("INPUT_TYPE", key.input_type);
+ push_type_defines("OUTPUT_TYPE", key.output_type);
+
if (whcn) {
defines.push_back("WHCN");
}
auto push_type_defines = [&](const char * prefix, ggml_type type) {
std::string s_prefix = prefix;
if (type == GGML_TYPE_F32) {
- defines.push_back(s_prefix + "_F32");
+ defines.push_back(s_prefix + "=f32");
} else if (type == GGML_TYPE_F16) {
- defines.push_back(s_prefix + "_F16");
+ defines.push_back(s_prefix + "=f16");
} else {
GGML_ABORT("Unsupported type for IM2COL shader");
}
};
- push_type_defines("INPUT", key.input_type);
- push_type_defines("OUTPUT", key.output_type);
+ push_type_defines("INPUT_TYPE", key.input_type);
+ push_type_defines("OUTPUT_TYPE", key.output_type);
defines.push_back(std::string("WG_SIZE=") + std::to_string(context.max_wg_size));
(uint32_t) src1->ne[0],
(uint32_t) dst->ne[2],
- (uint32_t) dst->ne[3],
};
std::vector<wgpu::BindGroupEntry> entries = {
(uint32_t) ggml_nelements(dst),
(uint32_t) dst->ne[2],
- (uint32_t) dst->ne[3],
(uint32_t) dst->ne[0],
(uint32_t) dst->ne[1],
(uint32_t) src1->ne[0],
(uint32_t) src0->ne[2],
(uint32_t) src4->ne[1],
(uint32_t) src1->ne[2],
- (uint32_t) src1->ne[3],
(uint32_t) ggml_nelements(src1),
};
const ggml_tensor * K,
const ggml_tensor * V) {
const size_t storage_offset_alignment = global_ctx->capabilities.limits.minStorageBufferOffsetAlignment;
- const bool k_float_vec4_aligned = (K->type != GGML_TYPE_F16 && K->type != GGML_TYPE_F32) ||
- ggml_webgpu_flash_attn_float_vec4_aligned(K, storage_offset_alignment);
- const bool v_float_vec4_aligned = (V->type != GGML_TYPE_F16 && V->type != GGML_TYPE_F32) ||
- ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment);
- const bool k_vec_type_supported =
- K->type == GGML_TYPE_F32 || K->type == GGML_TYPE_F16 || K->type == GGML_TYPE_Q4_0 || K->type == GGML_TYPE_Q8_0;
- const bool v_vec_type_supported =
- V->type == GGML_TYPE_F32 || V->type == GGML_TYPE_F16 || V->type == GGML_TYPE_Q4_0 || V->type == GGML_TYPE_Q8_0;
- const uint32_t k_vec_head_align = (K->type == GGML_TYPE_F32 || K->type == GGML_TYPE_F16) ?
- GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH :
- (uint32_t) ggml_blck_size(K->type);
- const uint32_t v_vec_head_align = (V->type == GGML_TYPE_F32 || V->type == GGML_TYPE_F16) ?
- GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH :
- (uint32_t) ggml_blck_size(V->type);
- const bool kv_vec_head_dims_aligned = Q->ne[0] % k_vec_head_align == 0 && V->ne[0] % v_vec_head_align == 0;
+
+ const bool k_float_vec4_aligned = (K->type != GGML_TYPE_F16 && K->type != GGML_TYPE_F32) ||
+ ggml_webgpu_flash_attn_float_vec4_aligned(K, storage_offset_alignment);
+ const bool v_float_vec4_aligned = (V->type != GGML_TYPE_F16 && V->type != GGML_TYPE_F32) ||
+ ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment);
+
+ const uint32_t k_vec_head_align =
+ ggml_is_quantized(K->type) ? ggml_blck_size(K->type) : GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH;
+ const uint32_t v_vec_head_align =
+ ggml_is_quantized(V->type) ? ggml_blck_size(V->type) : GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH;
+ const bool kv_vec_head_dims_aligned = Q->ne[0] % k_vec_head_align == 0 && V->ne[0] % v_vec_head_align == 0;
return global_ctx->capabilities.supports_subgroups && (Q->ne[1] < GGML_WEBGPU_FLASH_ATTN_VEC_MAX_SEQ_LEN) &&
- kv_vec_head_dims_aligned && k_vec_type_supported && v_vec_type_supported && k_float_vec4_aligned &&
- v_float_vec4_aligned;
+ kv_vec_head_dims_aligned && k_float_vec4_aligned && v_float_vec4_aligned;
}
static ggml_webgpu_flash_attn_op ggml_webgpu_flash_attn_prepare(webgpu_context & ctx,
(uint32_t) dst->ne[0],
(uint32_t) dst->ne[1],
(uint32_t) dst->ne[2],
- (uint32_t) dst->ne[3],
dim,
(uint32_t) src0->ne[dim] };
(uint32_t) dst->ne[0],
(uint32_t) dst->ne[1],
(uint32_t) dst->ne[2],
- (uint32_t) dst->ne[3],
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(rn_dst, 0)) // epsilon, treated as f32 in the shader
};
(uint32_t) src->ne[0],
(uint32_t) src->ne[1],
(uint32_t) src->ne[2],
- (uint32_t) src->ne[3],
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 0)) // epsilon, treated as f32 in the shader
};
(uint32_t) (dst->nb[1] / ggml_type_size(dst->type)),
(uint32_t) (dst->nb[2] / ggml_type_size(dst->type)),
(uint32_t) (dst->nb[3] / ggml_type_size(dst->type)),
- (uint32_t) ggml_nelements(dst),
(uint32_t) src0->ne[0],
(uint32_t) src0->ne[1],
(uint32_t) src0->ne[2],
ne0: u32,
ne1: u32,
ne2: u32,
- ne3: u32,
dim: u32,
src0_nedim: u32
enable f16;
@group(0) @binding(0)
-#if defined(WEIGHT_F32)
-var<storage, read_write> weights: array<f32>;
-#elif defined(WEIGHT_F16)
-var<storage, read_write> weights: array<f16>;
-#endif
-
+var<storage, read_write> weights: array<WEIGHT_TYPE>;
@group(0) @binding(1)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
@group(0) @binding(2)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
struct Params {
offset_w: u32,
@group(0) @binding(3)
var<uniform> params: Params;
-fn load_weight(idx: u32) -> f32 {
- #if defined(WEIGHT_F32)
- return weights[idx];
- #elif defined(WEIGHT_F16)
- return f32(weights[idx]);
- #endif
-}
-
-fn load_input(idx: u32) -> f32 {
- #if defined(INPUT_F32)
- return input[idx];
- #elif defined(INPUT_F16)
- return f32(input[idx]);
- #endif
-}
-
-fn store_output(idx: u32, val: f32) {
- #if defined(OUTPUT_F32)
- output[idx] = val;
- #elif defined(OUTPUT_F16)
- output[idx] = f16(val);
- #endif
-}
-
fn ceil_div_u32(x: u32, y: u32) -> u32 {
return (x + y - 1) / y;
}
// entire receptive field is out of bounds
if (kw_begin >= kw_end || kh_begin >= kh_end) {
let out_idx = params.offset_o + ow * params.so0 + oh * params.so1 + oc * params.so2 + n * params.so3;
- store_output(out_idx, 0.0);
+ output[out_idx] = OUTPUT_TYPE(0.0);
return;
}
let iw = u32(ow_base + i32(kw * params.d0));
let w_idx = w_row_base + kw * params.sw0;
let in_idx = in_row_base + iw * params.si0;
- sum += load_weight(w_idx) * load_input(in_idx);
+ sum += f32(weights[w_idx]) * f32(input[in_idx]);
}
}
}
let out_idx = params.offset_o + ow * params.so0 + oh * params.so1 + oc * params.so2 + n * params.so3;
- store_output(out_idx, sum);
+ output[out_idx] = OUTPUT_TYPE(sum);
}
// weight (src0) is [KW,KH,1,C]; output matches the input layout.
@group(0) @binding(0)
-#if defined(WEIGHT_F32)
-var<storage, read_write> weights: array<f32>;
-#elif defined(WEIGHT_F16)
-var<storage, read_write> weights: array<f16>;
-#endif
-
+var<storage, read_write> weights: array<WEIGHT_TYPE>;
@group(0) @binding(1)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
@group(0) @binding(2)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
struct Params {
offset_w: u32,
ne: u32,
channels: u32,
- batches: u32,
dst_w: u32, dst_h: u32,
src_w: u32, src_h: u32,
knl_w: u32, knl_h: u32,
@group(0) @binding(3)
var<uniform> params: Params;
-fn load_weight(idx: u32) -> f32 {
- #if defined(WEIGHT_F32)
- return weights[idx];
- #elif defined(WEIGHT_F16)
- return f32(weights[idx]);
- #endif
-}
-fn load_input(idx: u32) -> f32 {
- #if defined(INPUT_F32)
- return input[idx];
- #elif defined(INPUT_F16)
- return f32(input[idx]);
- #endif
-}
-fn store_output(idx: u32, val: f32) {
- #if defined(OUTPUT_F32)
- output[idx] = val;
- #elif defined(OUTPUT_F16)
- output[idx] = f16(val);
- #endif
-}
-
#if defined(WHCN)
// Input/output/kernel contiguous in [W, H, C, N] order (kernel [KW,KH,C]).
fn conv_2d_dw(idx: u32) -> f32 {
for (var kx: u32 = 0u; kx < params.knl_w; kx += 1u) {
let src_x = i32(dst_x) * params.stride_x + i32(kx) * params.dilation_x - params.pad_x;
if (src_x < 0 || src_x >= i32(params.src_w)) { continue; }
- let v = load_input(src_i + u32(src_y) * params.src_w + u32(src_x));
- let k = load_weight(knl_i + ky * params.knl_w + kx);
+ let v = f32(input[src_i + u32(src_y) * params.src_w + u32(src_x)]);
+ let k = f32(weights[knl_i + ky * params.knl_w + kx]);
sum += v * k;
}
}
for (var kx: u32 = 0u; kx < params.knl_w; kx += 1u) {
let src_x = i32(dst_x) * params.stride_x + i32(kx) * params.dilation_x - params.pad_x;
if (src_x < 0 || src_x >= i32(params.src_w)) { continue; }
- let v = load_input(src_i + u32(src_y) * src_row + u32(src_x) * params.channels + c);
- let k = load_weight(params.offset_w + ky * knl_row + kx * params.channels + c);
+ let v = f32(input[src_i + u32(src_y) * src_row + u32(src_x) * params.channels + c]);
+ let k = f32(weights[params.offset_w + ky * knl_row + kx * params.channels + c]);
sum += v * k;
}
}
) {
let idx = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
if (idx >= params.ne) { return; }
- store_output(params.offset_o + idx, conv_2d_dw(idx));
+ output[params.offset_o + idx] = OUTPUT_TYPE(conv_2d_dw(idx));
}
#define BYTE_HELPERS
#include "common_decls.tmpl"
-#ifdef K_F32
-#define K_TYPE f32
-#elif defined(K_Q4_0) || defined(K_Q8_0)
-#define K_TYPE u32
-#else
-#define K_TYPE f16
-#endif
-
-#ifdef V_F32
-#define V_TYPE f32
-#elif defined(V_Q4_0) || defined(V_Q8_0)
-#define V_TYPE u32
-#else
-#define V_TYPE f16
-#endif
+#define FLASH_ATTN_SCALAR_KV
+#include "flash_attn_decls.tmpl"
// Default values
+// The actual values are defined in shader-lib.
#define HEAD_DIM_QK 64
#define HEAD_DIM_V 64
-
// The number of rows/columns/k in a subgroup matrix. MxK * KxN = MxN
// Note that the "K" here does not correspond to the K in attention's Q/K/V, it's just the common dimension.
#define SG_MAT_M 8
#define SG_MAT_N 8
#define SG_MAT_K 8
-
// Each workgroup processes one subgroup matrix of Q rows
#define Q_TILE SG_MAT_M
#define KV_TILE 16
// Number of subgroup-matrix-width blocks that span the KV tile. SG_MAT_N must divide KV_TILE.
#define KV_BLOCKS (KV_TILE / SG_MAT_N)
-struct Params {
- offset_q: u32,
- offset_k: u32,
- offset_v: u32,
- offset_mask: u32,
- offset_sinks: u32,
- offset_dst: u32,
-
- // shapes of Q/K/V
- n_heads: u32,
- seq_len_q: u32,
- seq_len_kv: u32,
-
- // strides (in elements)
- stride_q1: u32,
- stride_q2: u32,
- stride_q3: u32,
- stride_k1: u32,
- stride_k2: u32,
- stride_k3: u32,
- stride_v1: u32,
- stride_v2: u32,
- stride_v3: u32,
- stride_mask3: u32,
-
- // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA
- q_per_kv: u32,
-
- // softmax params
- scale: f32,
- max_bias: f32,
- logit_softcap: f32,
- n_head_log2: f32,
- m0: f32,
- m1: f32,
-};
-
-@group(0) @binding(0) var<storage, read_write> Q: array<f32>;
-#ifdef KV_OVERLAP
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#define V K
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#endif
-
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#elif defined(MASK)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#elif defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#else
-#ifdef KV_OVERLAP
-#define DST_BINDING 2
-#define PARAMS_BINDING 3
-#else
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#endif
-#endif
-
-@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<f32>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
-
-// Just a very small float value.
-const FLOAT_MIN: f32 = -1.0e9;
-
// The number of Q rows processed per workgroup
var<workgroup> q_shmem: array<f16, Q_TILE * HEAD_DIM_QK>;
#if !defined(K_DIRECT) || !defined(V_DIRECT)
+#define STAGING_SHMEM kv_shmem
+#define STAGING_OUT_TYPE f16
+#include "flash_attn_staging.tmpl"
const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V);
// we can reuse the same shmem for K and V since we only need one at a time
var<workgroup> kv_shmem: array<f16, kv_shmem_size>;
return v;
}
-fn load_f32x4(buf: ptr<storage, array<vec4<f32>>, read_write>, scalar_index: u32) -> vec4<f32> {
- return (*buf)[scalar_index >> 2u];
-}
-
-fn load_kx4(buf: ptr<storage, array<vec4<K_TYPE>>, read_write>, scalar_index: u32) -> vec4<K_TYPE> {
- return (*buf)[scalar_index >> 2u];
-}
-
-#if !defined(K_DIRECT) || !defined(V_DIRECT)
-#define QUANT_SHMEM kv_shmem
-#define QUANT_OUT_TYPE f16
-#include "flash_attn_quant_staging.tmpl"
-
-#if !defined(K_DIRECT) && !defined(K_Q4_0) && !defined(K_Q8_0)
-fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
- for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE) {
- let k_row = elem_idx / HEAD_DIM_QK;
- let k_col = elem_idx % HEAD_DIM_QK;
- let global_k_row = kv_tile + k_row;
- let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1;
- kv_shmem[elem_idx] = f16(select(
- 0.0,
- K[global_k_row_offset + k_col],
- global_k_row < params.seq_len_kv && k_col < HEAD_DIM_QK));
- }
-}
-#endif
-
-#if !defined(V_DIRECT) && !defined(V_Q4_0) && !defined(V_Q8_0)
-fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
- for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE) {
- let v_row = elem_idx / HEAD_DIM_V;
- let v_col = elem_idx % HEAD_DIM_V;
- let global_v_row = kv_tile + v_row;
- let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1;
- kv_shmem[elem_idx] = f16(select(
- 0.0,
- V[global_v_row_offset + v_col],
- global_v_row < params.seq_len_kv && v_col < HEAD_DIM_V));
- }
-}
-#endif
-#endif
-
@compute @workgroup_size(WG_SIZE)
fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
@builtin(local_invocation_id) local_id: vec3<u32>,
--- /dev/null
+#ifdef Q_F32
+#define Q_TYPE f32
+#else
+#define Q_TYPE f16
+#endif
+
+#ifdef K_F32
+#define K_TYPE f32
+#elif defined(K_Q4_0) || defined(K_Q8_0)
+#define K_TYPE u32
+#else
+#define K_TYPE f16
+#endif
+
+#ifdef V_F32
+#define V_TYPE f32
+#elif defined(V_Q4_0) || defined(V_Q8_0)
+#define V_TYPE u32
+#else
+#define V_TYPE f16
+#endif
+
+#ifdef DST_F32
+#define DST_TYPE f32
+#else
+#define DST_TYPE f16
+#endif
+
+#if defined(FLASH_ATTN_SCALAR_KV) || defined(K_Q4_0) || defined(K_Q8_0)
+#define K_STORAGE_TYPE K_TYPE
+#else
+#define K_STORAGE_TYPE vec4<K_TYPE>
+#endif
+
+#if defined(FLASH_ATTN_SCALAR_KV) || defined(V_Q4_0) || defined(V_Q8_0)
+#define V_STORAGE_TYPE V_TYPE
+#else
+#define V_STORAGE_TYPE vec4<V_TYPE>
+#endif
+
+// Just a very small float value.
+const FLOAT_MIN: f32 = -1.0e9;
+
+struct Params {
+ offset_q: u32,
+ offset_k: u32,
+ offset_v: u32,
+ offset_mask: u32,
+ offset_sinks: u32,
+ offset_dst: u32,
+
+ // shapes of Q/K/V
+ n_heads: u32,
+ seq_len_q: u32,
+ seq_len_kv: u32,
+
+ // strides (in elements)
+ stride_q1: u32,
+ stride_q2: u32,
+ stride_q3: u32,
+ stride_k1: u32,
+ stride_k2: u32,
+ stride_k3: u32,
+ stride_v1: u32,
+ stride_v2: u32,
+ stride_v3: u32,
+ stride_mask3: u32,
+
+ // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA
+ q_per_kv: u32,
+
+ // softmax params
+ scale: f32,
+ max_bias: f32,
+ logit_softcap: f32,
+ n_head_log2: f32,
+ m0: f32,
+ m1: f32,
+
+#ifdef FLASH_ATTN_VEC_SPLIT
+#ifdef BLK
+ blk_base: u32,
+ blk_nblk0: u32,
+ blk_nblk1: u32,
+#endif
+
+ tmp_data_base: u32,
+ tmp_stats_base: u32,
+ nwg: u32,
+#endif
+};
+
+@group(0) @binding(0) var<storage, read_write> Q: array<Q_TYPE>;
+@group(0) @binding(1) var<storage, read_write> K: array<K_STORAGE_TYPE>;
+#ifdef KV_OVERLAP
+#define V K
+#define MASK_BINDING 2
+#else
+@group(0) @binding(2) var<storage, read_write> V: array<V_STORAGE_TYPE>;
+#define MASK_BINDING 3
+#endif // KV_OVERLAP
+
+#ifdef MASK
+@group(0) @binding(MASK_BINDING) var<storage, read_write> mask: array<f16>;
+#define SINKS_BINDING (MASK_BINDING + 1)
+#else
+#define SINKS_BINDING MASK_BINDING
+#endif
+
+#ifdef SINKS
+@group(0) @binding(SINKS_BINDING) var<storage, read_write> sinks: array<f32>;
+#define BLK_BINDING (SINKS_BINDING + 1)
+#else
+#define BLK_BINDING SINKS_BINDING
+#endif
+
+#ifdef FLASH_ATTN_VEC_SPLIT
+#ifdef BLK
+@group(0) @binding(BLK_BINDING) var<storage, read_write> blk: array<u32>;
+#define TMP_BINDING (BLK_BINDING + 1)
+#else
+#define TMP_BINDING BLK_BINDING
+#endif
+
+@group(0) @binding(TMP_BINDING) var<storage, read_write> tmp: array<f32>;
+#define DST_BINDING (TMP_BINDING + 1)
+#else
+#define DST_BINDING BLK_BINDING
+#endif // FLASH_ATTN_VEC_SPLIT
+
+@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;
+
+#define PARAMS_BINDING (DST_BINDING + 1)
+@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
+++ /dev/null
-#include "quant_inner_loops.tmpl"
-
-#define BLOCK_SIZE 32
-#define BLOCKS_K ((HEAD_DIM_QK + BLOCK_SIZE - 1) / BLOCK_SIZE)
-#define BLOCKS_V ((HEAD_DIM_V + BLOCK_SIZE - 1) / BLOCK_SIZE)
-
-#if defined(K_Q4_0)
-#define K_NQ 16
-#define K_BLOCK_SIZE_BYTES 18u
-#define K_BYTES_PER_THREAD 8u
-#define K_BYTES_PER_INNER_LOOP 4u
-#elif defined(K_Q8_0)
-#define K_NQ 16
-#define K_BLOCK_SIZE_BYTES 34u
-#define K_BYTES_PER_THREAD 16u
-#define K_BYTES_PER_INNER_LOOP 4u
-#endif
-
-#if defined(V_Q4_0)
-#define V_NQ 16
-#define V_BLOCK_SIZE_BYTES 18u
-#define V_BYTES_PER_THREAD 8u
-#define V_BYTES_PER_INNER_LOOP 4u
-#elif defined(V_Q8_0)
-#define V_NQ 16
-#define V_BLOCK_SIZE_BYTES 34u
-#define V_BYTES_PER_THREAD 16u
-#define V_BYTES_PER_INNER_LOOP 4u
-#endif
-
-#if defined(K_Q4_0) || defined(K_Q8_0)
-fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
- for (var elem_idx = local_x * K_NQ; elem_idx < kv_count * HEAD_DIM_QK; elem_idx += WG_SIZE * K_NQ) {
- let blck_idx = elem_idx / BLOCK_SIZE;
- let block_offset = (elem_idx % BLOCK_SIZE) / K_NQ;
- let k_row = blck_idx / BLOCKS_K;
- let global_k_row = kv_tile + k_row;
- let block_k = blck_idx % BLOCKS_K;
- let row_offset = k_row * HEAD_DIM_QK;
- let global_block_idx = k_head_offset + global_k_row * params.stride_k1 + block_k;
- let block_byte_base = global_block_idx * K_BLOCK_SIZE_BYTES;
- let d = f16_from_u16(load_k_u16_at(block_byte_base));
- let thread_byte_offset = block_offset * K_BYTES_PER_THREAD;
- let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
- for (var j = 0u; j < K_BYTES_PER_THREAD / K_BYTES_PER_INNER_LOOP; j += 1u) {
- let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * K_BYTES_PER_INNER_LOOP;
- let q_packed = load_k_u32_at(q_byte_offset);
-#if defined(K_Q4_0)
- dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP);
-#elif defined(K_Q8_0)
- dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP);
-#endif
- }
- }
-}
-#endif
-
-#if defined(V_Q4_0) || defined(V_Q8_0)
-fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
- for (var elem_idx = local_x * V_NQ; elem_idx < kv_count * HEAD_DIM_V; elem_idx += WG_SIZE * V_NQ) {
- let blck_idx = elem_idx / BLOCK_SIZE;
- let block_offset = (elem_idx % BLOCK_SIZE) / V_NQ;
- let v_row = blck_idx / BLOCKS_V;
- let global_v_row = kv_tile + v_row;
- let block_k = blck_idx % BLOCKS_V;
- let row_offset = v_row * HEAD_DIM_V;
- let global_block_idx = v_head_offset + global_v_row * params.stride_v1 + block_k;
- let block_byte_base = global_block_idx * V_BLOCK_SIZE_BYTES;
- let d = f16_from_u16(load_v_u16_at(block_byte_base));
- let thread_byte_offset = block_offset * V_BYTES_PER_THREAD;
- let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
- for (var j = 0u; j < V_BYTES_PER_THREAD / V_BYTES_PER_INNER_LOOP; j += 1u) {
- let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * V_BYTES_PER_INNER_LOOP;
- let q_packed = load_v_u32_at(q_byte_offset);
-#if defined(V_Q4_0)
- dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP);
-#elif defined(V_Q8_0)
- dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP);
-#endif
- }
- }
-}
-#endif
--- /dev/null
+#if defined(K_Q4_0) || defined(K_Q8_0) || defined(V_Q4_0) || defined(V_Q8_0)
+#define QUANT_SHMEM STAGING_SHMEM
+#define QUANT_OUT_TYPE STAGING_OUT_TYPE
+#include "quant_inner_loops.tmpl"
+#undef QUANT_SHMEM
+#undef QUANT_OUT_TYPE
+#define BLOCK_SIZE 32
+#define BLOCKS_K ((HEAD_DIM_QK + BLOCK_SIZE - 1) / BLOCK_SIZE)
+#define BLOCKS_V ((HEAD_DIM_V + BLOCK_SIZE - 1) / BLOCK_SIZE)
+#endif
+
+#if defined(K_Q4_0)
+#define K_NQ 16
+#define K_BLOCK_SIZE_BYTES 18u
+#define K_BYTES_PER_THREAD 8u
+#define K_BYTES_PER_INNER_LOOP 4u
+#define DEQUANT_K_PACKED_TO_SHMEM dequant_q4_0_packed_to_shmem
+#elif defined(K_Q8_0)
+#define K_NQ 16
+#define K_BLOCK_SIZE_BYTES 34u
+#define K_BYTES_PER_THREAD 16u
+#define K_BYTES_PER_INNER_LOOP 4u
+#define DEQUANT_K_PACKED_TO_SHMEM dequant_q8_0_packed_to_shmem
+#endif
+
+#if defined(V_Q4_0)
+#define V_NQ 16
+#define V_BLOCK_SIZE_BYTES 18u
+#define V_BYTES_PER_THREAD 8u
+#define V_BYTES_PER_INNER_LOOP 4u
+#define DEQUANT_V_PACKED_TO_SHMEM dequant_q4_0_packed_to_shmem
+#elif defined(V_Q8_0)
+#define V_NQ 16
+#define V_BLOCK_SIZE_BYTES 34u
+#define V_BYTES_PER_THREAD 16u
+#define V_BYTES_PER_INNER_LOOP 4u
+#define DEQUANT_V_PACKED_TO_SHMEM dequant_q8_0_packed_to_shmem
+#endif
+
+#ifndef K_DIRECT
+fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
+#if defined(K_Q4_0) || defined(K_Q8_0)
+ for (var elem_idx = local_x * K_NQ; elem_idx < kv_count * HEAD_DIM_QK; elem_idx += WG_SIZE * K_NQ) {
+ let blck_idx = elem_idx / BLOCK_SIZE;
+ let block_offset = (elem_idx % BLOCK_SIZE) / K_NQ;
+ let k_row = blck_idx / BLOCKS_K;
+ let global_k_row = kv_tile + k_row;
+ let block_k = blck_idx % BLOCKS_K;
+ let row_offset = k_row * HEAD_DIM_QK;
+ let global_block_idx = k_head_offset + global_k_row * params.stride_k1 + block_k;
+ let block_byte_base = global_block_idx * K_BLOCK_SIZE_BYTES;
+ let d = f16_from_u16(load_k_u16_at(block_byte_base));
+ let thread_byte_offset = block_offset * K_BYTES_PER_THREAD;
+ let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
+ for (var j = 0u; j < K_BYTES_PER_THREAD / K_BYTES_PER_INNER_LOOP; j += 1u) {
+ let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * K_BYTES_PER_INNER_LOOP;
+ let q_packed = load_k_u32_at(q_byte_offset);
+ DEQUANT_K_PACKED_TO_SHMEM(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP);
+ }
+ }
+#elif defined(FLASH_ATTN_SCALAR_KV)
+ for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE) {
+ let k_row = elem_idx / HEAD_DIM_QK;
+ let k_col = elem_idx % HEAD_DIM_QK;
+ let global_k_row = kv_tile + k_row;
+ let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1;
+ STAGING_SHMEM[elem_idx] = STAGING_OUT_TYPE(select(
+ 0.0,
+ K[global_k_row_offset + k_col],
+ global_k_row < params.seq_len_kv && k_col < HEAD_DIM_QK));
+ }
+#else
+ for (var vec_idx_local = local_x; vec_idx_local < kv_count * Q_CHUNKS; vec_idx_local += WG_SIZE) {
+ let kv_local = vec_idx_local / Q_CHUNKS;
+ let chunk = vec_idx_local % Q_CHUNKS;
+ let global_k_row = kv_tile + kv_local;
+ let k_vec_index = (k_head_offset + global_k_row * params.stride_k1 + chunk * 4u) >> 2u;
+ let k4 = K[k_vec_index];
+ let kv_off = kv_local * HEAD_DIM_QK + chunk * 4u;
+ STAGING_SHMEM[kv_off + 0u] = STAGING_OUT_TYPE(k4.x);
+ STAGING_SHMEM[kv_off + 1u] = STAGING_OUT_TYPE(k4.y);
+ STAGING_SHMEM[kv_off + 2u] = STAGING_OUT_TYPE(k4.z);
+ STAGING_SHMEM[kv_off + 3u] = STAGING_OUT_TYPE(k4.w);
+ }
+#endif
+}
+#endif // !defined(K_DIRECT)
+
+#ifndef V_DIRECT
+fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
+#if defined(V_Q4_0) || defined(V_Q8_0)
+ for (var elem_idx = local_x * V_NQ; elem_idx < kv_count * HEAD_DIM_V; elem_idx += WG_SIZE * V_NQ) {
+ let blck_idx = elem_idx / BLOCK_SIZE;
+ let block_offset = (elem_idx % BLOCK_SIZE) / V_NQ;
+ let v_row = blck_idx / BLOCKS_V;
+ let global_v_row = kv_tile + v_row;
+ let block_k = blck_idx % BLOCKS_V;
+ let row_offset = v_row * HEAD_DIM_V;
+ let global_block_idx = v_head_offset + global_v_row * params.stride_v1 + block_k;
+ let block_byte_base = global_block_idx * V_BLOCK_SIZE_BYTES;
+ let d = f16_from_u16(load_v_u16_at(block_byte_base));
+ let thread_byte_offset = block_offset * V_BYTES_PER_THREAD;
+ let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset;
+ for (var j = 0u; j < V_BYTES_PER_THREAD / V_BYTES_PER_INNER_LOOP; j += 1u) {
+ let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * V_BYTES_PER_INNER_LOOP;
+ let q_packed = load_v_u32_at(q_byte_offset);
+ DEQUANT_V_PACKED_TO_SHMEM(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP);
+ }
+ }
+#elif defined(FLASH_ATTN_SCALAR_KV)
+ for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE) {
+ let v_row = elem_idx / HEAD_DIM_V;
+ let v_col = elem_idx % HEAD_DIM_V;
+ let global_v_row = kv_tile + v_row;
+ let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1;
+ STAGING_SHMEM[elem_idx] = STAGING_OUT_TYPE(select(
+ 0.0,
+ V[global_v_row_offset + v_col],
+ global_v_row < params.seq_len_kv && v_col < HEAD_DIM_V));
+ }
+#else
+ for (var vec_idx_local = local_x; vec_idx_local < kv_count * V_CHUNKS; vec_idx_local += WG_SIZE) {
+ let kv_local = vec_idx_local / V_CHUNKS;
+ let chunk = vec_idx_local % V_CHUNKS;
+ let global_v_row = kv_tile + kv_local;
+ let v_vec_index = (v_head_offset + global_v_row * params.stride_v1 + chunk * 4u) >> 2u;
+ let v4 = V[v_vec_index];
+ let kv_off = kv_local * HEAD_DIM_V + chunk * 4u;
+ STAGING_SHMEM[kv_off + 0u] = STAGING_OUT_TYPE(v4.x);
+ STAGING_SHMEM[kv_off + 1u] = STAGING_OUT_TYPE(v4.y);
+ STAGING_SHMEM[kv_off + 2u] = STAGING_OUT_TYPE(v4.z);
+ STAGING_SHMEM[kv_off + 3u] = STAGING_OUT_TYPE(v4.w);
+ }
+#endif
+}
+#endif // !defined(V_DIRECT)
#define BYTE_HELPERS
#include "common_decls.tmpl"
+#include "flash_attn_decls.tmpl"
-#ifdef Q_F16
-#define Q_TYPE f16
-#else
-#define Q_TYPE f32
-#endif
-
-#ifdef K_F32
-#define K_TYPE f32
-#elif defined(K_Q4_0) || defined(K_Q8_0)
-#define K_TYPE u32
-#else
-#define K_TYPE f16
-#endif
-
-#ifdef V_F32
-#define V_TYPE f32
-#elif defined(V_Q4_0) || defined(V_Q8_0)
-#define V_TYPE u32
-#else
-#define V_TYPE f16
-#endif
-
-#ifdef DST_F16
-#define DST_TYPE f16
-#else
-#define DST_TYPE f32
-#endif
-
+// Default values
+// The actual values are defined in shader-lib.
#define HEAD_DIM_QK 64
#define HEAD_DIM_V 64
#define Q_TILE 4
#define KV_TILE 64
#define WG_SIZE 128
-#ifndef MIN_SUBGROUP_SIZE
-#define MIN_SUBGROUP_SIZE MAX_SUBGROUP_SIZE
-#endif
-struct Params {
- offset_q: u32,
- offset_k: u32,
- offset_v: u32,
- offset_mask: u32,
- offset_sinks: u32,
- offset_dst: u32,
-
- n_heads: u32,
- seq_len_q: u32,
- seq_len_kv: u32,
-
- stride_q1: u32,
- stride_q2: u32,
- stride_q3: u32,
- stride_k1: u32,
- stride_k2: u32,
- stride_k3: u32,
- stride_v1: u32,
- stride_v2: u32,
- stride_v3: u32,
- stride_mask3: u32,
-
- q_per_kv: u32,
-
- scale: f32,
- max_bias: f32,
- logit_softcap: f32,
- n_head_log2: f32,
- m0: f32,
- m1: f32,
-};
-
-@group(0) @binding(0) var<storage, read_write> Q: array<Q_TYPE>;
-#ifdef KV_OVERLAP
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#define V K
-#else
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#if defined(V_Q4_0) || defined(V_Q8_0)
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#else
-@group(0) @binding(2) var<storage, read_write> V: array<vec4<V_TYPE>>;
-#endif
-#endif
-
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#elif defined(MASK)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#elif defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#else
-#ifdef KV_OVERLAP
-#define DST_BINDING 2
-#define PARAMS_BINDING 3
-#else
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#endif
-#endif
-
-@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
-
-const FLOAT_MIN: f32 = -1.0e9;
const Q_CHUNKS: u32 = HEAD_DIM_QK / 4u;
const V_CHUNKS: u32 = HEAD_DIM_V / 4u;
const SCORE_REGS_PER_LANE: u32 = (KV_TILE + MIN_SUBGROUP_SIZE - 1u) / MIN_SUBGROUP_SIZE;
const OUT_REGS_PER_LANE: u32 = (V_CHUNKS + MIN_SUBGROUP_SIZE - 1u) / MIN_SUBGROUP_SIZE;
-const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V);
-var<workgroup> q_shmem: array<Q_TYPE, Q_TILE * HEAD_DIM_QK>;
+#if !defined(K_DIRECT) || !defined(V_DIRECT)
+#define STAGING_SHMEM kv_shmem
+#define STAGING_OUT_TYPE f16
+#include "flash_attn_staging.tmpl"
+const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V);
var<workgroup> kv_shmem: array<f16, kv_shmem_size>;
-var<workgroup> p_shmem: array<f16, Q_TILE * KV_TILE>;
-
-#define QUANT_SHMEM kv_shmem
-#define QUANT_OUT_TYPE f16
-#include "flash_attn_quant_staging.tmpl"
-
-#if !defined(K_Q4_0) && !defined(K_Q8_0)
-fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
- for (var vec_idx_local = local_x; vec_idx_local < kv_count * Q_CHUNKS; vec_idx_local += WG_SIZE) {
- let kv_local = vec_idx_local / Q_CHUNKS;
- let chunk = vec_idx_local % Q_CHUNKS;
- let global_k_row = kv_tile + kv_local;
- let k_vec_index = (k_head_offset + global_k_row * params.stride_k1 + chunk * 4u) >> 2u;
- let k4 = K[k_vec_index];
- let kv_off = kv_local * HEAD_DIM_QK + chunk * 4u;
- kv_shmem[kv_off + 0u] = f16(k4.x);
- kv_shmem[kv_off + 1u] = f16(k4.y);
- kv_shmem[kv_off + 2u] = f16(k4.z);
- kv_shmem[kv_off + 3u] = f16(k4.w);
- }
-}
#endif
-#if !defined(V_Q4_0) && !defined(V_Q8_0)
-fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
- for (var vec_idx_local = local_x; vec_idx_local < kv_count * V_CHUNKS; vec_idx_local += WG_SIZE) {
- let kv_local = vec_idx_local / V_CHUNKS;
- let chunk = vec_idx_local % V_CHUNKS;
- let global_v_row = kv_tile + kv_local;
- let v_vec_index = (v_head_offset + global_v_row * params.stride_v1 + chunk * 4u) >> 2u;
- let v4 = V[v_vec_index];
- let kv_off = kv_local * HEAD_DIM_V + chunk * 4u;
- kv_shmem[kv_off + 0u] = f16(v4.x);
- kv_shmem[kv_off + 1u] = f16(v4.y);
- kv_shmem[kv_off + 2u] = f16(v4.z);
- kv_shmem[kv_off + 3u] = f16(v4.w);
- }
-}
-#endif
+var<workgroup> q_shmem: array<Q_TYPE, Q_TILE * HEAD_DIM_QK>;
+var<workgroup> p_shmem: array<f16, Q_TILE * KV_TILE>;
@compute @workgroup_size(WG_SIZE)
fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
#define BYTE_HELPERS
#include "common_decls.tmpl"
+#define FLASH_ATTN_VEC_SPLIT
+#include "flash_attn_decls.tmpl"
-#ifdef K_F32
-#define K_TYPE f32
-#elif defined(K_Q4_0) || defined(K_Q8_0)
-#define K_TYPE u32
-#else
-#define K_TYPE f16
-#endif
-
-#ifdef V_F32
-#define V_TYPE f32
-#elif defined(V_Q4_0) || defined(V_Q8_0)
-#define V_TYPE u32
-#else
-#define V_TYPE f16
-#endif
-
-#ifdef Q_F16
-#define Q_TYPE f16
-#else
-#define Q_TYPE f32
-#endif
-
-#ifdef DST_F16
-#define DST_TYPE f16
-#else
-#define DST_TYPE f32
-#endif
-
+// Default values
+// The actual values are defined in shader-lib.
#define HEAD_DIM_QK 64
#define HEAD_DIM_V 64
-
-#define KV_GRANULARITY 8
#define KV_TILE 16
#define WG_SIZE 64
-#define KV_BLOCKS (KV_TILE / KV_GRANULARITY)
-
-struct Params {
- offset_q: u32,
- offset_k: u32,
- offset_v: u32,
- offset_mask: u32,
- offset_sinks: u32,
- offset_dst: u32,
-
- // shapes of Q/K/V
- n_heads: u32,
- seq_len_q: u32,
- seq_len_kv: u32,
-
- // strides (in elements)
- stride_q1: u32,
- stride_q2: u32,
- stride_q3: u32,
- stride_k1: u32,
- stride_k2: u32,
- stride_k3: u32,
- stride_v1: u32,
- stride_v2: u32,
- stride_v3: u32,
- stride_mask3: u32,
-
- // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA
- q_per_kv: u32,
-
- // softmax params
- scale: f32,
- max_bias: f32,
- logit_softcap: f32,
- n_head_log2: f32,
- m0: f32,
- m1: f32,
-
-#ifdef BLK
- blk_base: u32,
- blk_nblk0: u32,
- blk_nblk1: u32,
-#endif
-
- tmp_data_base: u32,
- tmp_stats_base: u32,
- nwg: u32,
-};
-
-@group(0) @binding(0) var<storage, read_write> Q: array<Q_TYPE>;
-#ifdef KV_OVERLAP
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#define V K
-#else
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#if defined(V_Q4_0) || defined(V_Q8_0)
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#else
-@group(0) @binding(2) var<storage, read_write> V: array<vec4<V_TYPE>>;
-#endif
-#endif
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#ifdef BLK
-#define BLK_BINDING 4
-#define TMP_BINDING 5
-#define DST_BINDING 6
-#define PARAMS_BINDING 7
-#else
-#define TMP_BINDING 4
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#ifdef BLK
-#define BLK_BINDING 5
-#define TMP_BINDING 6
-#define DST_BINDING 7
-#define PARAMS_BINDING 8
-#else
-#define TMP_BINDING 5
-#define DST_BINDING 6
-#define PARAMS_BINDING 7
-#endif
-#endif
-#elif defined(MASK)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-#ifdef BLK
-#define BLK_BINDING 3
-#define TMP_BINDING 4
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#else
-#define TMP_BINDING 3
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-#ifdef BLK
-#define BLK_BINDING 4
-#define TMP_BINDING 5
-#define DST_BINDING 6
-#define PARAMS_BINDING 7
-#else
-#define TMP_BINDING 4
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#endif
-#elif defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> sinks: array<f32>;
-#define TMP_BINDING 3
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define TMP_BINDING 4
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#else
-#ifdef KV_OVERLAP
-#define TMP_BINDING 2
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-#define TMP_BINDING 3
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#endif
-
-#ifdef BLK
-@group(0) @binding(BLK_BINDING) var<storage, read_write> blk: array<u32>;
-#endif
-@group(0) @binding(TMP_BINDING) var<storage, read_write> tmp: array<f32>;
-@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
-
-// Just a very small float value.
-const FLOAT_MIN: f32 = -1.0e9;
+const Q_CHUNKS: u32 = HEAD_DIM_QK / 4u;
+const V_CHUNKS: u32 = HEAD_DIM_V / 4u;
const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V);
-var<workgroup> q_shmem: array<f32, HEAD_DIM_QK>;
-var<workgroup> o_shmem: array<f32, HEAD_DIM_V>;
-// note that we reuse the same storage for both since we only need one at a time
-var<workgroup> inter_shmem: array<f32, KV_TILE>;
-
-#ifdef MASK
-// storage for mask values
-var<workgroup> mask_shmem: array<f32, KV_TILE>;
-#endif
-
#if defined(K_DIRECT) || defined(V_DIRECT)
// Shared memory for scale factor (d) in quantized K/V. Multiple threads use the same value,
// so caching it is more efficient, even on the direct path.
// K/V shared memory handling
#if !defined(K_DIRECT) || !defined(V_DIRECT)
-
+#define STAGING_SHMEM kv_shmem
+#define STAGING_OUT_TYPE f32
+#include "flash_attn_staging.tmpl"
// we can reuse the same shmem for K and V since we only need one at a time
var<workgroup> kv_shmem: array<f32, kv_shmem_size>;
-
-#define QUANT_SHMEM kv_shmem
-#define QUANT_OUT_TYPE f32
-#include "flash_attn_quant_staging.tmpl"
-
-#if !defined(K_DIRECT) && !defined(K_Q4_0) && !defined(K_Q8_0)
-fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) {
- for (var elem_idx = local_x * 4u; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE * 4u) {
- let k_row = elem_idx / HEAD_DIM_QK;
- let k_col = elem_idx % HEAD_DIM_QK;
- let global_k_row = kv_tile + k_row;
- let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1;
- let in_bounds = global_k_row < params.seq_len_kv && (k_col + 3u) < HEAD_DIM_QK;
- let vec_idx = (global_k_row_offset + k_col) >> 2u;
- let k4 = select(vec4<K_TYPE>(0.0), K[vec_idx], in_bounds);
- kv_shmem[elem_idx + 0u] = f32(k4.x);
- kv_shmem[elem_idx + 1u] = f32(k4.y);
- kv_shmem[elem_idx + 2u] = f32(k4.z);
- kv_shmem[elem_idx + 3u] = f32(k4.w);
- }
-}
#endif
-#if !defined(V_DIRECT) && !defined(V_Q4_0) && !defined(V_Q8_0)
-fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) {
- for (var elem_idx = local_x * 4u; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE * 4u) {
- let v_row = elem_idx / HEAD_DIM_V;
- let v_col = elem_idx % HEAD_DIM_V;
- let global_v_row = kv_tile + v_row;
- let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1;
- let in_bounds = global_v_row < params.seq_len_kv && (v_col + 3u) < HEAD_DIM_V;
- let vec_idx = (global_v_row_offset + v_col) >> 2u;
- let v4 = select(vec4<V_TYPE>(0.0), V[vec_idx], in_bounds);
- kv_shmem[elem_idx + 0u] = f32(v4.x);
- kv_shmem[elem_idx + 1u] = f32(v4.y);
- kv_shmem[elem_idx + 2u] = f32(v4.z);
- kv_shmem[elem_idx + 3u] = f32(v4.w);
- }
-}
+var<workgroup> q_shmem: array<f32, HEAD_DIM_QK>;
+var<workgroup> o_shmem: array<f32, HEAD_DIM_V>;
+// note that we reuse the same storage for both since we only need one at a time
+var<workgroup> inter_shmem: array<f32, KV_TILE>;
+
+#ifdef MASK
+// storage for mask values
+var<workgroup> mask_shmem: array<f32, KV_TILE>;
#endif
-#endif // !defined(K_DIRECT) || !defined(V_DIRECT)
// Storage for row max and exp sum during online softmax
fn calc_softmax_term(kv_idx: u32, slope: f32, has_bias: bool, apply_mask: bool) -> f32 {
-#include "common_decls.tmpl"
enable f16;
@group(0) @binding(0)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
@group(0) @binding(1)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
struct Params {
offset_i: u32,
@group(0) @binding(2)
var<uniform> params: Params;
-fn load_input(idx: u32) -> f32 {
- #if defined(INPUT_F32)
- return input[idx];
- #elif defined(INPUT_F16)
- return f32(input[idx]);
- #endif
-}
-
-fn store_output(idx: u32, val: f32) {
- #if defined(OUTPUT_F32)
- output[idx] = val;
- #elif defined(OUTPUT_F16)
- output[idx] = f16(val);
- #endif
-}
-
@compute @workgroup_size(WG_SIZE)
fn main(
@builtin(global_invocation_id) gid: vec3<u32>,
let iw_i32 = i32(ow * params.s0 + kw * params.d0) - i32(params.p0);
let ih_i32 = i32(oh * params.s1 + kh * params.d1) - i32(params.p1);
+ let output_idx = params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3;
+
if (iw_i32 >= 0 && iw_i32 < i32(params.IW) && ih_i32 >= 0 && ih_i32 < i32(params.IH)) {
let iw = u32(iw_i32);
let ih = u32(ih_i32);
let in_idx = params.offset_i + iw * params.si0 + ih * params.si1 + ic * params.si2 + n * params.si3;
- store_output(params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3, load_input(in_idx));
+ output[output_idx] = OUTPUT_TYPE(input[in_idx]);
} else {
- store_output(params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3, 0.0);
+ output[output_idx] = OUTPUT_TYPE(0.0);
}
}
ne0: u32,
ne1: u32,
ne2: u32,
- ne3: u32,
eps: f32
};
ne0: u32,
ne1: u32,
ne2: u32,
- ne3: u32,
eps: f32
};
stride_dst3: u32,
// shape of src0/dst
- ne: u32,
ne0: u32,
ne1: u32,
ne2: u32,
m1: f32,
};
-@group(0) @binding(0)
+#define SRC_BINDING 0
+@group(0) @binding(SRC_BINDING)
var<storage, read_write> src: array<f32>;
#ifdef HAS_MASK
-#ifdef HAS_SINK
-@group(0) @binding(1)
+#define MASK_BINDING SRC_BINDING + 1
+@group(0) @binding(MASK_BINDING)
var<storage, read_write> mask: array<MaskType>;
-@group(0) @binding(2)
-var<storage, read_write> sinks: array<f32>;
-
-#ifdef INPLACE
-@group(0) @binding(3)
-var<uniform> params: Params;
-
#else
-@group(0) @binding(3)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(4)
-var<uniform> params: Params;
+#define MASK_BINDING SRC_BINDING
#endif
+#ifdef HAS_SINK
+#define SINKS_BINDING MASK_BINDING + 1
+@group(0) @binding(SINKS_BINDING)
+var<storage, read_write> sinks: array<f32>;
#else
-@group(0) @binding(1)
-var<storage, read_write> mask: array<MaskType>;
-
-#ifdef INPLACE
-@group(0) @binding(2)
-var<uniform> params: Params;
-
-#else
-@group(0) @binding(2)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(3)
-var<uniform> params: Params;
-#endif
+#define SINKS_BINDING MASK_BINDING
#endif
-#else
-#ifdef HAS_SINK
-@group(0) @binding(1)
-var<storage, read_write> sinks: array<f32>;
+#define DST_BINDING SINKS_BINDING + 1
+@group(0) @binding(DST_BINDING)
+var<storage, read_write> dst: array<f32>;
#ifdef INPLACE
-@group(0) @binding(2)
-var<uniform> params: Params;
-
+#define PARAMS_BINDING DST_BINDING
#else
-@group(0) @binding(2)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(3)
-var<uniform> params: Params;
+#define PARAMS_BINDING (DST_BINDING + 1)
#endif
-#else
-#ifdef INPLACE
-@group(0) @binding(1)
-var<uniform> params: Params;
-#else
-@group(0) @binding(1)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(2)
+@group(0) @binding(PARAMS_BINDING)
var<uniform> params: Params;
-#endif
-#endif
-#endif
#ifdef INPLACE
fn inter_value(i: u32) -> f32 {
col += WG_SIZE;
}
}
-
k: u32,
ne2: u32,
- ne3: u32,
};
@group(0) @binding(3)
n_head: u32,
n_group: u32,
n_seq_tokens: u32,
- n_seqs: u32,
y_elems: u32,
};