use std::path::Path;
use anyhow::Result;
use crate::forward::{make_bg, pipeline, uni};
use crate::qwen3tts::{chatterbox_st, rope_1d_tables, Qwen3TtsConfig};
use crate::whisper_gpu::ADD_SRC;
use crate::GpuCtx;
const MATVEC_SRC: &str = r#"
struct Meta { m: u32, n: u32, k: u32, p: u32 }
@group(0) @binding(0) var<storage, read> x: array<vec4<f32>>; // [m, k/4]
@group(0) @binding(1) var<storage, read> w: array<vec4<f32>>; // [n, k/4]
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) tid: u32) {
let j = wg.x;
if (j >= mt.n) { return; }
let k4 = mt.k >> 2u; // k is always a multiple of 4 (h/kv-width/inter all are)
let wbase = j * k4;
for (var r = 0u; r < mt.m; r = r + 1u) {
let xbase = r * k4;
var acc = vec4<f32>(0.0);
for (var t = tid; t < k4; t = t + 256u) { acc = acc + x[xbase + t] * w[wbase + t]; }
red[tid] = acc.x + acc.y + acc.z + acc.w;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (tid < s) { red[tid] = red[tid] + red[tid + s]; } workgroupBarrier(); }
if (tid == 0u) { y[r * mt.n + j] = red[0]; }
workgroupBarrier();
}
}
"#;
const RMSNORM_SRC: &str = r#"
struct Meta { h: u32, eps_bits: 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_write> outp: array<f32>;
@group(0) @binding(3) 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 ss = 0.0;
for (var i = t; i < mt.h; i = i + 256u) { let v = x[base + i]; ss = ss + v * v; }
sh[t] = ss;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let eps = bitcast<f32>(mt.eps_bits);
let inv = 1.0 / sqrt(sh[0] / f32(mt.h) + eps);
for (var i = t; i < mt.h; i = i + 256u) { outp[base + i] = x[base + i] * inv * w[i]; }
}
"#;
const QK_NORM_SRC: &str = r#"
struct Meta { hd: u32, eps_bits: u32, nheads: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) 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) t: u32) {
let base = wg.x * mt.hd;
var v = 0.0;
if (t < mt.hd) { v = x[base + t]; }
sh[t] = v * v;
workgroupBarrier();
for (var s = 64u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let eps = bitcast<f32>(mt.eps_bits);
let inv = 1.0 / sqrt(sh[0] / f32(mt.hd) + eps);
if (t < mt.hd) { x[base + t] = v * inv * w[t]; }
}
"#;
const ROPE_SRC: &str = r#"
struct Meta { hd: u32, nheads: u32, qbase: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<storage, read> cos: array<f32>;
@group(0) @binding(2) var<storage, read> sin: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
var<workgroup> xs: array<f32, 128>;
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let hd = mt.hd;
let row = wg.x / mt.nheads;
let pos = mt.qbase + row;
let base = wg.x * hd;
let half = hd / 2u;
if (t < hd) { xs[t] = x[base + t]; }
workgroupBarrier();
if (t < hd) {
var rot = 0.0;
if (t < half) { rot = -xs[t + half]; } else { rot = xs[t - half]; }
let c = cos[pos * hd + t];
let s = sin[pos * hd + t];
x[base + t] = xs[t] * c + rot * s;
}
}
"#;
const GQA_ATTN_SRC: &str = r#"
struct Meta { nq: u32, qheads: u32, hd: u32, kvheads: u32, lk: u32, qbase: u32, group: u32, p: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> kc: array<f32>;
@group(0) @binding(2) var<storage, read> vc: array<f32>;
@group(0) @binding(3) var<storage, read_write> outp: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> q_sh: array<f32, 128>;
var<workgroup> red: array<f32, 128>;
var<workgroup> sc: array<f32, 4096>;
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let s = wg.x;
let head = wg.y;
let hd = mt.hd;
let qw = mt.qheads * hd;
let kvw = mt.kvheads * hd;
let kvh = head / mt.group;
let scale = 1.0 / sqrt(f32(hd));
let limit = mt.qbase + s + 1u; // causal: cache row j valid iff j < limit
let qptr = s * qw + head * hd;
if (t < hd) { q_sh[t] = q[qptr + t]; }
workgroupBarrier();
// scores sc[0..limit] = (q · k_j) * scale, reduced over hd
for (var j = 0u; j < limit; j = j + 1u) {
var p = 0.0;
if (t < hd) { p = q_sh[t] * kc[j * kvw + kvh * hd + t]; }
red[t] = p;
workgroupBarrier();
for (var st = 64u; st > 0u; st = st >> 1u) { if (t < st) { red[t] = red[t] + red[t + st]; } workgroupBarrier(); }
if (t == 0u) { sc[j] = red[0] * scale; }
workgroupBarrier();
}
// softmax over sc[0..limit] on a single thread — CPU's exact ascending-j order
if (t == 0u) {
var mx = sc[0];
for (var j = 1u; j < limit; j = j + 1u) { mx = max(mx, sc[j]); }
var den = 0.0;
for (var j = 0u; j < limit; j = j + 1u) { let e = exp(sc[j] - mx); sc[j] = e; den = den + e; }
for (var j = 0u; j < limit; j = j + 1u) { sc[j] = sc[j] / den; }
}
workgroupBarrier();
// out[t] = Σ_j p_j · v_j[t] — sequential over j, matching the CPU accumulation order
if (t < hd) {
var acc = 0.0;
for (var j = 0u; j < limit; j = j + 1u) { acc = acc + sc[j] * vc[j * kvw + kvh * hd + t]; }
outp[qptr + t] = acc;
}
}
"#;
const SWIGLU_SRC: &str = r#"
struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
@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> outp: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let i = gid.x;
if (i < mt.n) { let g = gate[i]; outp[i] = (g / (1.0 + exp(-g))) * up[i]; }
}
"#;
const APPEND_SRC: &str = r#"
struct Meta { n: u32, base: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> cache: 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) { cache[mt.base + gid.x] = src[gid.x]; }
}
"#;
const MAX_KV: usize = 4096;
struct Mat {
w: wgpu::Buffer,
n: u32,
k: u32,
}
struct GpuLayer {
input_ln: wgpu::Buffer,
q: Mat,
k: Mat,
v: Mat,
o: Mat,
q_norm: wgpu::Buffer,
k_norm: wgpu::Buffer,
post_ln: wgpu::Buffer,
gate: Mat,
up: Mat,
down: Mat,
}
struct StepPlan {
passes: Vec<(wgpu::ComputePipeline, wgpu::BindGroup, u32, u32)>,
rope_q: wgpu::Buffer,
rope_k: wgpu::Buffer,
attn: wgpu::Buffer,
append_k: wgpu::Buffer,
append_v: wgpu::Buffer,
_keep: Vec<wgpu::Buffer>,
}
pub struct TalkerGpuState {
kc: Vec<wgpu::Buffer>,
vc: Vec<wgpu::Buffer>,
len: usize,
plan: Option<StepPlan>,
}
impl TalkerGpuState {
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
}
#[derive(Clone, Copy)]
struct Dims {
h: usize,
layers: usize,
heads: usize,
kv_heads: usize,
hd: usize,
inter: usize,
eps: f32,
theta: f32,
}
pub struct TalkerGpu {
d: Dims,
layers: Vec<GpuLayer>,
final_norm: wgpu::Buffer,
cos: wgpu::Buffer,
sin: wgpu::Buffer,
max_pos: usize,
max_m: usize,
matvec: wgpu::ComputePipeline,
rmsnorm: wgpu::ComputePipeline,
qk_norm: wgpu::ComputePipeline,
rope: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
swiglu: wgpu::ComputePipeline,
append: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
nb: wgpu::Buffer,
qb: wgpu::Buffer,
kb: wgpu::Buffer,
vb: wgpu::Buffer,
attn_out: wgpu::Buffer,
ob: wgpu::Buffer,
gateb: wgpu::Buffer,
upb: wgpu::Buffer,
actb: wgpu::Buffer,
downb: wgpu::Buffer,
xb: wgpu::Buffer,
outb: wgpu::Buffer,
}
impl TalkerGpu {
pub fn load(
ctx: &GpuCtx,
dir: &Path,
config: &Qwen3TtsConfig,
max_pos: usize,
max_m: usize,
) -> Result<Self> {
anyhow::ensure!(
max_pos <= MAX_KV,
"max_pos {max_pos} exceeds MAX_KV {MAX_KV}"
);
let t = &config.top.talker_config;
anyhow::ensure!(t.head_dim <= 128, "head_dim {} > 128", t.head_dim);
let d = Dims {
h: t.hidden_size,
layers: t.num_hidden_layers,
heads: t.num_attention_heads,
kv_heads: t.num_key_value_heads,
hd: t.head_dim,
inter: t.intermediate_size,
eps: t.rms_norm_eps,
theta: t.rope_theta,
};
let (h, qw, kvw) = (d.h, d.heads * d.hd, d.kv_heads * d.hd);
let st = chatterbox_st::St::open(&dir.join("model.safetensors"))?;
let lin = |name: &str, n: usize, k: usize| -> Result<Mat> {
let w = st.mat(name, n, k)?; Ok(Mat {
w: ctx.storage(&w),
n: n as u32,
k: k as u32,
})
};
let norm = |name: &str| -> Result<wgpu::Buffer> { Ok(ctx.storage(&st.f32(name)?)) };
let mut layers = Vec::with_capacity(d.layers);
for l in 0..d.layers {
let p = format!("talker.model.layers.{l}");
layers.push(GpuLayer {
input_ln: norm(&format!("{p}.input_layernorm.weight"))?,
q: lin(&format!("{p}.self_attn.q_proj.weight"), qw, h)?,
k: lin(&format!("{p}.self_attn.k_proj.weight"), kvw, h)?,
v: lin(&format!("{p}.self_attn.v_proj.weight"), kvw, h)?,
o: lin(&format!("{p}.self_attn.o_proj.weight"), h, qw)?,
q_norm: norm(&format!("{p}.self_attn.q_norm.weight"))?,
k_norm: norm(&format!("{p}.self_attn.k_norm.weight"))?,
post_ln: norm(&format!("{p}.post_attention_layernorm.weight"))?,
gate: lin(&format!("{p}.mlp.gate_proj.weight"), d.inter, h)?,
up: lin(&format!("{p}.mlp.up_proj.weight"), d.inter, h)?,
down: lin(&format!("{p}.mlp.down_proj.weight"), h, d.inter)?,
});
}
let final_norm = norm("talker.model.norm.weight")?;
let mut cos = vec![0f32; max_pos * d.hd];
let mut sin = vec![0f32; max_pos * d.hd];
for pos in 0..max_pos {
let (c, s) = rope_1d_tables(d.hd, d.theta, pos as u32);
cos[pos * d.hd..(pos + 1) * d.hd].copy_from_slice(&c);
sin[pos * d.hd..(pos + 1) * d.hd].copy_from_slice(&s);
}
Ok(Self {
d,
layers,
final_norm,
cos: ctx.storage(&cos),
sin: ctx.storage(&sin),
max_pos,
max_m,
matvec: pipeline(ctx, "tts_matvec", MATVEC_SRC),
rmsnorm: pipeline(ctx, "tts_rmsnorm", RMSNORM_SRC),
qk_norm: pipeline(ctx, "tts_qk_norm", QK_NORM_SRC),
rope: pipeline(ctx, "tts_rope", ROPE_SRC),
attn: pipeline(ctx, "tts_attn", GQA_ATTN_SRC),
swiglu: pipeline(ctx, "tts_swiglu", SWIGLU_SRC),
append: pipeline(ctx, "tts_append", APPEND_SRC),
add: pipeline(ctx, "tts_add", ADD_SRC),
nb: ctx.empty(max_m * h),
qb: ctx.empty(max_m * qw),
kb: ctx.empty(max_m * kvw),
vb: ctx.empty(max_m * kvw),
attn_out: ctx.empty(max_m * qw),
ob: ctx.empty(max_m * h),
gateb: ctx.empty(max_m * d.inter),
upb: ctx.empty(max_m * d.inter),
actb: ctx.empty(max_m * d.inter),
downb: ctx.empty(max_m * h),
xb: ctx.empty(max_m * h),
outb: ctx.empty(max_m * h),
})
}
pub fn hidden_size(&self) -> usize {
self.d.h
}
pub fn new_state(&self, ctx: &GpuCtx) -> TalkerGpuState {
let kvw = self.d.kv_heads * self.d.hd;
TalkerGpuState {
kc: (0..self.d.layers)
.map(|_| ctx.empty(self.max_pos * kvw))
.collect(),
vc: (0..self.d.layers)
.map(|_| ctx.empty(self.max_pos * kvw))
.collect(),
len: 0,
plan: None,
}
}
pub fn prefill(
&self,
ctx: &GpuCtx,
state: &mut TalkerGpuState,
embeds: &[f32],
) -> Result<Vec<f32>> {
let seq = embeds.len() / self.d.h;
anyhow::ensure!(embeds.len() == seq * self.d.h, "embeds not a multiple of h");
self.forward_rows(ctx, state, embeds, seq, state.len as u32)
}
pub fn step(
&self,
ctx: &GpuCtx,
state: &mut TalkerGpuState,
embed: &[f32],
pos: u32,
) -> Result<Vec<f32>> {
let d = self.d;
let (h, kvw) = (d.h, d.kv_heads * d.hd);
anyhow::ensure!(embed.len() == h, "embed dim {} != h", embed.len());
anyhow::ensure!(
pos as usize == state.len,
"step pos {pos} != cache len {} (positions must be contiguous)",
state.len
);
let lk = state.len + 1;
anyhow::ensure!(
lk <= self.max_pos,
"cache len {lk} exceeds max_pos {}",
self.max_pos
);
anyhow::ensure!(lk <= MAX_KV, "cache len {lk} exceeds MAX_KV {MAX_KV}");
if state.plan.is_none() {
let p = self.build_step_plan(ctx, state);
state.plan = Some(p);
}
let group = (d.heads / d.kv_heads) as u32;
let base = (state.len * kvw) as u32;
let out = {
let plan = state.plan.as_ref().unwrap();
ctx.queue.write_buffer(
&plan.rope_q,
0,
bytemuck::cast_slice(&[d.hd as u32, d.heads as u32, pos, 0u32]),
);
ctx.queue.write_buffer(
&plan.rope_k,
0,
bytemuck::cast_slice(&[d.hd as u32, d.kv_heads as u32, pos, 0u32]),
);
ctx.queue.write_buffer(
&plan.attn,
0,
bytemuck::cast_slice(&[
1u32,
d.heads as u32,
d.hd as u32,
d.kv_heads as u32,
lk as u32,
pos,
group,
0u32,
]),
);
ctx.queue.write_buffer(
&plan.append_k,
0,
bytemuck::cast_slice(&[kvw as u32, base, 0u32, 0u32]),
);
ctx.queue.write_buffer(
&plan.append_v,
0,
bytemuck::cast_slice(&[kvw as u32, base, 0u32, 0u32]),
);
ctx.queue
.write_buffer(&self.xb, 0, bytemuck::cast_slice(embed));
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("qwen3tts_talker_step"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("qwen3tts_talker_step"),
timestamp_writes: None,
});
for (pl, bg, gx, gy) in &plan.passes {
cpass.set_pipeline(pl);
cpass.set_bind_group(0, bg, &[]);
cpass.dispatch_workgroups(*gx, *gy, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
ctx.read(&self.outb, h)?
};
state.len += 1;
Ok(out)
}
fn build_step_plan(&self, ctx: &GpuCtx, state: &TalkerGpuState) -> StepPlan {
let d = self.d;
let (h, kvw) = (d.h, d.kv_heads * d.hd);
let m = 1usize;
let eps_bits = d.eps.to_bits();
let group = (d.heads / d.kv_heads) as u32;
let rope_q = uni(
ctx,
bytemuck::cast_slice(&[d.hd as u32, d.heads as u32, 0u32, 0u32]),
);
let rope_k = uni(
ctx,
bytemuck::cast_slice(&[d.hd as u32, d.kv_heads as u32, 0u32, 0u32]),
);
let attn_u = uni(
ctx,
bytemuck::cast_slice(&[
1u32,
d.heads as u32,
d.hd as u32,
d.kv_heads as u32,
0u32,
0u32,
group,
0u32,
]),
);
let append_k = uni(ctx, bytemuck::cast_slice(&[kvw as u32, 0u32, 0u32, 0u32]));
let append_v = uni(ctx, bytemuck::cast_slice(&[kvw as u32, 0u32, 0u32, 0u32]));
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) => {{
let lw: &Mat = $lw;
let meta = uni(ctx, bytemuck::cast_slice(&[m as u32, lw.n, lw.k, 0u32]));
let bg = make_bg(ctx, &self.matvec, &[$x, &lw.w, $y], &meta);
passes.push((self.matvec.clone(), bg, lw.n, 1));
keep.push(meta);
}};
}
macro_rules! rmsnorm {
($x:expr, $w:expr, $y:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[h as u32, eps_bits, 0u32, 0u32]));
let bg = make_bg(ctx, &self.rmsnorm, &[$x, $w, $y], &meta);
passes.push((self.rmsnorm.clone(), bg, m as u32, 1));
keep.push(meta);
}};
}
macro_rules! qknorm {
($x:expr, $w:expr, $nheads:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[d.hd as u32, eps_bits, $nheads as u32, 0u32]),
);
let bg = make_bg(ctx, &self.qk_norm, &[$x, $w], &meta);
passes.push((self.qk_norm.clone(), bg, (m * $nheads) as u32, 1));
keep.push(meta);
}};
}
macro_rules! rope {
($x:expr, $uni:expr) => {{
let bg = make_bg(ctx, &self.rope, &[$x, &self.cos, &self.sin], $uni);
passes.push((self.rope.clone(), bg, (m * d.heads) as u32, 1));
}};
}
macro_rules! append {
($src:expr, $cache:expr, $uni:expr) => {{
let bg = make_bg(ctx, &self.append, &[$src, $cache], $uni);
passes.push((self.append.clone(), bg, ((m * kvw) as u32).div_ceil(256), 1));
}};
}
macro_rules! swiglu {
($g:expr, $u:expr, $y:expr, $n:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[$n as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.swiglu, &[$g, $u, $y], &meta);
passes.push((self.swiglu.clone(), bg, ($n as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
macro_rules! add {
($dst:expr, $src:expr, $n:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[$n as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.add, &[$dst, $src], &meta);
passes.push((self.add.clone(), bg, ($n as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
for (l, layer) in self.layers.iter().enumerate() {
rmsnorm!(&self.xb, &layer.input_ln, &self.nb);
gemm!(&self.nb, &layer.q, &self.qb);
gemm!(&self.nb, &layer.k, &self.kb);
gemm!(&self.nb, &layer.v, &self.vb);
qknorm!(&self.qb, &layer.q_norm, d.heads);
qknorm!(&self.kb, &layer.k_norm, d.kv_heads);
rope!(&self.qb, &rope_q);
{
let bg = make_bg(ctx, &self.rope, &[&self.kb, &self.cos, &self.sin], &rope_k);
passes.push((self.rope.clone(), bg, (m * d.kv_heads) as u32, 1));
}
append!(&self.kb, &state.kc[l], &append_k);
append!(&self.vb, &state.vc[l], &append_v);
{
let bg = make_bg(
ctx,
&self.attn,
&[&self.qb, &state.kc[l], &state.vc[l], &self.attn_out],
&attn_u,
);
passes.push((self.attn.clone(), bg, m as u32, d.heads as u32));
}
gemm!(&self.attn_out, &layer.o, &self.ob);
add!(&self.xb, &self.ob, m * h);
rmsnorm!(&self.xb, &layer.post_ln, &self.nb);
gemm!(&self.nb, &layer.gate, &self.gateb);
gemm!(&self.nb, &layer.up, &self.upb);
swiglu!(&self.gateb, &self.upb, &self.actb, m * d.inter);
gemm!(&self.actb, &layer.down, &self.downb);
add!(&self.xb, &self.downb, m * h);
}
rmsnorm!(&self.xb, &self.final_norm, &self.outb);
StepPlan {
passes,
rope_q,
rope_k,
attn: attn_u,
append_k,
append_v,
_keep: keep,
}
}
fn forward_rows(
&self,
ctx: &GpuCtx,
state: &mut TalkerGpuState,
embeds: &[f32],
m: usize,
qbase: u32,
) -> Result<Vec<f32>> {
let d = self.d;
let (h, kvw) = (d.h, d.kv_heads * d.hd);
let lk = state.len + m; anyhow::ensure!(m <= self.max_m, "m {m} exceeds max_m {}", self.max_m);
anyhow::ensure!(
lk <= self.max_pos,
"cache len {lk} exceeds max_pos {}",
self.max_pos
);
anyhow::ensure!(lk <= MAX_KV, "cache len {lk} exceeds MAX_KV {MAX_KV}");
ctx.queue
.write_buffer(&self.xb, 0, bytemuck::cast_slice(embeds));
type Pass<'a> = (&'a wgpu::ComputePipeline, wgpu::BindGroup, u32, u32);
let mut passes: Vec<Pass> = Vec::new();
let mut keep: Vec<wgpu::Buffer> = Vec::new();
let eps_bits = d.eps.to_bits();
macro_rules! gemm {
($x:expr, $lw:expr, $y:expr) => {{
let lw: &Mat = $lw;
let meta = uni(ctx, bytemuck::cast_slice(&[m as u32, lw.n, lw.k, 0u32]));
let bg = make_bg(ctx, &self.matvec, &[$x, &lw.w, $y], &meta);
passes.push((&self.matvec, bg, lw.n, 1));
keep.push(meta);
}};
}
macro_rules! rmsnorm {
($x:expr, $w:expr, $y:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[h as u32, eps_bits, 0u32, 0u32]));
let bg = make_bg(ctx, &self.rmsnorm, &[$x, $w, $y], &meta);
passes.push((&self.rmsnorm, bg, m as u32, 1));
keep.push(meta);
}};
}
macro_rules! qknorm {
($x:expr, $w:expr, $nheads:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[d.hd as u32, eps_bits, $nheads as u32, 0u32]),
);
let bg = make_bg(ctx, &self.qk_norm, &[$x, $w], &meta);
passes.push((&self.qk_norm, bg, (m * $nheads) as u32, 1));
keep.push(meta);
}};
}
macro_rules! rope {
($x:expr, $nheads:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[d.hd as u32, $nheads as u32, qbase, 0u32]),
);
let bg = make_bg(ctx, &self.rope, &[$x, &self.cos, &self.sin], &meta);
passes.push((&self.rope, bg, (m * $nheads) as u32, 1));
keep.push(meta);
}};
}
macro_rules! append {
($src:expr, $cache:expr, $w:expr) => {{
let n = (m * $w) as u32;
let base = (state.len * $w) as u32;
let meta = uni(ctx, bytemuck::cast_slice(&[n, base, 0u32, 0u32]));
let bg = make_bg(ctx, &self.append, &[$src, $cache], &meta);
passes.push((&self.append, bg, n.div_ceil(256), 1));
keep.push(meta);
}};
}
macro_rules! swiglu {
($g:expr, $u:expr, $y:expr, $n:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[$n as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.swiglu, &[$g, $u, $y], &meta);
passes.push((&self.swiglu, bg, ($n as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
macro_rules! add {
($dst:expr, $src:expr, $n:expr) => {{
let meta = uni(ctx, bytemuck::cast_slice(&[$n as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.add, &[$dst, $src], &meta);
passes.push((&self.add, bg, ($n as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
for (l, layer) in self.layers.iter().enumerate() {
rmsnorm!(&self.xb, &layer.input_ln, &self.nb);
gemm!(&self.nb, &layer.q, &self.qb);
gemm!(&self.nb, &layer.k, &self.kb);
gemm!(&self.nb, &layer.v, &self.vb);
qknorm!(&self.qb, &layer.q_norm, d.heads);
qknorm!(&self.kb, &layer.k_norm, d.kv_heads);
rope!(&self.qb, d.heads);
rope!(&self.kb, d.kv_heads);
append!(&self.kb, &state.kc[l], kvw);
append!(&self.vb, &state.vc[l], kvw);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
m as u32,
d.heads as u32,
d.hd as u32,
d.kv_heads as u32,
lk as u32,
qbase,
(d.heads / d.kv_heads) as u32,
0u32,
]),
);
let bg = make_bg(
ctx,
&self.attn,
&[&self.qb, &state.kc[l], &state.vc[l], &self.attn_out],
&meta,
);
passes.push((&self.attn, bg, m as u32, d.heads as u32));
keep.push(meta);
}
gemm!(&self.attn_out, &layer.o, &self.ob);
add!(&self.xb, &self.ob, m * h);
rmsnorm!(&self.xb, &layer.post_ln, &self.nb);
gemm!(&self.nb, &layer.gate, &self.gateb);
gemm!(&self.nb, &layer.up, &self.upb);
swiglu!(&self.gateb, &self.upb, &self.actb, m * d.inter);
gemm!(&self.actb, &layer.down, &self.downb);
add!(&self.xb, &self.downb, m * h);
}
rmsnorm!(&self.xb, &self.final_norm, &self.outb);
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("qwen3tts_talker"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("qwen3tts_talker"),
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(&self.outb, m * h)?;
drop(keep);
state.len += m;
Ok(out)
}
}