use std::path::Path;
use anyhow::{Context, Result};
use rayon::prelude::*;
use crate::weights::LazySt;
fn linear(x: &[f32], w: &[f32], b: &[f32], m: usize, n: usize, k: usize) -> Vec<f32> {
let mut out = vec![0f32; m * n];
out.par_chunks_mut(n).enumerate().for_each(|(i, orow)| {
let xr = &x[i * k..][..k];
for o in 0..n {
let wr = &w[o * k..][..k];
let mut acc = b[o];
for c in 0..k {
acc += xr[c] * wr[c];
}
orow[o] = acc;
}
});
out
}
fn layer_norm(x: &[f32], d: usize, w: &[f32], b: &[f32], eps: f32) -> Vec<f32> {
let mut out = vec![0f32; x.len()];
out.par_chunks_mut(d)
.zip(x.par_chunks(d))
.for_each(|(orow, row)| {
let mean = row.iter().sum::<f32>() / d as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / d as f32;
let inv = 1.0 / (var + eps).sqrt();
for i in 0..d {
orow[i] = (row[i] - mean) * inv * w[i] + b[i];
}
});
out
}
fn gelu(v: f32) -> f32 {
0.5 * v * (1.0 + libm::erff(v * std::f32::consts::FRAC_1_SQRT_2))
}
struct Linear {
w: Vec<f32>,
b: Vec<f32>,
n: usize,
k: usize,
}
impl Linear {
fn load(st: &LazySt, prefix: &str, n: usize, k: usize) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
let b = st.tensor_f32(&format!("{prefix}.bias"))?;
anyhow::ensure!(w.len() == n * k, "{prefix}.weight {} != {n}x{k}", w.len());
Ok(Self { w, b, n, k })
}
fn forward(&self, x: &[f32], m: usize) -> Vec<f32> {
linear(x, &self.w, &self.b, m, self.n, self.k)
}
}
struct Norm {
w: Vec<f32>,
b: Vec<f32>,
}
impl Norm {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
Ok(Self {
w: st.tensor_f32(&format!("{prefix}.weight"))?,
b: st.tensor_f32(&format!("{prefix}.bias"))?,
})
}
fn forward(&self, x: &[f32], d: usize) -> Vec<f32> {
layer_norm(x, d, &self.w, &self.b, 1e-12)
}
}
struct Block {
ln1: Norm,
q: Linear,
k: Linear,
v: Linear,
attn_out: Linear,
ln2: Norm,
fc1: Linear,
fc2: Linear,
}
const FRAME_LEN: usize = 400;
const HOP: usize = 160;
const FFT_LEN: usize = 512;
const MEL_FLOOR: f64 = 1.192_092_955_078_125e-07;
pub struct Ast {
window: Vec<f32>, mel_filters: Vec<f32>, dft_cos: Vec<f64>, dft_sin: Vec<f64>,
n_freq: usize,
mean: f32,
std: f32,
patch_w: Vec<f32>, patch_b: Vec<f32>, cls: Vec<f32>,
dist: Vec<f32>,
pos: Vec<f32>, blocks: Vec<Block>,
ln_post: Norm,
head_ln: Norm,
head: Linear,
hidden: usize,
heads: usize,
hd: usize,
patch: usize,
n_mels: usize,
max_len: usize,
fstride: usize,
tstride: usize,
ft: usize, ff: usize, pub n_labels: usize,
pub id_labels: Vec<String>,
}
impl Ast {
pub fn load(dir: &Path) -> Result<Self> {
let v: serde_json::Value =
serde_json::from_slice(&std::fs::read(dir.join("config.json"))?).context("config")?;
let g = |k: &str| v.get(k).and_then(|x| x.as_u64()).unwrap_or(0) as usize;
let st = LazySt::open(dir)?;
let hidden = g("hidden_size");
let heads = g("num_attention_heads");
let inter = g("intermediate_size");
let patch = g("patch_size");
let n_mels = g("num_mel_bins");
let max_len = g("max_length");
let fstride = g("frequency_stride");
let tstride = g("time_stride");
let ft = (max_len - patch) / tstride + 1;
let ff = (n_mels - patch) / fstride + 1;
let p = "audio_spectrogram_transformer";
let blocks = (0..g("num_hidden_layers"))
.map(|i| {
let b = format!("{p}.layers.{i}");
Ok(Block {
ln1: Norm::load(&st, &format!("{b}.layernorm_before"))?,
q: Linear::load(&st, &format!("{b}.attention.q_proj"), hidden, hidden)?,
k: Linear::load(&st, &format!("{b}.attention.k_proj"), hidden, hidden)?,
v: Linear::load(&st, &format!("{b}.attention.v_proj"), hidden, hidden)?,
attn_out: Linear::load(&st, &format!("{b}.attention.o_proj"), hidden, hidden)?,
ln2: Norm::load(&st, &format!("{b}.layernorm_after"))?,
fc1: Linear::load(&st, &format!("{b}.mlp.fc1"), inter, hidden)?,
fc2: Linear::load(&st, &format!("{b}.mlp.fc2"), hidden, inter)?,
})
})
.collect::<Result<Vec<_>>>()?;
let id_labels = v
.get("id2label")
.and_then(|m| m.as_object())
.map(|m| {
let mut labels = vec![String::new(); m.len()];
for (k, val) in m {
if let (Ok(i), Some(s)) = (k.parse::<usize>(), val.as_str())
&& i < labels.len()
{
labels[i] = s.to_string();
}
}
labels
})
.unwrap_or_default();
let mel_filters = st.tensor_f32("ast_mel_filters")?; let n_freq = FFT_LEN / 2 + 1;
anyhow::ensure!(mel_filters.len() == n_freq * n_mels, "mel_filters shape");
let mut dft_cos = vec![0f64; n_freq * FRAME_LEN];
let mut dft_sin = vec![0f64; n_freq * FRAME_LEN];
for b in 0..n_freq {
for j in 0..FRAME_LEN {
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / FFT_LEN as f64;
dft_cos[b * FRAME_LEN + j] = ang.cos();
dft_sin[b * FRAME_LEN + j] = ang.sin();
}
}
let getf = |k: &str| v.get(k).and_then(|x| x.as_f64()).unwrap_or(0.0) as f32;
Ok(Self {
window: st.tensor_f32("ast_window")?,
mel_filters,
dft_cos,
dft_sin,
n_freq,
mean: getf("fe_mean"),
std: getf("fe_std"),
patch_w: st.tensor_f32(&format!(
"{p}.embeddings.patch_embeddings.projection.weight"
))?,
patch_b: st.tensor_f32(&format!("{p}.embeddings.patch_embeddings.projection.bias"))?,
cls: st.tensor_f32(&format!("{p}.embeddings.cls_token"))?,
dist: st.tensor_f32(&format!("{p}.embeddings.distillation_token"))?,
pos: st.tensor_f32(&format!("{p}.embeddings.position_embeddings"))?,
ln_post: Norm::load(&st, &format!("{p}.layernorm"))?,
head_ln: Norm::load(&st, "classifier.layernorm")?,
head: Linear::load(&st, "classifier.dense", g("num_labels"), hidden)?,
blocks,
hidden,
heads,
hd: hidden / heads,
patch,
n_mels,
max_len,
fstride,
tstride,
ft,
ff,
n_labels: g("num_labels"),
id_labels,
})
}
pub fn logmel(&self, samples: &[f32]) -> Vec<f32> {
let (nm, nf) = (self.n_mels, self.n_freq);
let num_frames = if samples.len() >= FRAME_LEN {
1 + (samples.len() - FRAME_LEN) / HOP
} else {
0
};
let mut fbank = vec![0f64; num_frames * nm];
fbank.par_chunks_mut(nm).enumerate().for_each(|(f, mrow)| {
let mut buf = [0f64; FRAME_LEN];
let base = f * HOP;
for j in 0..FRAME_LEN {
buf[j] = samples[base + j] as f64;
}
let mean = buf.iter().sum::<f64>() / FRAME_LEN as f64;
for x in buf.iter_mut() {
*x -= mean;
}
for i in (1..FRAME_LEN).rev() {
buf[i] -= 0.97 * buf[i - 1];
}
buf[0] *= 1.0 - 0.97;
for j in 0..FRAME_LEN {
buf[j] *= self.window[j] as f64;
}
let mut power = vec![0f64; nf];
for (b, pw) in power.iter_mut().enumerate() {
let (cb, sb) = (
&self.dft_cos[b * FRAME_LEN..],
&self.dft_sin[b * FRAME_LEN..],
);
let (mut re, mut im) = (0f64, 0f64);
for j in 0..FRAME_LEN {
re += buf[j] * cb[j];
im -= buf[j] * sb[j];
}
*pw = re * re + im * im;
}
for m in 0..nm {
let mut acc = 0f64;
for b in 0..nf {
acc += self.mel_filters[b * nm + m] as f64 * power[b];
}
mrow[m] = acc.max(MEL_FLOOR).ln();
}
});
let denom = self.std * 2.0;
let mut out = vec![0f32; self.max_len * nm];
for i in 0..self.max_len * nm {
let v = if i < fbank.len() {
fbank[i] as f32
} else {
0.0
};
out[i] = (v - self.mean) / denom;
}
out
}
fn patch_embed(&self, mel: &[f32]) -> Vec<f32> {
let (h, ps) = (self.hidden, self.patch);
let n_patches = self.ft * self.ff;
let mut out = vec![0f32; n_patches * h];
out.par_chunks_mut(h).enumerate().for_each(|(p, orow)| {
let (fp, tp) = (p / self.ft, p % self.ft); let (f0, t0) = (fp * self.fstride, tp * self.tstride);
for o in 0..h {
let wbase = o * ps * ps;
let mut acc = self.patch_b[o];
for kh in 0..ps {
for kw in 0..ps {
acc += mel[(t0 + kw) * self.n_mels + (f0 + kh)]
* self.patch_w[wbase + kh * ps + kw];
}
}
orow[o] = acc;
}
});
out
}
fn attention(&self, blk: &Block, x: &[f32], t: usize) -> Vec<f32> {
let (h, heads, hd) = (self.hidden, self.heads, self.hd);
let q = blk.q.forward(x, t);
let k = blk.k.forward(x, t);
let v = blk.v.forward(x, t);
let scale = 1.0 / (hd as f32).sqrt();
let ctx: Vec<Vec<f32>> = (0..heads)
.into_par_iter()
.map(|head| {
let mut oh = vec![0f32; t * hd];
let mut srow = vec![0f32; t];
for i in 0..t {
let qi = &q[i * h + head * hd..][..hd];
for j in 0..t {
let kj = &k[j * h + head * hd..][..hd];
srow[j] = qi.iter().zip(kj).map(|(a, b)| a * b).sum::<f32>() * scale;
}
let max = srow.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for s in srow.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
let inv = 1.0 / sum;
let orow = &mut oh[i * hd..][..hd];
for j in 0..t {
let w = srow[j] * inv;
let vj = &v[j * h + head * hd..][..hd];
for c in 0..hd {
orow[c] += w * vj[c];
}
}
}
oh
})
.collect();
let mut merged = vec![0f32; t * h];
for (head, oh) in ctx.iter().enumerate() {
for i in 0..t {
merged[i * h + head * hd..][..hd].copy_from_slice(&oh[i * hd..][..hd]);
}
}
blk.attn_out.forward(&merged, t)
}
pub fn stem(&self, mel: &[f32]) -> (Vec<f32>, usize) {
let h = self.hidden;
let patches = self.patch_embed(mel);
let t = self.ft * self.ff + 2;
let mut x = vec![0f32; t * h];
x[..h].copy_from_slice(&self.cls);
x[h..2 * h].copy_from_slice(&self.dist);
x[2 * h..].copy_from_slice(&patches);
for (xi, pi) in x.iter_mut().zip(&self.pos) {
*xi += pi;
}
(x, t)
}
pub fn head_logits(&self, enc: &[f32]) -> Vec<f32> {
let h = self.hidden;
let mut pooled = vec![0f32; h];
for i in 0..h {
pooled[i] = (enc[i] + enc[h + i]) * 0.5;
}
let normed = self.head_ln.forward(&pooled, h);
self.head.forward(&normed, 1)
}
pub fn encode(&self, mel: &[f32]) -> Vec<f32> {
let h = self.hidden;
let (mut x, t) = self.stem(mel);
for blk in &self.blocks {
let normed = blk.ln1.forward(&x, h);
let a = self.attention(blk, &normed, t);
for (xi, ai) in x.iter_mut().zip(a) {
*xi += ai;
}
let normed = blk.ln2.forward(&x, h);
let mut mid = blk.fc1.forward(&normed, t);
for m in mid.iter_mut() {
*m = gelu(*m);
}
let f = blk.fc2.forward(&mid, t);
for (xi, fi) in x.iter_mut().zip(f) {
*xi += fi;
}
}
self.ln_post.forward(&x, h)
}
pub fn classify(&self, mel: &[f32]) -> Vec<f32> {
self.head_logits(&self.encode(mel))
}
pub(crate) fn dims(&self) -> (usize, usize, usize, usize) {
(self.hidden, self.heads, self.hd, self.blocks[0].fc1.n)
}
pub(crate) fn gpu_blocks(&self) -> &[Block] {
&self.blocks
}
pub(crate) fn gpu_ln_post(&self) -> &Norm {
&self.ln_post
}
}
use crate::GpuCtx;
use crate::encoder::{GEMM3_TILES, act_code, enc_gemm3_src, gemm3_tier, gemm3_tile};
use crate::encoder_weights::Act;
use crate::forward::{make_bg, pipeline, uni};
use crate::whisper_gpu::{ADD_SRC, GpuLinear, GpuNorm, LN_SRC, VANILLA_ATTN};
struct AstGpuBlock {
ln1: GpuNorm,
q: GpuLinear,
k: GpuLinear,
v: GpuLinear,
o: GpuLinear,
ln2: GpuNorm,
fc1: GpuLinear,
fc2: GpuLinear,
}
pub struct AstGpu {
cpu: Ast,
gemm3: Vec<wgpu::ComputePipeline>,
ln: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
blocks: Vec<AstGpuBlock>,
ln_post: GpuNorm,
hidden: usize,
heads: usize,
hd: usize,
ff: usize,
}
impl AstGpu {
pub fn new(ctx: &GpuCtx, cpu: Ast) -> Result<Self> {
let (hidden, heads, hd, ff) = cpu.dims();
anyhow::ensure!(hd <= 128, "attention supports head_dim ≤ 128");
let norm = |n: &Norm| GpuNorm {
w: ctx.storage(&n.w),
b: ctx.storage(&n.b),
};
let lin = |l: &Linear| GpuLinear::new_v3(ctx, &l.w, &l.b, l.n, l.k);
let blocks = cpu
.gpu_blocks()
.iter()
.map(|b| AstGpuBlock {
ln1: norm(&b.ln1),
q: lin(&b.q),
k: lin(&b.k),
v: lin(&b.v),
o: lin(&b.attn_out),
ln2: norm(&b.ln2),
fc1: lin(&b.fc1),
fc2: lin(&b.fc2),
})
.collect();
let ln_post = norm(cpu.gpu_ln_post());
Ok(Self {
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "ast_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
ln: pipeline(ctx, "ast_ln", LN_SRC),
add: pipeline(ctx, "ast_add", ADD_SRC),
attn: pipeline(ctx, "ast_attn", VANILLA_ATTN),
blocks,
ln_post,
hidden,
heads,
hd,
ff,
cpu,
})
}
pub fn cpu(&self) -> &Ast {
&self.cpu
}
pub fn encode(&self, ctx: &GpuCtx, mel: &[f32]) -> Result<Vec<f32>> {
let (stem, t) = self.cpu.stem(mel);
let (d, ff) = (self.hidden, 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 tile = gemm3_tile(t, lw.n as usize);
let pl = &self.gemm3[gemm3_tier(tile)];
let bg = make_bg(ctx, pl, &[$x, &lw.w, &lw.b, $y], &meta);
passes.push((
pl,
bg,
lw.n.div_ceil(tile.1 as u32),
(t as u32).div_ceil(tile.0 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.ln1, &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.o, &sb, None);
add!(&xb, &sb);
ln!(&xb, b.ln2, &nb);
gemm!(&nb, &b.fc1, &hb, Some(Act::GeluErf));
gemm!(&hb, &b.fc2, &sb, None);
add!(&xb, &sb);
}
ln!(&xb, self.ln_post, &nb);
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("ast_enc"),
});
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
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(&nb, t * d)?;
drop(keep);
Ok(out)
}
pub fn classify(&self, ctx: &GpuCtx, mel: &[f32]) -> Result<Vec<f32>> {
Ok(self.cpu.head_logits(&self.encode(ctx, mel)?))
}
}