Skip to content
Merged
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
844 changes: 844 additions & 0 deletions backends/webgpu/runtime/WebGPUGraph.cpp

Large diffs are not rendered by default.

37 changes: 37 additions & 0 deletions backends/webgpu/runtime/WebGPUGraph.h
Original file line number Diff line number Diff line change
Expand Up @@ -450,6 +450,43 @@ 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
size_t sep_dispatch[3] = {
0,
0,
0}; // their dispatch indices (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);

// SwiGLU fusion: emit ONE fused elementwise dispatch
// computing out = (gate * sigmoid(gate)) * up, replacing the sigmoid + 2
// muls. `out` is repointed to a private pooled buffer (aliasing guard);
// `gate` is likewise given a private pooled buffer at its producer op by the
// build() walk (the planner reuse-aliases up onto gate's slot, so up_proj
// would stomp gate before the fused reads it). Only used during build(); the
// detection maps are empty (inert) when no SwiGLU triple matches.
void add_swiglu_fused_dispatch(int gate_id, int up_id, int out_id);
};

} // namespace executorch::backends::webgpu
24 changes: 24 additions & 0 deletions backends/webgpu/runtime/ops/mul/silu_mul_fused.wgsl
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
@group(0) @binding(0) var<storage, read> gate: array<f32>;
@group(0) @binding(1) var<storage, read> up: array<f32>;
@group(0) @binding(2) var<storage, read_write> output: array<f32>;

struct Params {
num_elements: u32,
}
@group(0) @binding(3) var<uniform> params: Params;

// Fused SwiGLU activation: output = (g * sigmoid(g)) * up, folding the separate
// sigmoid(gate) -> mul(gate,sig)=silu -> mul(silu,up) triple into one dispatch.
// sigmoid + silu are computed in registers (never written to memory), so gate + up
// are read once and one output is written. The sigmoid form (1/(1+exp(-x))) and the
// multiply order match the original ops -> bit-exact.
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let idx = gid.x;
if (idx >= params.num_elements) {
return;
}
let g = gate[idx];
let sig = 1.0 / (1.0 + exp(-g));
output[idx] = (g * sig) * up[idx];
}
48 changes: 48 additions & 0 deletions backends/webgpu/runtime/ops/mul/silu_mul_fused_wgsl.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/

#pragma once

#include <cstdint>

namespace executorch::backends::webgpu {

// @generated from silu_mul_fused.wgsl - DO NOT EDIT.
// wgsl-sha256: 4b8ede66c5dbc9829ff48f745eb9ad48fa5a5200058baa532fbf34f78ec2f560
inline constexpr const char* kSiluMulFusedWGSL = R"(
@group(0) @binding(0) var<storage, read> gate: array<f32>;
@group(0) @binding(1) var<storage, read> up: array<f32>;
@group(0) @binding(2) var<storage, read_write> output: array<f32>;

struct Params {
num_elements: u32,
}
@group(0) @binding(3) var<uniform> params: Params;

// Fused SwiGLU activation: output = (g * sigmoid(g)) * up, folding the separate
// sigmoid(gate) -> mul(gate,sig)=silu -> mul(silu,up) triple into one dispatch.
// sigmoid + silu are computed in registers (never written to memory), so gate + up
// are read once and one output is written. The sigmoid form (1/(1+exp(-x))) and the
// multiply order match the original ops -> bit-exact.
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let idx = gid.x;
if (idx >= params.num_elements) {
return;
}
let g = gate[idx];
let sig = 1.0 / (1.0 + exp(-g));
output[idx] = (g * sig) * up[idx];
}
)";

inline constexpr uint32_t kSiluMulFusedWorkgroupSizeX = 64;
inline constexpr uint32_t kSiluMulFusedWorkgroupSizeY = 1;
inline constexpr uint32_t kSiluMulFusedWorkgroupSizeZ = 1;

} // namespace executorch::backends::webgpu
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
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>;
@group(0) @binding(3) var<storage, read> t_input: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read> t_weight: array<u32>;
@group(0) @binding(5) var<storage, read> t_scales: array<f32>;
@group(0) @binding(6) var<storage, read> t_bias: array<f32>;
struct Params {
M: u32,
N: u32,
K: u32,
K_packed: u32,
group_size: u32,
padded_N: u32,
has_bias: u32,
_pad: u32,
}
@group(0) @binding(7) var<uniform> params: Params;
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u;
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>;
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid: vec3<u32>) {
let nbN = (params.N + BN - 1u) / BN;
let bx = wid.x % nbN;
let by = wid.x / nbN;
let row0 = by * BM;
let col0 = bx * BN;
let tid = lid.y * 16u + lid.x;
var acc: array<array<f16, 4>, 4>;
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0h; }
}
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);
} 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;
}
if (tid < BN) {
let c = tid;
let n = col0 + c;
if (n < params.N) {
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 w1 = t_weight[base_word + 1u];
for (var br: u32 = 0u; br < BK; br = br + 1u) {
let word = select(w1, w0, br < 8u);
let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu;
Bs[br * BN + c] = f16(i32(nib) - 8) * scale;
}
} else {
for (var br: u32 = 0u; br < BK; br = br + 1u) { Bs[br * BN + c] = 0.0h; }
}
}
workgroupBarrier();
for (var k: u32 = 0u; k < BK; k = k + 1u) {
var a: array<f16, 4>;
var bvec: array<f16, 4>;
for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; }
for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; }
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = fma(a[m], bvec[n], acc[m][n]); }
}
}
workgroupBarrier();
k0 = k0 + BK;
}
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)
if (r < params.M && c < params.N) {
var val = f32(acc[m][n]);
if (params.has_bias != 0u) { val = val + t_bias[c]; }
if (c < N_Q) { t_out_q[r * N_Q + c] = val; }
else if (c < N_QK) { t_out_k[r * N_KV + (c - N_Q)] = val; }
else { t_out_v[r * N_KV + (c - N_QK)] = val; }
}
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,128 @@
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/

