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::parakeet::Parakeet;
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];
}
}
"#;
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]; }
}
"#;
const CP_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] = src[gid.x]; }
}
"#;
struct GpuLinear {
w: wgpu::Buffer,
b: wgpu::Buffer,
n: u32,
k: u32,
v3: bool,
}
impl GpuLinear {
fn new(ctx: &GpuCtx, w: &[f32], b: &[f32], n: usize, k: usize, scale: f32) -> Self {
let (ws, bs): (Vec<f32>, Vec<f32>) = if scale == 1.0 {
(w.to_vec(), b.to_vec())
} else {
(
w.iter().map(|x| x * scale).collect(),
b.iter().map(|x| x * scale).collect(),
)
};
Self {
w: ctx.storage(&ws),
b: ctx.storage(&bs),
n: n as u32,
k: k as u32,
v3: false,
}
}
fn new_v3(ctx: &GpuCtx, w: &[f32], b: &[f32], n: usize, k: usize, scale: f32) -> 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] * scale;
}
}
let bs: Vec<f32> = b.iter().map(|x| x * scale).collect();
Self {
w: ctx.storage(&wt),
b: ctx.storage(&bs),
n: n as u32,
k: k as u32,
v3: true,
}
}
}
struct GpuNorm {
w: wgpu::Buffer,
b: wgpu::Buffer,
}
struct GpuBlock {
n_ff1: GpuNorm,
ff1_1: GpuLinear,
ff1_2: GpuLinear, n_attn: GpuNorm,
lq: GpuLinear,
lk: GpuLinear,
lv: GpuLinear,
lpos: GpuLinear,
lout: GpuLinear,
ub: wgpu::Buffer, vbb: wgpu::Buffer,
n_conv: GpuNorm,
pw1: GpuLinear,
dw: wgpu::Buffer,
bn_scale: wgpu::Buffer,
bn_shift: wgpu::Buffer,
pw2: GpuLinear,
n_ff2: GpuNorm,
ff2_1: GpuLinear,
ff2_2: GpuLinear, n_out: GpuNorm,
}
pub struct ParakeetGpu {
cpu: Parakeet,
gemm2: Vec<wgpu::ComputePipeline>,
gemm3: Vec<wgpu::ComputePipeline>,
ln: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
cp: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
glu: wgpu::ComputePipeline,
conv_dw: wgpu::ComputePipeline,
blocks: Vec<GpuBlock>,
d: usize,
heads: usize,
hd: usize,
conv_k: usize,
ff: usize,
max_t: usize,
}
impl ParakeetGpu {
pub fn new(ctx: &GpuCtx, cpu: Parakeet, max_t: usize) -> Result<Self> {
let d = cpu.d_model;
let (heads, hd) = (cpu.layers[0].attn.heads, cpu.layers[0].attn.head_dim);
let conv_k = cpu.layers[0].conv.k;
let ff = cpu.layers[0].ff1.linear1.n;
anyhow::ensure!(hd <= 128, "rel-pos attention supports head_dim ≤ 128");
let norm = |ctx: &GpuCtx, ln: &crate::parakeet::LayerNorm| GpuNorm {
w: ctx.storage(&ln.w),
b: ctx.storage(&ln.b),
};
let lin = |ctx: &GpuCtx, l: &crate::parakeet::Linear, scale: f32| {
GpuLinear::new_v3(ctx, &l.w, &l.b, l.n, l.k, scale)
};
let blocks = cpu
.layers
.iter()
.map(|b| GpuBlock {
n_ff1: norm(ctx, &b.norm_ff1),
ff1_1: lin(ctx, &b.ff1.linear1, 1.0),
ff1_2: lin(ctx, &b.ff1.linear2, 0.5),
n_attn: norm(ctx, &b.norm_attn),
lq: lin(ctx, &b.attn.linear_q, 1.0),
lk: lin(ctx, &b.attn.linear_k, 1.0),
lv: lin(ctx, &b.attn.linear_v, 1.0),
lpos: lin(ctx, &b.attn.linear_pos, 1.0),
lout: lin(ctx, &b.attn.linear_out, 1.0),
ub: ctx.storage(&b.attn.pos_bias_u),
vbb: ctx.storage(&b.attn.pos_bias_v),
n_conv: norm(ctx, &b.norm_conv),
pw1: lin(ctx, &b.conv.pointwise1, 1.0),
dw: ctx.storage(&b.conv.dw),
bn_scale: ctx.storage(&b.conv.bn_scale),
bn_shift: ctx.storage(&b.conv.bn_shift),
pw2: lin(ctx, &b.conv.pointwise2, 1.0),
n_ff2: norm(ctx, &b.norm_ff2),
ff2_1: lin(ctx, &b.ff2.linear1, 1.0),
ff2_2: lin(ctx, &b.ff2.linear2, 0.5),
n_out: norm(ctx, &b.norm_out),
})
.collect();
Ok(Self {
gemm2: GEMM2_TILES
.iter()
.map(|&(bm, bn)| pipeline(ctx, "pk_gemm2", &enc_gemm2_src(false, bm, bn)))
.collect(),
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "pk_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
ln: pipeline(ctx, "pk_ln", LN_SRC),
add: pipeline(ctx, "pk_add", ADD_SRC),
cp: pipeline(ctx, "pk_cp", CP_SRC),
attn: pipeline(ctx, "pk_attn", REL_POS_ATTN),
glu: pipeline(ctx, "pk_glu", CONV_GLU),
conv_dw: pipeline(ctx, "pk_conv_dw", CONV_DW_BN_SILU),
cpu,
blocks,
d,
heads,
hd,
conv_k,
ff,
max_t,
})
}
pub fn cpu(&self) -> &Parakeet {
&self.cpu
}
pub fn transcribe(
&self,
ctx: &GpuCtx,
samples_16k: &[f32],
) -> Result<(Vec<crate::parakeet::TokenEvent>, String)> {
let (mel, t) = self.cpu.logmel(samples_16k);
let (features, t_enc) = self.forward_encoder(ctx, &mel, t)?;
let (tokens, _) = self.cpu.decode_greedy(&features, t_enc);
let text: String = tokens.iter().map(|t| t.text.as_str()).collect();
Ok((tokens, text.trim().to_string()))
}
pub fn forward_encoder(
&self,
ctx: &GpuCtx,
mel: &[f32],
t: usize,
) -> Result<(Vec<f32>, usize)> {
let (x0, pe, t_enc) = self.cpu.subsample_and_posemb(mel, t);
anyhow::ensure!(
t_enc <= self.max_t,
"t_enc {t_enc} exceeds max_t {}",
self.max_t
);
let (d, ff, plen) = (self.d, self.ff, 2 * t_enc - 1);
let xb = ctx.storage(&x0); let peb = ctx.storage(&pe);
let nb = ctx.empty(t_enc * d); let hb = ctx.empty(t_enc * ff); let sb = ctx.empty(t_enc * d); let qb = ctx.empty(t_enc * d);
let kb = ctx.empty(t_enc * d);
let vb = ctx.empty(t_enc * d);
let pb = ctx.empty(plen * d);
let g2b = ctx.empty(t_enc * 2 * d); let glb = ctx.empty(t_enc * d);
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, $m:expr, $act:expr) => {{
let lw: &GpuLinear = $lw;
let m: usize = $m;
let flags = 1u32 | (act_code($act) << 8); let meta = uni(ctx, bytemuck::cast_slice(&[m as u32, lw.n, lw.k, flags]));
let (pl, bm, bn) = if lw.v3 {
let tile = gemm3_tile(m, lw.n as usize);
(&self.gemm3[gemm3_tier(tile)], tile.0, tile.1)
} else {
let tile = gemm2_tile(m, 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),
(m 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_enc as u32, 1));
keep.push(meta);
}};
}
macro_rules! add {
($dst:expr, $src:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[(t_enc * d) as u32, 0u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.add, &[$dst, $src], &meta);
passes.push((&self.add, bg, ((t_enc * d) as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
for b in &self.blocks {
ln!(&xb, b.n_ff1, &nb);
gemm!(&nb, &b.ff1_1, &hb, t_enc, Some(Act::Silu));
gemm!(&hb, &b.ff1_2, &sb, t_enc, None);
add!(&xb, &sb);
ln!(&xb, b.n_attn, &nb);
gemm!(&nb, &b.lq, &qb, t_enc, None);
gemm!(&nb, &b.lk, &kb, t_enc, None);
gemm!(&nb, &b.lv, &vb, t_enc, None);
gemm!(&peb, &b.lpos, &pb, plen, None);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
t_enc as u32,
self.heads as u32,
self.hd as u32,
plen as u32,
]),
);
let bg = make_bg(
ctx,
&self.attn,
&[&qb, &kb, &vb, &pb, &b.ub, &b.vbb, &nb],
&meta,
);
passes.push((&self.attn, bg, t_enc as u32, self.heads as u32));
keep.push(meta);
}
gemm!(&nb, &b.lout, &sb, t_enc, None);
add!(&xb, &sb);
ln!(&xb, b.n_conv, &nb);
gemm!(&nb, &b.pw1, &g2b, t_enc, None); {
let meta = uni(
ctx,
bytemuck::cast_slice(&[t_enc as u32, d as u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.glu, &[&g2b, &glb], &meta);
passes.push((&self.glu, bg, t_enc as u32, 1));
keep.push(meta);
}
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
t_enc as u32,
d as u32,
self.conv_k as u32,
((self.conv_k - 1) / 2) as u32,
]),
);
let bg = make_bg(
ctx,
&self.conv_dw,
&[&glb, &b.dw, &b.bn_scale, &b.bn_shift, &nb],
&meta,
);
passes.push((&self.conv_dw, bg, t_enc as u32, 1));
keep.push(meta);
}
gemm!(&nb, &b.pw2, &sb, t_enc, None);
add!(&xb, &sb);
ln!(&xb, b.n_ff2, &nb);
gemm!(&nb, &b.ff2_1, &hb, t_enc, Some(Act::Silu));
gemm!(&hb, &b.ff2_2, &sb, t_enc, None);
add!(&xb, &sb);
ln!(&xb, b.n_out, &sb);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[(t_enc * d) as u32, 0u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.cp, &[&xb, &sb], &meta);
passes.push((&self.cp, bg, ((t_enc * d) as u32).div_ceil(256), 1));
keep.push(meta);
}
}
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("parakeet_enc"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("parakeet_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(&xb, t_enc * d)?;
drop(keep);
Ok((out, t_enc))
}
}
const REL_POS_ATTN: &str = r#"
struct Meta { t: u32, heads: u32, hd: u32, plen: 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> p: array<f32>;
@group(0) @binding(4) var<storage, read> ub: array<f32>; // pos_bias_u [heads*hd]
@group(0) @binding(5) var<storage, read> vbb: array<f32>; // pos_bias_v [heads*hd]
@group(0) @binding(6) var<storage, read_write> outp: array<f32>;
@group(0) @binding(7) var<uniform> mt: Meta;
var<workgroup> sh: array<f32, 128>;
@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;
let hb = head * hd;
let on = tid < hd;
let qi_u = select(0.0, q[qbase + tid] + ub[hb + tid], on);
let qi_v = select(0.0, q[qbase + tid] + vbb[hb + tid], on);
var acc = 0.0;
var run_max = -1e30;
var run_den = 0.0;
for (var j = 0u; j < t; j = j + 1u) {
let m = (t - 1u) - i + j; // rel-shift index into p; always in [0, plen)
let kbase = j * d + head * hd;
let pbase = m * d + head * hd;
var partial = 0.0;
if (on) { partial = qi_u * k[kbase + tid] + qi_v * p[pbase + tid]; }
sh[tid] = partial;
workgroupBarrier();
for (var s = 64u; s > 0u; s = s >> 1u) {
if (tid < s) { sh[tid] = sh[tid] + sh[tid + s]; }
workgroupBarrier();
}
let score = sh[0] * scale;
workgroupBarrier(); // sh[0] read before the next sh[tid] write
let nm = max(run_max, score);
let corr = exp(run_max - nm);
let pj = exp(score - nm);
run_den = run_den * corr + pj;
run_max = nm;
if (on) { acc = acc * corr + pj * v[kbase + tid]; }
}
if (on) { outp[qbase + tid] = acc / run_den; }
}
"#;
const CONV_GLU: &str = r#"
struct Meta { t: u32, d: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> gated: array<f32>;
@group(0) @binding(1) var<storage, read_write> outp: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) tid: u32) {
let row = wg.x;
let d = mt.d;
let base = row * 2u * d;
for (var c = tid; c < d; c = c + 256u) {
let a = gated[base + c];
let g = gated[base + d + c];
outp[row * d + c] = a * (1.0 / (1.0 + exp(-g))); // GLU: a · sigmoid(g)
}
}
"#;
const CONV_DW_BN_SILU: &str = r#"
struct Meta { t: u32, d: u32, k: u32, pad: u32 }
@group(0) @binding(0) var<storage, read> glu: array<f32>;
@group(0) @binding(1) var<storage, read> dw: array<f32>; // [k*d] tap-major
@group(0) @binding(2) var<storage, read> scale: array<f32>; // [d]
@group(0) @binding(3) var<storage, read> shift: array<f32>; // [d]
@group(0) @binding(4) var<storage, read_write> outp: array<f32>;
@group(0) @binding(5) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) tid: u32) {
let i = wg.x;
let d = mt.d;
let t = mt.t;
for (var c = tid; c < d; c = c + 256u) {
var acc = 0.0;
for (var tap = 0u; tap < mt.k; tap = tap + 1u) {
let j = i32(i) + i32(tap) - i32(mt.pad);
if (j >= 0 && j < i32(t)) {
acc = acc + dw[tap * d + c] * glu[u32(j) * d + c];
}
}
let bn = acc * scale[c] + shift[c];
outp[i * d + c] = bn / (1.0 + exp(-bn)); // SiLU
}
}
"#;
#[allow(clippy::too_many_arguments)]
pub fn conv_glu_dw_bn_silu_gpu(
ctx: &GpuCtx,
gated: &[f32],
dw: &[f32],
bn_scale: &[f32],
bn_shift: &[f32],
t: usize,
d: usize,
k: usize,
) -> Result<Vec<f32>> {
let gatedb = ctx.storage(gated);
let glub = ctx.empty(t * d);
let dwb = ctx.storage(dw);
let scaleb = ctx.storage(bn_scale);
let shiftb = ctx.storage(bn_shift);
let outb = ctx.empty(t * d);
let m_glu = uni(ctx, bytemuck::cast_slice(&[t as u32, d as u32, 0u32, 0u32]));
let m_dw = uni(
ctx,
bytemuck::cast_slice(&[t as u32, d as u32, k as u32, ((k - 1) / 2) as u32]),
);
let pl_glu = pipeline(ctx, "parakeet_conv_glu", CONV_GLU);
let pl_dw = pipeline(ctx, "parakeet_conv_dw", CONV_DW_BN_SILU);
let bg_glu = make_bg(ctx, &pl_glu, &[&gatedb, &glub], &m_glu);
let bg_dw = make_bg(ctx, &pl_dw, &[&glub, &dwb, &scaleb, &shiftb, &outb], &m_dw);
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("parakeet_conv"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("parakeet_conv"),
timestamp_writes: None,
});
cpass.set_pipeline(&pl_glu);
cpass.set_bind_group(0, &bg_glu, &[]);
cpass.dispatch_workgroups(t as u32, 1, 1);
cpass.set_pipeline(&pl_dw);
cpass.set_bind_group(0, &bg_dw, &[]);
cpass.dispatch_workgroups(t as u32, 1, 1);
}
ctx.queue.submit(Some(enc.finish()));
ctx.read(&outb, t * d)
}
#[allow(clippy::too_many_arguments)]
pub fn relpos_attention_gpu(
ctx: &GpuCtx,
q: &[f32],
k: &[f32],
v: &[f32],
p: &[f32],
pos_bias_u: &[f32],
pos_bias_v: &[f32],
heads: usize,
head_dim: usize,
t: usize,
) -> Result<Vec<f32>> {
anyhow::ensure!(head_dim <= 128, "rel-pos attention supports head_dim ≤ 128");
let d = heads * head_dim;
let qb = ctx.storage(q);
let kb = ctx.storage(k);
let vb = ctx.storage(v);
let pb = ctx.storage(p);
let ubb = ctx.storage(pos_bias_u);
let vbbb = ctx.storage(pos_bias_v);
let outb = ctx.empty(t * d);
let meta = uni(
ctx,
bytemuck::cast_slice(&[t as u32, heads as u32, head_dim as u32, (2 * t - 1) as u32]),
);
let pl = pipeline(ctx, "parakeet_relpos_attn", REL_POS_ATTN);
let bg = make_bg(ctx, &pl, &[&qb, &kb, &vb, &pb, &ubb, &vbbb, &outb], &meta);
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("parakeet_attn"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("parakeet_attn"),
timestamp_writes: None,
});
cpass.set_pipeline(&pl);
cpass.set_bind_group(0, &bg, &[]);
cpass.dispatch_workgroups(t as u32, heads as u32, 1);
}
ctx.queue.submit(Some(enc.finish()));
ctx.read(&outb, t * d)
}
const PD_EMBED: &str = r#"
struct Meta { tok: u32, ph: u32, start: u32, pad: u32 }
@group(0) @binding(0) var<storage, read> embed: array<f32>;
@group(0) @binding(1) var<storage, read_write> x: 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.ph) { return; }
x[g.x] = select(embed[mt.tok * mt.ph + g.x], 0.0, mt.start != 0u);
}
"#;
const PD_LSTM_GATE: &str = r#"
struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> gates: array<f32>; // [4H] (i|f|g|o)
@group(0) @binding(1) var<storage, read_write> c: array<f32>; // [H]
@group(0) @binding(2) var<storage, read_write> h: array<f32>; // [H]
@group(0) @binding(3) var<uniform> mt: Meta;
fn sig(x: f32) -> f32 { return 1.0 / (1.0 + exp(-x)); }
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
let j = g.x;
if (j >= mt.h) { return; }
let hn = mt.h;
let i = sig(gates[j]);
let f = sig(gates[hn + j]);
let gg = tanh(clamp(gates[2u * hn + j], -20.0, 20.0));
let o = sig(gates[3u * hn + j]);
let cc = f * c[j] + i * gg;
c[j] = cc;
h[j] = o * tanh(clamp(cc, -20.0, 20.0));
}
"#;
const PD_JOINT_HIDDEN: &str = r#"
struct Meta { jh: u32, step: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> enc: array<f32>; // [t_enc, jh]
@group(0) @binding(1) var<storage, read> pred: array<f32>; // [jh]
@group(0) @binding(2) var<storage, read_write> hidden: array<f32>; // [jh]
@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.jh) { return; }
hidden[g.x] = max(enc[mt.step * mt.jh + g.x] + pred[g.x], 0.0);
}
"#;
const PD_ARGMAX: &str = r#"
struct Meta { offset: u32, len: u32, out_idx: u32, pad: u32 }
@group(0) @binding(0) var<storage, read> logits: array<f32>;
@group(0) @binding(1) var<storage, read_write> out: array<u32>;
@group(0) @binding(2) 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.len; j += 256u) {
let v = logits[mt.offset + j];
if (v > lv) { lv = v; 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) { out[mt.out_idx] = ib[0]; }
}
"#;
struct GpuLstm {
wx: GpuLinear, wh: GpuLinear, }
pub struct ParakeetDecoderGpu {
gemm3: Vec<wgpu::ComputePipeline>,
add: wgpu::ComputePipeline,
embed_pl: wgpu::ComputePipeline,
lstm_gate: wgpu::ComputePipeline,
joint_hidden: wgpu::ComputePipeline,
argmax: wgpu::ComputePipeline,
embed: wgpu::Buffer,
lstm: Vec<GpuLstm>,
joint_enc: GpuLinear,
joint_pred: GpuLinear,
joint_out: GpuLinear,
ph: usize,
jh: usize,
v1: usize,
nd: usize,
blank: usize,
vocabulary: Vec<String>,
durations: Vec<usize>,
max_symbols: Option<usize>,
}
impl ParakeetDecoderGpu {
pub fn new(ctx: &GpuCtx, cpu: &Parakeet) -> Result<Self> {
let ph = cpu.pred_hidden;
let jh = cpu.joint_enc.n;
let lin = |l: &crate::parakeet::Linear| GpuLinear::new_v3(ctx, &l.w, &l.b, l.n, l.k, 1.0);
let lstm = cpu
.lstm
.iter()
.map(|ly| {
let fourh = 4 * ly.hidden;
GpuLstm {
wx: GpuLinear::new_v3(ctx, &ly.wx, &ly.bias, fourh, ly.input, 1.0),
wh: GpuLinear::new_v3(ctx, &ly.wh, &vec![0.0; fourh], fourh, ly.hidden, 1.0),
}
})
.collect();
Ok(Self {
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "pd_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
add: pipeline(ctx, "pd_add", ADD_SRC),
embed_pl: pipeline(ctx, "pd_embed", PD_EMBED),
lstm_gate: pipeline(ctx, "pd_lstm_gate", PD_LSTM_GATE),
joint_hidden: pipeline(ctx, "pd_joint_hidden", PD_JOINT_HIDDEN),
argmax: pipeline(ctx, "pd_argmax", PD_ARGMAX),
embed: ctx.storage(&cpu.embed),
lstm,
joint_enc: lin(&cpu.joint_enc),
joint_pred: lin(&cpu.joint_pred),
joint_out: lin(&cpu.joint_out),
ph,
jh,
v1: cpu.vocabulary.len() + 1,
nd: cpu.durations.len(),
blank: cpu.vocabulary.len(),
vocabulary: cpu.vocabulary.clone(),
durations: cpu.durations.clone(),
max_symbols: cpu.max_symbols,
})
}
pub fn decode_greedy(
&self,
ctx: &GpuCtx,
features: &[f32],
t_enc: usize,
) -> Result<(
Vec<crate::parakeet::TokenEvent>,
Vec<crate::parakeet::TraceStep>,
)> {
let (ph, jh) = (self.ph, self.jh);
let nl = self.lstm.len();
let b_feat = ctx.storage(features);
let enc_proj = ctx.empty(t_enc * jh);
let xb = ctx.empty(ph);
let gates = ctx.empty(4 * ph);
let whb = ctx.empty(4 * ph);
let h: Vec<wgpu::Buffer> = (0..nl).map(|_| ctx.empty(ph)).collect();
let c: Vec<wgpu::Buffer> = (0..nl).map(|_| ctx.empty(ph)).collect();
let pred = ctx.empty(jh);
let hidden = ctx.empty(jh);
let logits = ctx.empty(self.v1 + self.nd);
let outb = ctx.empty(2);
type Pass<'a> = (&'a wgpu::ComputePipeline, wgpu::BindGroup, u32, u32);
macro_rules! gemm {
($ps:expr, $kp: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);
$ps.push((
pl,
bg,
l.n.div_ceil(tile.1 as u32),
(m as u32).div_ceil(tile.0 as u32),
));
$kp.push(meta);
}};
}
macro_rules! simple {
($ps:expr, $kp:expr, $pl:expr, $bufs:expr, $meta:expr, $groups:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice($meta));
let bg = make_bg(ctx, $pl, $bufs, &meta);
$ps.push(($pl, bg, $groups, 1u32));
$kp.push(meta);
}};
}
let submit = |ps: &[Pass]| {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("pk_dec"),
});
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for (pl, bg, gx, gy) in ps {
p.set_pipeline(pl);
p.set_bind_group(0, bg, &[]);
p.dispatch_workgroups(*gx, *gy, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
};
{
let mut ps: Vec<Pass> = Vec::new();
let mut kp: Vec<wgpu::Buffer> = Vec::new();
gemm!(ps, kp, &b_feat, &self.joint_enc, &enc_proj, t_enc, None);
submit(&ps);
}
let phg = (ph as u32).div_ceil(256);
let mut tokens = Vec::new();
let mut trace = Vec::new();
let mut last_token: Option<usize> = None;
let mut lstm_for: Option<Option<usize>> = None;
let mut step = 0usize;
let mut new_symbols = 0usize;
let jhg = (jh as u32).div_ceil(256);
while step < t_enc {
let mut ps: Vec<Pass> = Vec::new();
let mut kp: Vec<wgpu::Buffer> = Vec::new();
if lstm_for != Some(last_token) {
let (tok, start) = match last_token {
Some(t) => (t as u32, 0u32),
None => (0u32, 1u32),
};
simple!(
ps,
kp,
&self.embed_pl,
&[&self.embed, &xb],
&[tok, ph as u32, start, 0u32],
phg
);
for (li, ly) in self.lstm.iter().enumerate() {
let x_in = if li == 0 { &xb } else { &h[li - 1] };
gemm!(ps, kp, x_in, &ly.wx, &gates, 1, None); gemm!(ps, kp, &h[li], &ly.wh, &whb, 1, None); simple!(
ps,
kp,
&self.add,
&[&gates, &whb],
&[(4 * ph) as u32, 0u32, 0u32, 0u32],
((4 * ph) as u32).div_ceil(256)
);
simple!(
ps,
kp,
&self.lstm_gate,
&[&gates, &c[li], &h[li]],
&[ph as u32, 0u32, 0u32, 0u32],
phg
);
}
gemm!(ps, kp, &h[nl - 1], &self.joint_pred, &pred, 1, None);
lstm_for = Some(last_token);
}
simple!(
ps,
kp,
&self.joint_hidden,
&[&enc_proj, &pred, &hidden],
&[jh as u32, step as u32, 0u32, 0u32],
jhg
);
gemm!(ps, kp, &hidden, &self.joint_out, &logits, 1, None);
simple!(
ps,
kp,
&self.argmax,
&[&logits, &outb],
&[0u32, self.v1 as u32, 0u32, 0u32],
1
);
simple!(
ps,
kp,
&self.argmax,
&[&logits, &outb],
&[self.v1 as u32, self.nd as u32, 1u32, 0u32],
1
);
submit(&ps);
drop(kp);
let got = ctx.read(&outb, 2)?;
let token = got[0].to_bits() as usize;
let duration = self.durations[got[1].to_bits() as usize];
trace.push(crate::parakeet::TraceStep {
step,
token,
duration,
});
if token != self.blank {
tokens.push(crate::parakeet::TokenEvent {
id: token,
step,
duration_frames: duration,
text: self.vocabulary[token].replace('▁', " "),
});
last_token = Some(token);
}
step += duration;
new_symbols += 1;
if duration != 0 {
new_symbols = 0;
} else if let Some(max) = self.max_symbols
&& max <= new_symbols
{
step += 1;
new_symbols = 0;
}
}
Ok((tokens, trace))
}
}