use anyhow::Result;
use crate::GpuCtx;
use crate::encoder::{
GEMM2_TILES, GEMM3_TILES, act_code, enc_gemm2_src, enc_gemm3_src, gemm2_tier, gemm2_tile,
gemm3_tier, gemm3_tile,
};
use crate::encoder_weights::Act;
use crate::forward::{make_bg, pipeline, uni};
use crate::whisper::{Whisper, WhisperEncoder};
pub(crate) const VANILLA_ATTN: &str = r#"
struct Meta { t: u32, heads: u32, hd: u32, p0: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> k: array<f32>;
@group(0) @binding(2) var<storage, read> v: array<f32>;
@group(0) @binding(3) var<storage, read_write> outp: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
const TILE: u32 = 128u;
var<workgroup> q_sh: array<f32, 128>; // this row's q (hd ≤ 128)
var<workgroup> sc: array<f32, 128>; // tile scores, then tile probs
var<workgroup> red: array<f32, 128>; // reduction scratch (max, then sum)
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) tid: u32) {
let t = mt.t;
let hd = mt.hd;
let d = mt.heads * hd;
let i = wg.x;
let head = wg.y;
if (i >= t || head >= mt.heads) { return; }
let scale = 1.0 / sqrt(f32(hd));
let qbase = i * d + head * hd;
if (tid < hd) { q_sh[tid] = q[qbase + tid]; }
workgroupBarrier();
var acc = 0.0; // output for dim `tid` (tid < hd), across all tiles
var run_max = -1e30;
var run_den = 0.0;
var tile0 = 0u;
loop {
if (tile0 >= t) { break; }
let j = tile0 + tid;
// 1) score for this thread's key (thread ↔ key), dot over hd with no barrier
var s = -1e30;
if (j < t) {
let kbase = j * d + head * hd;
var dot = 0.0;
for (var c = 0u; c < hd; c = c + 1u) { dot = dot + q_sh[c] * k[kbase + c]; }
s = dot * scale;
}
sc[tid] = s;
red[tid] = s;
workgroupBarrier();
// 2) tile max
for (var st = 64u; st > 0u; st = st >> 1u) {
if (tid < st) { red[tid] = max(red[tid], red[tid + st]); }
workgroupBarrier();
}
let nm = max(run_max, red[0]);
let corr = exp(run_max - nm);
workgroupBarrier();
// 3) probs relative to the new running max
var p = 0.0;
if (j < t) { p = exp(sc[tid] - nm); }
sc[tid] = p;
red[tid] = p;
workgroupBarrier();
// 4) tile denom
for (var st = 64u; st > 0u; st = st >> 1u) {
if (tid < st) { red[tid] = red[tid] + red[tid + st]; }
workgroupBarrier();
}
run_den = run_den * corr + red[0];
// 5) accumulate output (thread ↔ head-dim); v read once per key
if (tid < hd) {
var a = acc * corr;
let n = min(TILE, t - tile0);
for (var r = 0u; r < n; r = r + 1u) {
a = a + sc[r] * v[(tile0 + r) * d + head * hd + tid];
}
acc = a;
}
run_max = nm;
workgroupBarrier(); // sc/red consumed before the next tile overwrites them
tile0 = tile0 + TILE;
}
if (tid < hd) { outp[qbase + tid] = acc / run_den; }
}
"#;
pub(crate) const LN_SRC: &str = r#"
struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) var<storage, read> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> outp: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> sh: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let base = wg.x * mt.h;
var sum = 0.0;
for (var i = t; i < mt.h; i = i + 256u) { sum = sum + x[base + i]; }
sh[t] = sum;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let mean = sh[0] / f32(mt.h);
workgroupBarrier();
var sq = 0.0;
for (var i = t; i < mt.h; i = i + 256u) { let d = x[base + i] - mean; sq = sq + d * d; }
sh[t] = sq;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let inv = 1.0 / sqrt(sh[0] / f32(mt.h) + 1e-5);
for (var i = t; i < mt.h; i = i + 256u) { outp[base + i] = (x[base + i] - mean) * inv * w[i] + b[i]; }
}
"#;
pub(crate) const ADD_SRC: &str = r#"
struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> dst: array<f32>;
@group(0) @binding(1) var<storage, read> src: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
if (gid.x < mt.n) { dst[gid.x] = dst[gid.x] + src[gid.x]; }
}
"#;
pub(crate) struct GpuLinear {
pub(crate) w: wgpu::Buffer,
pub(crate) b: wgpu::Buffer,
pub(crate) n: u32,
pub(crate) k: u32,
pub(crate) v3: bool,
}
impl GpuLinear {
pub(crate) fn new_v3(ctx: &GpuCtx, w: &[f32], b: &[f32], n: usize, k: usize) -> Self {
let mut wt = vec![0f32; n * k];
for nn in 0..n {
for kk in 0..k {
wt[kk * n + nn] = w[nn * k + kk];
}
}
Self {
w: ctx.storage(&wt),
b: ctx.storage(b),
n: n as u32,
k: k as u32,
v3: true,
}
}
}
pub(crate) struct GpuNorm {
pub(crate) w: wgpu::Buffer,
pub(crate) b: wgpu::Buffer,
}
struct GpuBlock {
n_attn: GpuNorm,
q: GpuLinear,
k: GpuLinear,
v: GpuLinear,
out: GpuLinear,
n_ff: GpuNorm,
fc1: GpuLinear,
fc2: GpuLinear,
}
pub struct WhisperEncoderGpu {
cpu: WhisperEncoder,
gemm2: Vec<wgpu::ComputePipeline>,
gemm3: Vec<wgpu::ComputePipeline>,
ln: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
ln_post: Option<GpuNorm>,
blocks: Vec<GpuBlock>,
d: usize,
heads: usize,
hd: usize,
ff: usize,
max_t: usize,
}
impl WhisperEncoderGpu {
pub fn new(ctx: &GpuCtx, cpu: WhisperEncoder, max_t: usize) -> Result<Self> {
let d = cpu.d;
let heads = cpu.blocks[0].attn.heads;
let hd = cpu.blocks[0].attn.hd;
anyhow::ensure!(hd <= 128, "attention supports head_dim ≤ 128");
let ff = cpu.blocks[0].fc1.n;
let norm = |g: &crate::whisper::LayerNorm| GpuNorm {
w: ctx.storage(&g.w),
b: ctx.storage(&g.b),
};
let lin = |l: &crate::whisper::Linear| GpuLinear::new_v3(ctx, &l.w, &l.b, l.n, l.k);
let blocks = cpu
.blocks
.iter()
.map(|b| GpuBlock {
n_attn: norm(&b.norm_attn),
q: lin(&b.attn.q),
k: lin(&b.attn.k),
v: lin(&b.attn.v),
out: lin(&b.attn.out),
n_ff: norm(&b.norm_ff),
fc1: lin(&b.fc1),
fc2: lin(&b.fc2),
})
.collect();
let ln_post = cpu.ln_post.as_ref().map(norm);
Ok(Self {
gemm2: GEMM2_TILES
.iter()
.map(|&(bm, bn)| pipeline(ctx, "w_gemm2", &enc_gemm2_src(false, bm, bn)))
.collect(),
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "w_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
ln: pipeline(ctx, "w_ln", LN_SRC),
add: pipeline(ctx, "w_add", ADD_SRC),
attn: pipeline(ctx, "w_attn", VANILLA_ATTN),
cpu,
ln_post,
blocks,
d,
heads,
hd,
ff,
max_t,
})
}
pub fn cpu(&self) -> &WhisperEncoder {
&self.cpu
}
pub fn forward(&self, ctx: &GpuCtx, mel: &[f32]) -> Result<Vec<f32>> {
let (stem, t) = self.cpu.stem(mel);
anyhow::ensure!(t <= self.max_t, "t {t} exceeds max_t {}", self.max_t);
let (d, ff) = (self.d, self.ff);
let xb = ctx.storage(&stem);
let nb = ctx.empty(t * d);
let qb = ctx.empty(t * d);
let kb = ctx.empty(t * d);
let vb = ctx.empty(t * d);
let sb = ctx.empty(t * d);
let hb = ctx.empty(t * ff);
let mut passes: Vec<(&wgpu::ComputePipeline, wgpu::BindGroup, u32, u32)> = Vec::new();
let mut keep: Vec<wgpu::Buffer> = Vec::new();
macro_rules! gemm {
($x:expr, $lw:expr, $y:expr, $act:expr) => {{
let lw: &GpuLinear = $lw;
let flags = 1u32 | (act_code($act) << 8);
let meta = uni(ctx, bytemuck::cast_slice(&[t as u32, lw.n, lw.k, flags]));
let (pl, bm, bn) = if lw.v3 {
let tile = gemm3_tile(t, lw.n as usize);
(&self.gemm3[gemm3_tier(tile)], tile.0, tile.1)
} else {
let tile = gemm2_tile(t, lw.n as usize);
(&self.gemm2[gemm2_tier(tile)], tile.0, tile.1)
};
let bg = make_bg(ctx, pl, &[$x, &lw.w, &lw.b, $y], &meta);
passes.push((
pl,
bg,
lw.n.div_ceil(bn as u32),
(t as u32).div_ceil(bm as u32),
));
keep.push(meta);
}};
}
macro_rules! ln {
($x:expr, $n:expr, $y:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[d as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.ln, &[$x, &$n.w, &$n.b, $y], &meta);
passes.push((&self.ln, bg, t as u32, 1));
keep.push(meta);
}};
}
macro_rules! add {
($dst:expr, $src:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[(t * d) as u32, 0u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.add, &[$dst, $src], &meta);
passes.push((&self.add, bg, ((t * d) as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
for b in &self.blocks {
ln!(&xb, b.n_attn, &nb);
gemm!(&nb, &b.q, &qb, None);
gemm!(&nb, &b.k, &kb, None);
gemm!(&nb, &b.v, &vb, None);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[t as u32, self.heads as u32, self.hd as u32, 0u32]),
);
let bg = make_bg(ctx, &self.attn, &[&qb, &kb, &vb, &nb], &meta);
passes.push((&self.attn, bg, t as u32, self.heads as u32));
keep.push(meta);
}
gemm!(&nb, &b.out, &sb, None);
add!(&xb, &sb);
ln!(&xb, b.n_ff, &nb);
gemm!(&nb, &b.fc1, &hb, Some(Act::GeluErf));
gemm!(&hb, &b.fc2, &sb, None);
add!(&xb, &sb);
}
let out_buf = match &self.ln_post {
Some(ln) => {
ln!(&xb, ln, &nb);
&nb
}
None => &xb,
};
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("whisper_enc"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("whisper_enc"),
timestamp_writes: None,
});
for (pl, bg, gx, gy) in &passes {
cpass.set_pipeline(pl);
cpass.set_bind_group(0, bg, &[]);
cpass.dispatch_workgroups(*gx, *gy, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
let out = ctx.read(out_buf, t * d)?;
drop(keep);
Ok(out)
}
}
const DEC_EMBED: &str = r#"
struct Meta { id: u32, p: u32, d: u32, pad: u32 }
@group(0) @binding(0) var<storage, read> et: array<f32>;
@group(0) @binding(1) var<storage, read> ep: array<f32>;
@group(0) @binding(2) var<storage, read_write> x: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= mt.d) { return; }
x[g.x] = et[mt.id * mt.d + g.x] + ep[mt.p * mt.d + g.x];
}
"#;
const DEC_WRITE: &str = r#"
struct Meta { p: u32, d: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= mt.d) { return; }
dst[mt.p * mt.d + g.x] = src[g.x];
}
"#;
const DEC_QATTN: &str = r#"
struct Meta { heads: u32, hd: u32, t: u32, pad: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>; // [hid]
@group(0) @binding(1) var<storage, read> k: array<f32>; // [t, hid]
@group(0) @binding(2) var<storage, read> v: array<f32>; // [t, hid]
@group(0) @binding(3) var<storage, read_write> o: array<f32>; // [hid]
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> q_sh: array<f32, 128>;
var<workgroup> sc: array<f32, 1600>; // scores/probs (>= 1500 encoder frames)
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let head = wg.x;
let hd = mt.hd;
let hid = mt.heads * hd;
let tk = mt.t;
let scale = 1.0 / sqrt(f32(hd));
if (t < hd) { q_sh[t] = q[head * hd + t]; }
workgroupBarrier();
for (var j = t; j < tk; j += 256u) {
let kb = j * hid + head * hd;
var d = 0.0;
for (var c = 0u; c < hd; c += 1u) { d += q_sh[c] * k[kb + c]; }
sc[j] = d * scale;
}
workgroupBarrier();
var lm = -3.0e38;
for (var j = t; j < tk; j += 256u) { lm = max(lm, sc[j]); }
red[t] = lm;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) { if (t < s) { red[t] = max(red[t], red[t + s]); } workgroupBarrier(); }
let m = red[0];
workgroupBarrier();
var ls = 0.0;
for (var j = t; j < tk; j += 256u) { let e = exp(sc[j] - m); sc[j] = e; ls += e; }
red[t] = ls;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) { if (t < s) { red[t] += red[t + s]; } workgroupBarrier(); }
let sm = red[0];
workgroupBarrier();
if (t < hd) {
var acc = 0.0;
for (var j = 0u; j < tk; j += 1u) { acc += sc[j] * v[j * hid + head * hd + t]; }
o[head * hd + t] = acc / sm;
}
}
"#;
const DEC_ARGMAX: &str = r#"
struct Meta { vocab: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> logits: array<f32>;
@group(0) @binding(1) var<storage, read> mask: array<u32>;
@group(0) @binding(2) var<storage, read_write> tid: array<u32>;
@group(0) @binding(3) var<uniform> mt: Meta;
var<workgroup> vb: array<f32, 256>;
var<workgroup> ib: array<u32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(local_invocation_index) t: u32) {
var lv = -3.0e38;
var li = 0u;
for (var j = t; j < mt.vocab; j += 256u) {
if (mask[j] == 0u && logits[j] > lv) { lv = logits[j]; li = j; }
}
vb[t] = lv;
ib[t] = li;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) {
if (t < s) {
if (vb[t + s] > vb[t] || (vb[t + s] == vb[t] && ib[t + s] < ib[t])) {
vb[t] = vb[t + s];
ib[t] = ib[t + s];
}
}
workgroupBarrier();
}
if (t == 0u) { tid[0] = ib[0]; }
}
"#;
struct DecGpuBlock {
n_self: GpuNorm,
sq: GpuLinear,
sk: GpuLinear,
sv: GpuLinear,
so: GpuLinear,
n_cross: GpuNorm,
cq: GpuLinear,
ck: GpuLinear,
cv: GpuLinear,
co: GpuLinear,
n_ff: GpuNorm,
fc1: GpuLinear,
fc2: GpuLinear,
}
pub struct WhisperDecoderGpu {
gemm3: Vec<wgpu::ComputePipeline>,
ln: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
embed_pl: wgpu::ComputePipeline,
write_pl: wgpu::ComputePipeline,
qattn_pl: wgpu::ComputePipeline,
argmax_pl: wgpu::ComputePipeline,
embed_tokens: wgpu::Buffer,
embed_pos: wgpu::Buffer,
blocks: Vec<DecGpuBlock>,
ln_post: GpuNorm,
lm_head: GpuLinear,
mask_gen: wgpu::Buffer,
mask_gen0: wgpu::Buffer,
mask_lang: wgpu::Buffer,
d: usize,
heads: usize,
hd: usize,
ff: usize,
vocab: usize,
max_target: usize,
prompt: Vec<usize>,
lang_on: bool,
transcribe_token: usize,
notimestamps_token: usize,
eos: usize,
}
impl WhisperDecoderGpu {
pub fn new(ctx: &GpuCtx, cpu: &Whisper) -> Result<Self> {
let d = cpu.d;
let heads = cpu.blocks[0].self_attn.heads;
let hd = cpu.blocks[0].self_attn.hd;
anyhow::ensure!(hd <= 128, "decoder head_dim {hd} > 128");
let ff = cpu.blocks[0].fc1.n;
let vocab = cpu.vocab;
let norm = |g: &crate::whisper::LayerNorm| GpuNorm {
w: ctx.storage(&g.w),
b: ctx.storage(&g.b),
};
let lin = |l: &crate::whisper::Linear| GpuLinear::new_v3(ctx, &l.w, &l.b, l.n, l.k);
let blocks = cpu
.blocks
.iter()
.map(|b| DecGpuBlock {
n_self: norm(&b.norm_self),
sq: lin(&b.self_attn.q),
sk: lin(&b.self_attn.k),
sv: lin(&b.self_attn.v),
so: lin(&b.self_attn.out),
n_cross: norm(&b.norm_cross),
cq: lin(&b.cross_attn.q),
ck: lin(&b.cross_attn.k),
cv: lin(&b.cross_attn.v),
co: lin(&b.cross_attn.out),
n_ff: norm(&b.norm_ff),
fc1: lin(&b.fc1),
fc2: lin(&b.fc2),
})
.collect();
let mask_of = |set: &[usize], base_ones: bool| -> wgpu::Buffer {
let mut m = vec![if base_ones { 1u32 } else { 0u32 }; vocab];
for &t in set {
if t < vocab {
m[t] = if base_ones { 0 } else { 1 };
}
}
ctx.storage_bytes(bytemuck::cast_slice(&m))
};
let mask_gen = mask_of(&cpu.suppress, false);
let mut m0 = vec![0u32; vocab];
for &t in cpu.suppress.iter().chain(cpu.begin_suppress.iter()) {
if t < vocab {
m0[t] = 1;
}
}
let mask_gen0 = ctx.storage_bytes(bytemuck::cast_slice(&m0));
let mask_lang = mask_of(&cpu.lang_tokens, true);
Ok(Self {
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "wd_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
ln: pipeline(ctx, "wd_ln", LN_SRC),
add: pipeline(ctx, "wd_add", ADD_SRC),
embed_pl: pipeline(ctx, "wd_embed", DEC_EMBED),
write_pl: pipeline(ctx, "wd_write", DEC_WRITE),
qattn_pl: pipeline(ctx, "wd_qattn", DEC_QATTN),
argmax_pl: pipeline(ctx, "wd_argmax", DEC_ARGMAX),
embed_tokens: ctx.storage(&cpu.embed_tokens),
embed_pos: ctx.storage(&cpu.embed_pos),
ln_post: norm(&cpu.ln_post),
lm_head: lin(&cpu.lm_head),
blocks,
mask_gen,
mask_gen0,
mask_lang,
d,
heads,
hd,
ff,
vocab,
max_target: cpu.max_target,
prompt: cpu.prompt.clone(),
lang_on: !cpu.lang_tokens.is_empty(),
transcribe_token: cpu.transcribe_token,
notimestamps_token: cpu.notimestamps_token,
eos: cpu.eos,
})
}
pub fn generate(&self, ctx: &GpuCtx, enc: &[f32], max_new: usize) -> Result<Vec<usize>> {
let d = self.d;
let (heads, hd, ff, vocab) = (self.heads, self.hd, self.ff, self.vocab);
let n_cross = enc.len() / d;
let b_enc = ctx.storage(enc);
let ckv: Vec<(wgpu::Buffer, wgpu::Buffer)> = self
.blocks
.iter()
.map(|_| (ctx.empty(n_cross * d), ctx.empty(n_cross * d)))
.collect();
let kc: Vec<wgpu::Buffer> = self
.blocks
.iter()
.map(|_| ctx.empty(self.max_target * d))
.collect();
let vc: Vec<wgpu::Buffer> = self
.blocks
.iter()
.map(|_| ctx.empty(self.max_target * d))
.collect();
let x = ctx.empty(d);
let normed = ctx.empty(d);
let q = ctx.empty(d);
let k_new = ctx.empty(d);
let v_new = ctx.empty(d);
let sa = ctx.empty(d);
let ca = ctx.empty(d);
let ffb = ctx.empty(ff);
let logits = ctx.empty(vocab);
let tid = ctx.empty(1);
type Pass<'a> = (&'a wgpu::ComputePipeline, wgpu::BindGroup, u32, u32);
macro_rules! gemm {
($passes:expr, $keep:expr, $x:expr, $l:expr, $y:expr, $m:expr, $act:expr) => {{
let l: &GpuLinear = $l;
let m: usize = $m;
let flags = 1u32 | (act_code($act) << 8);
let meta = uni(ctx, bytemuck::cast_slice(&[m as u32, l.n, l.k, flags]));
let tile = gemm3_tile(m, l.n as usize);
let pl = &self.gemm3[gemm3_tier(tile)];
let bg = make_bg(ctx, pl, &[$x, &l.w, &l.b, $y], &meta);
$passes.push((
pl,
bg,
l.n.div_ceil(tile.1 as u32),
(m as u32).div_ceil(tile.0 as u32),
));
$keep.push(meta);
}};
}
macro_rules! ln {
($passes:expr, $keep:expr, $x:expr, $nm:expr, $y:expr) => {{
let nm: &GpuNorm = $nm;
let meta = uni(ctx, bytemuck::cast_slice(&[d as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.ln, &[$x, &nm.w, &nm.b, $y], &meta);
$passes.push((&self.ln, bg, 1u32, 1u32));
$keep.push(meta);
}};
}
macro_rules! add {
($passes:expr, $keep:expr, $dst:expr, $src:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[d as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.add, &[$dst, $src], &meta);
$passes.push((&self.add, bg, (d as u32).div_ceil(256), 1u32));
$keep.push(meta);
}};
}
macro_rules! write {
($passes:expr, $keep:expr, $src:expr, $dst:expr, $pos:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[$pos as u32, d as u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.write_pl, &[$src, $dst], &meta);
$passes.push((&self.write_pl, bg, (d as u32).div_ceil(256), 1u32));
$keep.push(meta);
}};
}
macro_rules! qattn {
($passes:expr, $keep:expr, $q:expr, $k:expr, $v:expr, $o:expr, $tk:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[heads as u32, hd as u32, $tk as u32, 0u32]),
);
let bg = make_bg(ctx, &self.qattn_pl, &[$q, $k, $v, $o], &meta);
$passes.push((&self.qattn_pl, bg, heads as u32, 1u32));
$keep.push(meta);
}};
}
{
let mut passes: Vec<Pass> = Vec::new();
let mut keep: Vec<wgpu::Buffer> = Vec::new();
for (b, (cbk, cbv)) in self.blocks.iter().zip(&ckv) {
gemm!(passes, keep, &b_enc, &b.ck, cbk, n_cross, None);
gemm!(passes, keep, &b_enc, &b.cv, cbv, n_cross, None);
}
record_submit(ctx, &passes);
drop(keep);
}
let step =
|token: usize, pos: usize, emit: Option<&wgpu::Buffer>| -> Result<Option<usize>> {
let mut passes: Vec<Pass> = Vec::new();
let mut keep: Vec<wgpu::Buffer> = Vec::new();
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[token as u32, pos as u32, d as u32, 0u32]),
);
let bg = make_bg(
ctx,
&self.embed_pl,
&[&self.embed_tokens, &self.embed_pos, &x],
&meta,
);
passes.push((&self.embed_pl, bg, (d as u32).div_ceil(256), 1u32));
keep.push(meta);
}
for (bi, b) in self.blocks.iter().enumerate() {
ln!(passes, keep, &x, &b.n_self, &normed);
gemm!(passes, keep, &normed, &b.sq, &q, 1, None);
gemm!(passes, keep, &normed, &b.sk, &k_new, 1, None);
gemm!(passes, keep, &normed, &b.sv, &v_new, 1, None);
write!(passes, keep, &k_new, &kc[bi], pos);
write!(passes, keep, &v_new, &vc[bi], pos);
qattn!(passes, keep, &q, &kc[bi], &vc[bi], &sa, pos + 1);
gemm!(passes, keep, &sa, &b.so, &normed, 1, None);
add!(passes, keep, &x, &normed);
ln!(passes, keep, &x, &b.n_cross, &normed);
gemm!(passes, keep, &normed, &b.cq, &q, 1, None);
qattn!(passes, keep, &q, &ckv[bi].0, &ckv[bi].1, &ca, n_cross);
gemm!(passes, keep, &ca, &b.co, &normed, 1, None);
add!(passes, keep, &x, &normed);
ln!(passes, keep, &x, &b.n_ff, &normed);
gemm!(passes, keep, &normed, &b.fc1, &ffb, 1, Some(Act::GeluErf));
gemm!(passes, keep, &ffb, &b.fc2, &normed, 1, None);
add!(passes, keep, &x, &normed);
}
if let Some(mask) = emit {
ln!(passes, keep, &x, &self.ln_post, &normed);
gemm!(passes, keep, &normed, &self.lm_head, &logits, 1, None);
let meta = uni(ctx, bytemuck::cast_slice(&[vocab as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.argmax_pl, &[&logits, mask, &tid], &meta);
passes.push((&self.argmax_pl, bg, 1u32, 1u32));
keep.push(meta);
}
record_submit(ctx, &passes);
drop(keep);
if emit.is_some() {
Ok(Some(ctx.read(&tid, 1)?[0].to_bits() as usize))
} else {
Ok(None)
}
};
let sot = self.prompt[0];
let lang = if self.lang_on {
step(sot, 0, Some(&self.mask_lang))?.expect("emit")
} else {
self.prompt[1]
};
let prompt = vec![sot, lang, self.transcribe_token, self.notimestamps_token];
let mut ids = prompt.clone();
for (i, &t) in prompt.iter().enumerate().take(prompt.len() - 1) {
step(t, i, None)?; }
let cap = (prompt.len() + max_new).min(self.max_target);
let mut pos = prompt.len() - 1;
let mut tok = prompt[pos];
let mut first = true;
loop {
let mask = if first {
&self.mask_gen0
} else {
&self.mask_gen
};
let nxt = step(tok, pos, Some(mask))?.expect("emit");
ids.push(nxt);
if nxt == self.eos || ids.len() >= cap {
break;
}
pos += 1;
tok = nxt;
first = false;
}
Ok(ids)
}
}
fn record_submit(ctx: &GpuCtx, passes: &[(&wgpu::ComputePipeline, wgpu::BindGroup, u32, u32)]) {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("whisper_dec"),
});
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for (pl, bg, gx, gy) in passes {
p.set_pipeline(pl);
p.set_bind_group(0, bg, &[]);
p.dispatch_workgroups(*gx, *gy, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
}