#pragma once

#include <cstdint>

namespace executorch::backends::webgpu {

// @generated from q4gsw_linear_gemm_qkv_fused.wgsl - DO NOT EDIT.
// wgsl-sha256: 93e127e8ee4609d846015c8b75a600a29502e19a92bdf3a08e3429635f834085
inline constexpr const char* kQ4gswLinearGemmQkvFusedWGSL = 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>;
@group(0) @binding(3) var<storage, read> t_input: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read> t_weight: array<u32>;
@group(0) @binding(5) var<storage, read> t_scales: array<f32>;
@group(0) @binding(6) var<storage, read> t_bias: array<f32>;
struct Params {
M: u32,
N: u32,
K: u32,
K_packed: u32,
group_size: u32,
padded_N: u32,
has_bias: u32,
_pad: u32,
}
@group(0) @binding(7) var<uniform> params: Params;
const BM: u32 = 64u; const BN: u32 = 64u; const BK: u32 = 16u;
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>;
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_id) lid: vec3<u32>) {
let nbN = (params.N + BN - 1u) / BN;
let bx = wid.x % nbN;
let by = wid.x / nbN;
let row0 = by * BM;
let col0 = bx * BN;
let tid = lid.y * 16u + lid.x;
var acc: array<array<f16, 4>, 4>;
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = 0.0h; }
}
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);
} 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;
}
if (tid < BN) {
let c = tid;
let n = col0 + c;
if (n < params.N) {
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 w1 = t_weight[base_word + 1u];
for (var br: u32 = 0u; br < BK; br = br + 1u) {
let word = select(w1, w0, br < 8u);
let nib = (word >> ((br & 7u) * 4u)) & 0x0Fu;
Bs[br * BN + c] = f16(i32(nib) - 8) * scale;
}
} else {
for (var br: u32 = 0u; br < BK; br = br + 1u) { Bs[br * BN + c] = 0.0h; }
}
}
workgroupBarrier();
for (var k: u32 = 0u; k < BK; k = k + 1u) {
var a: array<f16, 4>;
var bvec: array<f16, 4>;
for (var m: u32 = 0u; m < 4u; m = m + 1u) { a[m] = As[(lid.y * 4u + m) * BK + k]; }
for (var n: u32 = 0u; n < 4u; n = n + 1u) { bvec[n] = Bs[k * BN + lid.x * 4u + n]; }
for (var m: u32 = 0u; m < 4u; m = m + 1u) {
for (var n: u32 = 0u; n < 4u; n = n + 1u) { acc[m][n] = fma(a[m], bvec[n], acc[m][n]); }
}
}
workgroupBarrier();
k0 = k0 + BK;
}
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)
if (r < params.M && c < params.N) {
var val = f32(acc[m][n]);
if (params.has_bias != 0u) { val = val + t_bias[c]; }
if (c < N_Q) { t_out_q[r * N_Q + c] = val; }
else if (c < N_QK) { t_out_k[r * N_KV + (c - N_Q)] = val; }
else { t_out_v[r * N_KV + (c - N_QK)] = val; }
}
}
}
}
)";

inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeX = 16;
inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeY = 16;
inline constexpr uint32_t kQ4gswLinearGemmQkvFusedWorkgroupSizeZ = 1;

} // namespace executorch::backends::webgpu
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
$if DTYPE == "half":
enable f16;
@group(0) @binding(0) var<storage, read_write> t_out: array<f32>;
@group(0) @binding(1) var<storage, read> t_input: array<f32>;
@group(0) @binding(1) var<storage, read> t_input: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read> t_weight: array<u32>;
@group(0) @binding(3) var<storage, read> t_scales: array<f32>;
@group(0) @binding(4) var<storage, read> t_bias: array<f32>;
Expand Down Expand Up @@ -63,10 +63,12 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
let arow = row0 + ar;
if (arow < params.M) {
let base = arow * params.K + k0 + ac;
As[ar * BK + ac + 0u] = ${buffer_scalar_type(DTYPE)}(t_input[base]);
As[ar * BK + ac + 1u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 1u]);
As[ar * BK + ac + 2u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 2u]);
As[ar * BK + ac + 3u] = ${buffer_scalar_type(DTYPE)}(t_input[base + 3u]);
// vec4<f32> coalesced load; base is 4-aligned on the steel route (K%16==0, ac/k0 multiples of 4).
let av = t_input[base >> 2u];
As[ar * BK + ac + 0u] = ${buffer_scalar_type(DTYPE)}(av.x);
As[ar * BK + ac + 1u] = ${buffer_scalar_type(DTYPE)}(av.y);
As[ar * BK + ac + 2u] = ${buffer_scalar_type(DTYPE)}(av.z);
As[ar * BK + ac + 3u] = ${buffer_scalar_type(DTYPE)}(av.w);
} else {
As[ar * BK + ac + 0u] = ${"0.0h" if PWDQ else "0.0"}; As[ar * BK + ac + 1u] = ${"0.0h" if PWDQ else "0.0"};
As[ar * BK + ac + 2u] = ${"0.0h" if PWDQ else "0.0"}; As[ar * BK + ac + 3u] = ${"0.0h" if PWDQ else "0.0"};
Expand Down
Loading
Loading