Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions backends/webgpu/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@ set(WEBGPU_SRCS
runtime/WebGPUExecutionOptions.cpp
runtime/WebGPUGraph.cpp
runtime/passes/SwiGLU.cpp
runtime/passes/QkvBk64.cpp
runtime/WebGPUDelegateHeader.cpp
runtime/WebGPUDevice.cpp
runtime/WebGPUQueryPool.cpp
Expand Down
560 changes: 49 additions & 511 deletions backends/webgpu/runtime/WebGPUGraph.cpp

Large diffs are not rendered by default.

34 changes: 6 additions & 28 deletions backends/webgpu/runtime/WebGPUGraph.h
Original file line number Diff line number Diff line change
Expand Up @@ -523,6 +523,12 @@ class WebGPUGraph {
return tensor_mem_obj_ids_[id];
}

// True if id is a prepack-routed constant with a recorded source (inline
// offset or named-data-map key); fusion passes require direct constants.
bool has_constant_source(int id) const {
return constant_sources_.count(id) != 0;
}

public:
// True when the sdpa K/V cache is stored f16-packed (runtime opt-in).
bool kv_f16() const {
Expand Down Expand Up @@ -660,34 +666,6 @@ class WebGPUGraph {
std::unordered_map<std::string, WGPUBindGroupLayout> bgl_cache_;

size_t uniform_buffer_bytes_ = 0;

// QKV-concat fusion: one detected attention q/k/v linear
// triple sharing an input activation (value ids + shapes), fused in build()
// into a single multi-output q4gsw GEMM that scatter-writes q/k/v. Only used
// during build(); inert (never populated) when no q/k/v triple matches.
struct QkvFusionGroup {
int input_id = -1;
int out_q = -1, out_k = -1, out_v = -1;
int weight_q = -1, weight_k = -1, weight_v = -1;
int scales_q = -1, scales_k = -1, scales_v = -1;
uint32_t Nq = 0, Nk = 0, Nv = 0; // 2048, 512, 512
uint32_t K = 0, K_packed = 0, group_size = 0, num_groups = 0;
uint32_t padded_N_q = 0, padded_N_k = 0, padded_N_v = 0;
unsigned op_idx[3] = {0, 0, 0}; // the 3 q/k/v linear op-chain indices
utils::DispatchRange sep_dispatch[3] = {
{0, 0},
{0, 0},
{0, 0}}; // each linear's complete route range (filled in build())
size_t fused_dispatch = 0; // the fused GEMM dispatch index
WGPUBuffer fused_params =
nullptr; // the fused params UBO (rewritten by the hook)
};
// Concat the 3 packed weights (row-stack) + scales (strided gather) into
// fused buffers, then record ONE fused-GEMM dispatch (bespoke 8-binding
// layout) that writes the 3 original q/k/v output buffers, plus a 3-output
// resize hook.
void add_qkv_fused_dispatch(QkvFusionGroup& g);
void add_qkv_fused_hook(const QkvFusionGroup& g);
};

} // namespace executorch::backends::webgpu
16 changes: 8 additions & 8 deletions backends/webgpu/runtime/WebGPUShaderRegistry.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -91,13 +91,13 @@
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/linear_dW_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_backward_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_coop4_bicol_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_qkv_fused_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_shmem_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_f16acc_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_pwdq_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_half_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_gemm_steel_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_linear_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_qkv_bk64_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_requant_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/quantized_linear/q4gsw_steel_bk64_wgsl.h>
#include <executorch/backends/webgpu/runtime/ops/reduce/reduce_wgsl.h>
Expand Down Expand Up @@ -709,13 +709,6 @@ constexpr std::array<WebGPUShaderInfo, 130> kShaderRegistry = {{
kQ4gswLinearCoop4BicolWorkgroupSizeY,
kQ4gswLinearCoop4BicolWorkgroupSizeZ,
},
{
"q4gsw_linear_gemm_qkv_fused",
kQ4gswLinearGemmQkvFusedWGSL,
kQ4gswLinearGemmQkvFusedWorkgroupSizeX,
kQ4gswLinearGemmQkvFusedWorkgroupSizeY,
kQ4gswLinearGemmQkvFusedWorkgroupSizeZ,
},
{
"q4gsw_linear_gemm_shmem",
kQ4gswLinearGemmShmemWGSL,
Expand Down Expand Up @@ -751,6 +744,13 @@ constexpr std::array<WebGPUShaderInfo, 130> kShaderRegistry = {{
kQ4gswLinearGemmSteelHalfPwdqF16accWorkgroupSizeY,
kQ4gswLinearGemmSteelHalfPwdqF16accWorkgroupSizeZ,
},
{
"q4gsw_qkv_bk64",
kQ4gswQkvBk64WGSL,
kQ4gswQkvBk64WorkgroupSizeX,
kQ4gswQkvBk64WorkgroupSizeY,
kQ4gswQkvBk64WorkgroupSizeZ,
},
{
"q4gsw_requant",
kQ4gswRequantWGSL,
Expand Down
Original file line number Diff line number Diff line change
@@ -1,12 +1,5 @@
enable f16;
// Fused QKV q4gsw GEMM (Llama attention projections): one [M, N=3072] pwdq + f16-accumulate GEMM
// (vec4<f32> activation load) that scatter-writes each output column range to a SEPARATE buffer --
// c<2048 -> q, [2048,2560) -> k, [2560,3072) -> v. Replaces the 3 separate q/k/v linear dispatches;
// fixes the N=512 K/V occupancy starvation (16 WGs -> 96 WGs at M~128). Boundaries are 64-tile-aligned
// so each 64-col tile maps to exactly one output (uniform branch per workgroup). Per-output ROW STRIDE:
// q=2048, k=v=512. BIT-EXACT to 3 separate pwdqf16acc linears (fusing along N does not change the
// per-column K-accumulation order). Validated on Canary M4 Pro: correct (maxRel ~1e-3), scatter overhead
// 1.02x (free), concat win 1.63x on the QKV block. Boundaries hardcoded for Llama-3.2-1B GQA (32Q/8KV).

@group(0) @binding(0) var<storage, read_write> t_out_q: array<f32>;
@group(0) @binding(1) var<storage, read_write> t_out_k: array<f32>;
@group(0) @binding(2) var<storage, read_write> t_out_v: array<f32>;
Expand All @@ -25,10 +18,12 @@ struct Params {
_pad: u32,
}
@group(0) @binding(7) var<uniform> params: Params;
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u;

// BK64 QKV variant: group_size=64 keeps one scale valid for all eight packed words.
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 64u;
const N_Q: u32 = 2048u; const N_QK: u32 = 2560u; const N_KV: u32 = 512u;
var<workgroup> As: array<f16, 1024>;
var<workgroup> Bs: array<f16, 1024>;
var<workgroup> As: array<f16, 4096>;
var<workgroup> Bs: array<f16, 4096>;
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid: vec3<u32>) {
Expand All @@ -44,18 +39,31 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
}
let ar = tid / 4u;
let ac = (tid % 4u) * 4u;

var k0: u32 = 0u;
loop {
if (k0 >= params.K) { break; }
let arow = row0 + ar;
if (arow < params.M) {
let base = arow * params.K + k0 + ac;
let av = t_input[base >> 2u];
As[ar * BK + ac + 0u] = f16(av.x); As[ar * BK + ac + 1u] = f16(av.y);
As[ar * BK + ac + 2u] = f16(av.z); As[ar * BK + ac + 3u] = f16(av.w);
let av0 = t_input[(base + 0u) >> 2u];
let av1 = t_input[(base + 16u) >> 2u];
let av2 = t_input[(base + 32u) >> 2u];
let av3 = t_input[(base + 48u) >> 2u];
As[ar * BK + ac + 0u] = f16(av0.x); As[ar * BK + ac + 1u] = f16(av0.y);
As[ar * BK + ac + 2u] = f16(av0.z); As[ar * BK + ac + 3u] = f16(av0.w);
As[ar * BK + ac + 16u] = f16(av1.x); As[ar * BK + ac + 17u] = f16(av1.y);
As[ar * BK + ac + 18u] = f16(av1.z); As[ar * BK + ac + 19u] = f16(av1.w);
As[ar * BK + ac + 32u] = f16(av2.x); As[ar * BK + ac + 33u] = f16(av2.y);
As[ar * BK + ac + 34u] = f16(av2.z); As[ar * BK + ac + 35u] = f16(av2.w);
As[ar * BK + ac + 48u] = f16(av3.x); As[ar * BK + ac + 49u] = f16(av3.y);
As[ar * BK + ac + 50u] = f16(av3.z); As[ar * BK + ac + 51u] = f16(av3.w);
} else {
As[ar * BK + ac + 0u] = 0.0h; As[ar * BK + ac + 1u] = 0.0h;
As[ar * BK + ac + 2u] = 0.0h; As[ar * BK + ac + 3u] = 0.0h;
for (var segment: u32 = 0u; segment < 4u; segment = segment + 1u) {
for (var ai: u32 = 0u; ai < 4u; ai = ai + 1u) {
As[ar * BK + ac + segment * 16u + ai] = 0.0h;
}
}
}
if (tid < BN) {
let c = tid;
Expand All @@ -64,10 +72,17 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
let scale_row = (k0 / params.group_size) * params.padded_N;
let scale = f16(t_scales[scale_row + n]);
let base_word = n * (params.K_packed >> 2u) + (k0 >> 3u);
let w0 = t_weight[base_word];
let w0 = t_weight[base_word + 0u];
let w1 = t_weight[base_word + 1u];
let w2 = t_weight[base_word + 2u];
let w3 = t_weight[base_word + 3u];
let w4 = t_weight[base_word + 4u];
let w5 = t_weight[base_word + 5u];
let w6 = t_weight[base_word + 6u];
let w7 = t_weight[base_word + 7u];
let words = array<u32, 8>(w0, w1, w2, w3, w4, w5, w6, w7);
for (var br: u32 = 0u; br < BK; br = br + 1u) {
let word = select(w1, w0, br < 8u);
let word = words[br >> 3u];
let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu;
Bs[br * BN + c] = f16(i32(nib) - 8) * scale;
}
Expand All @@ -91,7 +106,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) {
let r = row0 + lid.y * 4u + m;
let c = col0 + lid.x * 4u + n; // global fused column [0, 3072)
let c = col0 + lid.x * 4u + n;
if (r < params.M && c < params.N) {
var val = f32(acc[m][n]);
if (params.has_bias != 0u) { val = val + t_bias[c]; }
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,18 +12,11 @@

namespace executorch::backends::webgpu {

// @generated from q4gsw_linear_gemm_qkv_fused.wgsl - DO NOT EDIT.
// wgsl-sha256: 93e127e8ee4609d846015c8b75a600a29502e19a92bdf3a08e3429635f834085
inline constexpr const char* kQ4gswLinearGemmQkvFusedWGSL = R"(
// @generated from q4gsw_qkv_bk64.wgsl - DO NOT EDIT.
// wgsl-sha256: d738762f00f79ca16cf1549d47e6d1f51155f50805eec5e7e6df3bc07ee309ee
inline constexpr const char* kQ4gswQkvBk64WGSL = R"(
enable f16;
// Fused QKV q4gsw GEMM (Llama attention projections): one [M, N=3072] pwdq + f16-accumulate GEMM
// (vec4<f32> activation load) that scatter-writes each output column range to a SEPARATE buffer --
// c<2048 -> q, [2048,2560) -> k, [2560,3072) -> v. Replaces the 3 separate q/k/v linear dispatches;
// fixes the N=512 K/V occupancy starvation (16 WGs -> 96 WGs at M~128). Boundaries are 64-tile-aligned
// so each 64-col tile maps to exactly one output (uniform branch per workgroup). Per-output ROW STRIDE:
// q=2048, k=v=512. BIT-EXACT to 3 separate pwdqf16acc linears (fusing along N does not change the
// per-column K-accumulation order). Validated on Canary M4 Pro: correct (maxRel ~1e-3), scatter overhead
// 1.02x (free), concat win 1.63x on the QKV block. Boundaries hardcoded for Llama-3.2-1B GQA (32Q/8KV).

@group(0) @binding(0) var<storage, read_write> t_out_q: array<f32>;
@group(0) @binding(1) var<storage, read_write> t_out_k: array<f32>;
@group(0) @binding(2) var<storage, read_write> t_out_v: array<f32>;
Expand All @@ -42,10 +35,12 @@ struct Params {
_pad: u32,
}
@group(0) @binding(7) var<uniform> params: Params;
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u;

// BK64 QKV variant: group_size=64 keeps one scale valid for all eight packed words.
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 64u;
const N_Q: u32 = 2048u; const N_QK: u32 = 2560u; const N_KV: u32 = 512u;
var<workgroup> As: array<f16, 1024>;
var<workgroup> Bs: array<f16, 1024>;
var<workgroup> As: array<f16, 4096>;
var<workgroup> Bs: array<f16, 4096>;
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid: vec3<u32>) {
Expand All @@ -61,18 +56,31 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
}
let ar = tid / 4u;
let ac = (tid % 4u) * 4u;

var k0: u32 = 0u;
loop {
if (k0 >= params.K) { break; }
let arow = row0 + ar;
if (arow < params.M) {
let base = arow * params.K + k0 + ac;
let av = t_input[base >> 2u];
As[ar * BK + ac + 0u] = f16(av.x); As[ar * BK + ac + 1u] = f16(av.y);
As[ar * BK + ac + 2u] = f16(av.z); As[ar * BK + ac + 3u] = f16(av.w);
let av0 = t_input[(base + 0u) >> 2u];
let av1 = t_input[(base + 16u) >> 2u];
let av2 = t_input[(base + 32u) >> 2u];
let av3 = t_input[(base + 48u) >> 2u];
As[ar * BK + ac + 0u] = f16(av0.x); As[ar * BK + ac + 1u] = f16(av0.y);
As[ar * BK + ac + 2u] = f16(av0.z); As[ar * BK + ac + 3u] = f16(av0.w);
As[ar * BK + ac + 16u] = f16(av1.x); As[ar * BK + ac + 17u] = f16(av1.y);
As[ar * BK + ac + 18u] = f16(av1.z); As[ar * BK + ac + 19u] = f16(av1.w);
As[ar * BK + ac + 32u] = f16(av2.x); As[ar * BK + ac + 33u] = f16(av2.y);
As[ar * BK + ac + 34u] = f16(av2.z); As[ar * BK + ac + 35u] = f16(av2.w);
As[ar * BK + ac + 48u] = f16(av3.x); As[ar * BK + ac + 49u] = f16(av3.y);
As[ar * BK + ac + 50u] = f16(av3.z); As[ar * BK + ac + 51u] = f16(av3.w);
} else {
As[ar * BK + ac + 0u] = 0.0h; As[ar * BK + ac + 1u] = 0.0h;
As[ar * BK + ac + 2u] = 0.0h; As[ar * BK + ac + 3u] = 0.0h;
for (var segment: u32 = 0u; segment < 4u; segment = segment + 1u) {
for (var ai: u32 = 0u; ai < 4u; ai = ai + 1u) {
As[ar * BK + ac + segment * 16u + ai] = 0.0h;
}
}
}
if (tid < BN) {
let c = tid;
Expand All @@ -81,10 +89,17 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
let scale_row = (k0 / params.group_size) * params.padded_N;
let scale = f16(t_scales[scale_row + n]);
let base_word = n * (params.K_packed >> 2u) + (k0 >> 3u);
let w0 = t_weight[base_word];
let w0 = t_weight[base_word + 0u];
let w1 = t_weight[base_word + 1u];
let w2 = t_weight[base_word + 2u];
let w3 = t_weight[base_word + 3u];
let w4 = t_weight[base_word + 4u];
let w5 = t_weight[base_word + 5u];
let w6 = t_weight[base_word + 6u];
let w7 = t_weight[base_word + 7u];
let words = array<u32, 8>(w0, w1, w2, w3, w4, w5, w6, w7);
for (var br: u32 = 0u; br < BK; br = br + 1u) {
let word = select(w1, w0, br < 8u);
let word = words[br >> 3u];
let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu;
Bs[br * BN + c] = f16(i32(nib) - 8) * scale;
}
Expand All @@ -108,7 +123,7 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) {
let r = row0 + lid.y * 4u + m;
let c = col0 + lid.x * 4u + n; // global fused column [0, 3072)
let c = col0 + lid.x * 4u + n;
if (r < params.M && c < params.N) {
var val = f32(acc[m][n]);
if (params.has_bias != 0u) { val = val + t_bias[c]; }
Expand All @@ -121,8 +136,8 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
}
)";

inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeX = 16;
inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeY = 16;
inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeZ = 1;
inline constexpr uint32_t kQ4gswQkvBk64WorkgroupSizeX = 16;
inline constexpr uint32_t kQ4gswQkvBk64WorkgroupSizeY = 16;
inline constexpr uint32_t kQ4gswQkvBk64WorkgroupSizeZ = 1;

} // namespace executorch::backends::webgpu
Loading
Loading