use anyhow::Result;
use super::audio_encoder::{
AudioEncoderConfig, AudioEncoderWeights, ConformerLayerWeights, ConvStemWeights, POS_EMB_DIM,
relative_pos_emb,
};
use crate::model::weights::MmapWeight;
pub const MAX_AUDIO_TOKENS: usize = 1024;
pub const MAX_AUDIO_HEAD_DIM: usize = 128;
#[cfg(test)]
#[allow(clippy::items_after_test_module)]
mod const_sync_tests {
use super::{MAX_AUDIO_HEAD_DIM, MAX_AUDIO_TOKENS};
#[test]
fn max_audio_tokens_matches_attention_shader_scratch() {
let src = include_str!("../backend/shaders/slang/audio_xl_attention.slang");
let tokens_decl = format!("static const uint MAX_TOKENS = {MAX_AUDIO_TOKENS}u;");
let head_decl = format!("static const uint MAX_HEAD_DIM = {MAX_AUDIO_HEAD_DIM}u;");
assert!(
src.contains(&tokens_decl),
"audio_xl_attention.slang MAX_TOKENS != MAX_AUDIO_TOKENS ({MAX_AUDIO_TOKENS}); \
update the shader's `scores` array size to match"
);
assert!(
src.contains(&head_decl),
"audio_xl_attention.slang MAX_HEAD_DIM != MAX_AUDIO_HEAD_DIM ({MAX_AUDIO_HEAD_DIM}); \
update the shader's `qu`/`qv` array sizes to match"
);
}
}
fn over_capacity_msg(t_out: usize) -> String {
format!(
"post-stem length {t_out} exceeds MAX_AUDIO_TOKENS ({MAX_AUDIO_TOKENS}); \
caller should fall back to the CPU encoder"
)
}
fn unrunnable_stem_msg(n_frames: usize) -> String {
format!("audio encoder: conv stem cannot run on {n_frames} mel frames")
}
pub const STEM_LAYER_MODES: [(bool, usize, usize); 5] = [
(false, 2, 1), (true, 2, 1), (false, 1, 0), (true, 2, 1), (false, 1, 0), ];
pub const STEM_RELU_AFTER: [bool; 5] = [true, false, true, false, true];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Conv2dSpec {
pub in_ch: usize,
pub out_ch: usize,
pub h_in: usize,
pub w_in: usize,
pub kh: usize,
pub kw: usize,
pub stride_h: usize,
pub stride_w: usize,
pub pad_h: usize,
pub pad_w: usize,
pub h_out: usize,
pub w_out: usize,
pub groups: usize,
}
impl Conv2dSpec {
#[allow(clippy::too_many_arguments)]
pub fn padded(
in_ch: usize,
out_ch: usize,
h_in: usize,
w_in: usize,
(kh, kw): (usize, usize),
(stride_h, stride_w): (usize, usize),
(pad_h_lo, pad_h_hi): (usize, usize),
(pad_w_lo, pad_w_hi): (usize, usize),
groups: usize,
) -> Option<Self> {
if in_ch == 0 || out_ch == 0 || groups == 0 {
return None;
}
if stride_h == 0 || stride_w == 0 || kh == 0 || kw == 0 {
return None;
}
if !in_ch.is_multiple_of(groups) || !out_ch.is_multiple_of(groups) {
return None;
}
let padded_h = h_in.checked_add(pad_h_lo)?.checked_add(pad_h_hi)?;
let padded_w = w_in.checked_add(pad_w_lo)?.checked_add(pad_w_hi)?;
if padded_h < kh || padded_w < kw {
return None;
}
let h_out = (padded_h - kh) / stride_h + 1;
let w_out = (padded_w - kw) / stride_w + 1;
Some(Self {
in_ch,
out_ch,
h_in,
w_in,
kh,
kw,
stride_h,
stride_w,
pad_h: pad_h_lo,
pad_w: pad_w_lo,
h_out,
w_out,
groups,
})
}
pub fn out_len(&self) -> usize {
self.out_ch * self.h_out * self.w_out
}
}
pub trait AudioEncoderGpuOps {
type Buf;
type Weight;
fn upload(&self, data: &[f32]) -> Self::Buf;
fn download(&self, buf: &Self::Buf, len: usize) -> Vec<f32>;
fn upload_weight(&self, w: &MmapWeight) -> Self::Weight;
fn linear(
&self,
x: &Self::Buf,
w: &Self::Weight,
rows: usize,
out_dim: usize,
in_dim: usize,
) -> Self::Buf;
fn bias_add(&self, x: &Self::Buf, bias: &Self::Buf, rows: usize, dim: usize);
fn layernorm(
&self,
src: &Self::Buf,
weight: &Self::Buf,
bias: &Self::Buf,
eps: f32,
rows: usize,
dim: usize,
) -> Self::Buf;
fn relu(&self, x: &Self::Buf, len: usize);
fn silu(&self, x: &Self::Buf, len: usize);
fn gelu_erf(&self, x: &Self::Buf, len: usize);
fn add(&self, dst: &Self::Buf, src: &Self::Buf, len: usize);
fn scaled_add(&self, dst: &Self::Buf, src: &Self::Buf, len: usize, scale: f32);
fn conv2d(
&self,
input: &Self::Buf,
weight: &Self::Buf,
bias: &Self::Buf,
spec: &Conv2dSpec,
) -> Self::Buf;
fn transpose_blocked(&self, src: &Self::Buf, a: usize, b: usize, k: usize) -> Self::Buf;
fn glu_split(&self, src: &Self::Buf, rows: usize, n: usize) -> Self::Buf;
fn chan_affine_silu(
&self,
x: &Self::Buf,
w: &Self::Buf,
b: &Self::Buf,
channels: usize,
t: usize,
);
#[allow(clippy::too_many_arguments)]
fn xl_attention(
&self,
q: &Self::Buf,
k: &Self::Buf,
v: &Self::Buf,
p: &Self::Buf,
bias_u: &Self::Buf,
bias_v: &Self::Buf,
tokens: usize,
n_head: usize,
head_dim: usize,
) -> Self::Buf;
fn stft_frames(
&self,
pcm: &Self::Buf,
hann: &Self::Buf,
n_samples: usize,
n_frames: usize,
) -> Self::Buf;
fn power_spec(&self, frames: &Self::Buf, twiddle: &Self::Buf, n_frames: usize) -> Self::Buf;
fn mel_project(
&self,
power: &Self::Buf,
filters: &Self::Buf,
n_mel: usize,
n_frames: usize,
) -> Self::Buf;
fn mel_norm(
&self,
mel: &Self::Buf,
n_mel: usize,
n_frames: usize,
effective_n_len: usize,
) -> Self::Buf;
}
pub struct GpuMelFrontend<O: AudioEncoderGpuOps> {
hann: O::Buf,
twiddle: O::Buf,
filters: O::Buf,
n_mel: usize,
}
impl<O: AudioEncoderGpuOps> GpuMelFrontend<O> {
pub fn build(ops: &O, n_mel: usize) -> Result<Self> {
use crate::model::audio_encoder::{N_FFT, SAMPLE_RATE};
use crate::model::audio_preprocessor::{build_mel_filterbank, build_padded_hann_window};
anyhow::ensure!(n_mel > 0, "audio encoder config has n_mel_bins = 0");
Ok(Self {
hann: ops.upload(&build_padded_hann_window()),
twiddle: ops.upload(&dft_twiddles(N_FFT)),
filters: ops.upload(&build_mel_filterbank(n_mel, N_FFT, SAMPLE_RATE as usize)),
n_mel,
})
}
pub fn twiddle(&self) -> &O::Buf {
&self.twiddle
}
}
fn dft_twiddles(n_fft: usize) -> Vec<f32> {
(0..n_fft)
.flat_map(|m| {
let ang = -2.0 * std::f64::consts::PI * m as f64 / n_fft as f64;
[ang.cos() as f32, ang.sin() as f32]
})
.collect()
}
pub fn log_mel_spectrogram_gpu<O: AudioEncoderGpuOps>(
ops: &O,
fe: &GpuMelFrontend<O>,
pcm: &[f32],
) -> Result<Option<(O::Buf, usize)>> {
use crate::model::audio_encoder::N_FFT;
use crate::model::audio_preprocessor::{N_FFT_BINS, effective_n_len, n_frames_for};
let n_frames = n_frames_for(pcm.len());
if n_frames == 0 {
return Ok(None);
}
let fits = |n: usize| n <= u32::MAX as usize;
anyhow::ensure!(
fits(pcm.len())
&& [N_FFT, N_FFT_BINS, fe.n_mel]
.into_iter()
.all(|width| n_frames.checked_mul(width).is_some_and(fits)),
"audio front-end: {} samples ({n_frames} frames) exceed the kernels' u32 buffer indices",
pcm.len(),
);
let pcm_buf = ops.upload(pcm);
let frames = ops.stft_frames(&pcm_buf, &fe.hann, pcm.len(), n_frames);
let power = ops.power_spec(&frames, &fe.twiddle, n_frames);
let mel = ops.mel_project(&power, &fe.filters, fe.n_mel, n_frames);
let eff = effective_n_len(pcm.len(), n_frames);
Ok(Some((
ops.mel_norm(&mel, fe.n_mel, n_frames, eff),
n_frames,
)))
}
struct GpuStemLayer<O: AudioEncoderGpuOps> {
weight: O::Buf,
bias: O::Buf,
kernel: (usize, usize),
in_per_group: usize,
out_ch: usize,
}
struct GpuConformerBlock<O: AudioEncoderGpuOps> {
ffn_norm_w: O::Buf,
ffn_norm_b: O::Buf,
ffn_up_w: O::Weight,
ffn_up_b: O::Buf,
ffn_down_w: O::Weight,
ffn_down_b: O::Buf,
ln1_w: O::Buf,
ln1_b: O::Buf,
attn_q_w: O::Weight,
attn_q_b: O::Buf,
attn_k_w: O::Weight,
attn_k_b: O::Buf,
attn_v_w: O::Weight,
attn_v_b: O::Buf,
attn_o_w: O::Weight,
attn_o_b: O::Buf,
pos_bias_u: O::Buf,
pos_bias_v: O::Buf,
linear_pos_w: O::Weight,
norm_conv_w: O::Buf,
norm_conv_b: O::Buf,
conv_pw1_w: O::Weight,
conv_pw1_b: O::Buf,
conv_dw_w: O::Buf,
conv_dw_b: O::Buf,
conv_dw_k: usize,
conv_norm_w: O::Buf,
conv_norm_b: O::Buf,
conv_pw2_w: O::Weight,
conv_pw2_b: O::Buf,
ffn_norm_1_w: O::Buf,
ffn_norm_1_b: O::Buf,
ffn_up_1_w: O::Weight,
ffn_up_1_b: O::Buf,
ffn_down_1_w: O::Weight,
ffn_down_1_b: O::Buf,
ln2_w: O::Buf,
ln2_b: O::Buf,
}
fn validate_block(il: usize, b: &ConformerLayerWeights, cfg: &AudioEncoderConfig) -> Result<usize> {
let n_embd = cfg.n_embd;
let linears: [(&str, &MmapWeight, usize, usize); 10] = [
("ffn_up", &b.ffn_up_w, cfg.n_ff, n_embd),
("ffn_down", &b.ffn_down_w, n_embd, cfg.n_ff),
("ffn_up_1", &b.ffn_up_1_w, cfg.n_ff, n_embd),
("ffn_down_1", &b.ffn_down_1_w, n_embd, cfg.n_ff),
("attn_q", &b.attn_q_w, n_embd, n_embd),
("attn_k", &b.attn_k_w, n_embd, n_embd),
("attn_v", &b.attn_v_w, n_embd, n_embd),
("attn_o", &b.attn_o_w, n_embd, n_embd),
("conv_pw1", &b.conv_pw1_w, 2 * n_embd, n_embd),
("conv_pw2", &b.conv_pw2_w, n_embd, n_embd),
];
for (name, weight, rows, cols) in linears {
anyhow::ensure!(
weight.rows == rows && weight.cols == cols,
"audio encoder block {il}: {name} is [{}, {}], expected [{rows}, {cols}]",
weight.rows,
weight.cols,
);
}
anyhow::ensure!(
b.linear_pos_w.rows == n_embd && b.linear_pos_w.cols == POS_EMB_DIM,
"audio encoder block {il}: linear_pos is [{}, {}], expected \
[{n_embd}, {POS_EMB_DIM}]",
b.linear_pos_w.rows,
b.linear_pos_w.cols,
);
let vectors: [(&str, usize, usize); 25] = [
("ffn_norm_w", b.ffn_norm_w.len(), n_embd),
("ffn_norm_b", b.ffn_norm_b.len(), n_embd),
("ffn_norm_1_w", b.ffn_norm_1_w.len(), n_embd),
("ffn_norm_1_b", b.ffn_norm_1_b.len(), n_embd),
("ln1_w", b.ln1_w.len(), n_embd),
("ln1_b", b.ln1_b.len(), n_embd),
("ln2_w", b.ln2_w.len(), n_embd),
("ln2_b", b.ln2_b.len(), n_embd),
("norm_conv_w", b.norm_conv_w.len(), n_embd),
("norm_conv_b", b.norm_conv_b.len(), n_embd),
("pos_bias_u", b.pos_bias_u.len(), n_embd),
("pos_bias_v", b.pos_bias_v.len(), n_embd),
("conv_norm_w", b.conv_norm_w.len(), n_embd),
("conv_norm_b", b.conv_norm_b.len(), n_embd),
("attn_q_b", b.attn_q_b.len(), n_embd),
("attn_k_b", b.attn_k_b.len(), n_embd),
("attn_v_b", b.attn_v_b.len(), n_embd),
("attn_o_b", b.attn_o_b.len(), n_embd),
("ffn_up_b", b.ffn_up_b.len(), cfg.n_ff),
("ffn_down_b", b.ffn_down_b.len(), n_embd),
("ffn_up_1_b", b.ffn_up_1_b.len(), cfg.n_ff),
("ffn_down_1_b", b.ffn_down_1_b.len(), n_embd),
("conv_pw1_b", b.conv_pw1_b.len(), 2 * n_embd),
("conv_pw2_b", b.conv_pw2_b.len(), n_embd),
("conv_dw_b", b.conv_dw_b.len(), n_embd),
];
for (name, got, want) in vectors {
anyhow::ensure!(
got == want,
"audio encoder block {il}: {name} has {got} values, expected {want}",
);
}
let conv_dw_k = *b.conv_dw_shape.first().unwrap_or(&0);
anyhow::ensure!(
conv_dw_k > 0 && conv_dw_k * cfg.n_embd == b.conv_dw_w.len(),
"audio encoder block {il}: conv_dw shape {:?} disagrees with its \
{} weights at n_embd {}",
b.conv_dw_shape,
b.conv_dw_w.len(),
cfg.n_embd,
);
Ok(conv_dw_k)
}
pub struct GpuAudioWeights<O: AudioEncoderGpuOps> {
cfg: AudioEncoderConfig,
mel_frontend: GpuMelFrontend<O>,
stem: Vec<GpuStemLayer<O>>,
pre_encode_out_w: O::Weight,
pre_encode_out_b: O::Buf,
pre_encode_in_dim: usize,
blocks: Vec<GpuConformerBlock<O>>,
adapter_norm_w: O::Buf,
adapter_norm_b: O::Buf,
adapter_up_w: O::Weight,
adapter_up_b: O::Buf,
adapter_down_w: O::Weight,
adapter_down_b: O::Buf,
adapter_intermediate: usize,
}
impl<O: AudioEncoderGpuOps> GpuAudioWeights<O> {
pub fn build(ops: &O, w: &AudioEncoderWeights) -> Result<Self> {
let cfg = w.config.clone();
anyhow::ensure!(cfg.n_head > 0, "audio encoder config has n_head = 0");
anyhow::ensure!(
cfg.n_embd.is_multiple_of(cfg.n_head),
"audio encoder n_embd ({}) is not divisible by n_head ({})",
cfg.n_embd,
cfg.n_head,
);
let head_dim = cfg.n_embd / cfg.n_head;
anyhow::ensure!(
head_dim <= MAX_AUDIO_HEAD_DIM,
"audio encoder head_dim ({head_dim}) exceeds the attention kernel's \
MAX_AUDIO_HEAD_DIM ({MAX_AUDIO_HEAD_DIM}); caller should use the CPU encoder",
);
anyhow::ensure!(
w.layers.len() == cfg.n_layer,
"audio encoder config.n_layer ({}) != loaded blocks ({})",
cfg.n_layer,
w.layers.len(),
);
anyhow::ensure!(
w.conv_stem.layers.len() == STEM_LAYER_MODES.len(),
"audio encoder conv stem has {} layers, expected {}",
w.conv_stem.layers.len(),
STEM_LAYER_MODES.len(),
);
let stem = Self::build_stem(ops, &w.conv_stem)?;
let stem_in_ch = std::iter::once(1usize).chain(stem.iter().map(|l| l.out_ch));
for ((pos, layer), in_ch) in stem.iter().enumerate().zip(stem_in_ch) {
let (depthwise, ..) = STEM_LAYER_MODES[pos];
let groups = if depthwise { in_ch } else { 1 };
anyhow::ensure!(
layer.in_per_group * groups == in_ch,
"audio conv stem layer {pos}: in_per_group ({}) * groups ({groups}) != in_ch ({in_ch})",
layer.in_per_group,
);
}
anyhow::ensure!(
w.conv_stem.pre_encode_out_w.rows == cfg.n_embd
&& w.conv_stem.pre_encode_out_b.len() == cfg.n_embd,
"audio encoder pre_encode_out is [{}, {}] with {} bias values, expected {} rows",
w.conv_stem.pre_encode_out_w.rows,
w.conv_stem.pre_encode_out_w.cols,
w.conv_stem.pre_encode_out_b.len(),
cfg.n_embd,
);
let adapter_intermediate = w.mlp_adapter.up_w.rows;
anyhow::ensure!(
w.mlp_adapter.norm_w.len() == cfg.n_embd
&& w.mlp_adapter.norm_b.len() == cfg.n_embd
&& w.mlp_adapter.up_b.len() == adapter_intermediate
&& w.mlp_adapter.down_b.len() == cfg.llm_hidden_size,
"audio encoder MLP adapter vectors disagree with the config \
(norm {}/{}, up_b {}, down_b {}; n_embd {}, intermediate {}, llm_hidden {})",
w.mlp_adapter.norm_w.len(),
w.mlp_adapter.norm_b.len(),
w.mlp_adapter.up_b.len(),
w.mlp_adapter.down_b.len(),
cfg.n_embd,
adapter_intermediate,
cfg.llm_hidden_size,
);
anyhow::ensure!(
w.mlp_adapter.up_w.cols == cfg.n_embd
&& w.mlp_adapter.down_w.cols == adapter_intermediate
&& w.mlp_adapter.down_w.rows == cfg.llm_hidden_size,
"audio encoder MLP adapter shapes disagree with the config \
(up [{}, {}], down [{}, {}], n_embd {}, llm_hidden_size {})",
w.mlp_adapter.up_w.rows,
w.mlp_adapter.up_w.cols,
w.mlp_adapter.down_w.rows,
w.mlp_adapter.down_w.cols,
cfg.n_embd,
cfg.llm_hidden_size,
);
let blocks = w
.layers
.iter()
.enumerate()
.map(|(il, b)| {
let conv_dw_k = validate_block(il, b, &cfg)?;
Ok(GpuConformerBlock {
ffn_norm_w: ops.upload(&b.ffn_norm_w),
ffn_norm_b: ops.upload(&b.ffn_norm_b),
ffn_up_w: ops.upload_weight(&b.ffn_up_w),
ffn_up_b: ops.upload(&b.ffn_up_b),
ffn_down_w: ops.upload_weight(&b.ffn_down_w),
ffn_down_b: ops.upload(&b.ffn_down_b),
ln1_w: ops.upload(&b.ln1_w),
ln1_b: ops.upload(&b.ln1_b),
attn_q_w: ops.upload_weight(&b.attn_q_w),
attn_q_b: ops.upload(&b.attn_q_b),
attn_k_w: ops.upload_weight(&b.attn_k_w),
attn_k_b: ops.upload(&b.attn_k_b),
attn_v_w: ops.upload_weight(&b.attn_v_w),
attn_v_b: ops.upload(&b.attn_v_b),
attn_o_w: ops.upload_weight(&b.attn_o_w),
attn_o_b: ops.upload(&b.attn_o_b),
pos_bias_u: ops.upload(&b.pos_bias_u),
pos_bias_v: ops.upload(&b.pos_bias_v),
linear_pos_w: ops.upload_weight(&b.linear_pos_w),
norm_conv_w: ops.upload(&b.norm_conv_w),
norm_conv_b: ops.upload(&b.norm_conv_b),
conv_pw1_w: ops.upload_weight(&b.conv_pw1_w),
conv_pw1_b: ops.upload(&b.conv_pw1_b),
conv_dw_w: ops.upload(&b.conv_dw_w),
conv_dw_b: ops.upload(&b.conv_dw_b),
conv_dw_k,
conv_norm_w: ops.upload(&b.conv_norm_w),
conv_norm_b: ops.upload(&b.conv_norm_b),
conv_pw2_w: ops.upload_weight(&b.conv_pw2_w),
conv_pw2_b: ops.upload(&b.conv_pw2_b),
ffn_norm_1_w: ops.upload(&b.ffn_norm_1_w),
ffn_norm_1_b: ops.upload(&b.ffn_norm_1_b),
ffn_up_1_w: ops.upload_weight(&b.ffn_up_1_w),
ffn_up_1_b: ops.upload(&b.ffn_up_1_b),
ffn_down_1_w: ops.upload_weight(&b.ffn_down_1_w),
ffn_down_1_b: ops.upload(&b.ffn_down_1_b),
ln2_w: ops.upload(&b.ln2_w),
ln2_b: ops.upload(&b.ln2_b),
})
})
.collect::<Result<Vec<_>>>()?;
Ok(Self {
mel_frontend: GpuMelFrontend::build(ops, cfg.n_mel_bins)?,
cfg,
stem,
pre_encode_out_w: ops.upload_weight(&w.conv_stem.pre_encode_out_w),
pre_encode_out_b: ops.upload(&w.conv_stem.pre_encode_out_b),
pre_encode_in_dim: w.conv_stem.pre_encode_out_w.cols,
blocks,
adapter_norm_w: ops.upload(&w.mlp_adapter.norm_w),
adapter_norm_b: ops.upload(&w.mlp_adapter.norm_b),
adapter_up_w: ops.upload_weight(&w.mlp_adapter.up_w),
adapter_up_b: ops.upload(&w.mlp_adapter.up_b),
adapter_down_w: ops.upload_weight(&w.mlp_adapter.down_w),
adapter_down_b: ops.upload(&w.mlp_adapter.down_b),
adapter_intermediate,
})
}
fn build_stem(ops: &O, stem: &ConvStemWeights) -> Result<Vec<GpuStemLayer<O>>> {
stem.layers
.iter()
.enumerate()
.map(|(pos, layer)| {
anyhow::ensure!(
layer.shape.len() == 4,
"audio conv stem layer {pos} ({}): expected a 4-dim weight shape, got {:?}",
layer.name,
layer.shape,
);
let (kw, kh, in_per_group, out_ch) = (
layer.shape[0],
layer.shape[1],
layer.shape[2],
layer.shape[3],
);
anyhow::ensure!(
out_ch * in_per_group * kh * kw == layer.weight.len(),
"audio conv stem layer {pos} ({}): shape {:?} disagrees with its {} weights",
layer.name,
layer.shape,
layer.weight.len(),
);
anyhow::ensure!(
layer.bias.len() == out_ch,
"audio conv stem layer {pos} ({}): {} bias values for {out_ch} channels",
layer.name,
layer.bias.len(),
);
Ok(GpuStemLayer {
weight: ops.upload(&layer.weight),
bias: ops.upload(&layer.bias),
kernel: (kh, kw),
in_per_group,
out_ch,
})
})
.collect()
}
pub fn predict_t_out(&self, n_frames: usize) -> Option<usize> {
Some(self.stem_specs(n_frames)?.last()?.h_out)
}
fn stem_specs(&self, n_frames: usize) -> Option<Vec<Conv2dSpec>> {
self.stem
.iter()
.zip(&STEM_LAYER_MODES)
.try_fold(
(
Vec::with_capacity(STEM_LAYER_MODES.len()),
1usize,
n_frames,
self.cfg.n_mel_bins,
),
|(mut specs, ch, h, w), (layer, &(depthwise, stride, pad))| {
let spec = Conv2dSpec::padded(
ch,
layer.out_ch,
h,
w,
layer.kernel,
(stride, stride),
(pad, pad),
(pad, pad),
if depthwise { ch } else { 1 },
)?;
specs.push(spec);
Some((specs, spec.out_ch, spec.h_out, spec.w_out))
},
)
.map(|(specs, ..)| specs)
}
}
fn conv_stem_gpu<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
mel: &O::Buf,
n_frames: usize,
) -> Result<(O::Buf, usize)> {
let cfg = &gpu_w.cfg;
let specs = gpu_w
.stem_specs(n_frames)
.ok_or_else(|| anyhow::anyhow!("{}", unrunnable_stem_msg(n_frames)))?;
let no_layers = || anyhow::anyhow!("conv_stem_gpu: conv stem has no layers");
let last = specs.last().ok_or_else(no_layers)?;
let (t, f_out) = (last.h_out, last.w_out);
anyhow::ensure!(t > 0 && t <= MAX_AUDIO_TOKENS, "{}", over_capacity_msg(t));
let cur = gpu_w
.stem
.iter()
.zip(&specs)
.zip(STEM_RELU_AFTER)
.fold(None, |input, ((layer, spec), relu)| {
let out = ops.conv2d(
input.as_ref().unwrap_or(mel),
&layer.weight,
&layer.bias,
spec,
);
if relu {
ops.relu(&out, spec.out_len());
}
Some(out)
})
.ok_or_else(no_layers)?;
let cur_ch = last.out_ch;
let plane = cur_ch * f_out;
anyhow::ensure!(
plane == gpu_w.pre_encode_in_dim,
"audio conv stem produced {cur_ch}×{f_out} = {plane} features but pre_encode_out \
expects {}",
gpu_w.pre_encode_in_dim,
);
let flat = ops.transpose_blocked(&cur, cur_ch, t, f_out);
let x = ops.linear(&flat, &gpu_w.pre_encode_out_w, t, cfg.n_embd, plane);
ops.bias_add(&x, &gpu_w.pre_encode_out_b, t, cfg.n_embd);
Ok((x, t))
}
fn ensure_capacity<O: AudioEncoderGpuOps>(
gpu_w: &GpuAudioWeights<O>,
n_frames: usize,
) -> Result<()> {
let t_out = gpu_w
.predict_t_out(n_frames)
.ok_or_else(|| anyhow::anyhow!("{}", unrunnable_stem_msg(n_frames)))?;
anyhow::ensure!(
t_out > 0 && t_out <= MAX_AUDIO_TOKENS,
"{}",
over_capacity_msg(t_out)
);
Ok(())
}
fn upload_mel<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
mel: &[f32],
n_frames: usize,
) -> Result<Option<O::Buf>> {
anyhow::ensure!(
mel.len() == n_frames * gpu_w.cfg.n_mel_bins,
"audio encoder: mel.len() {} != n_frames * n_mel_bins ({n_frames} * {})",
mel.len(),
gpu_w.cfg.n_mel_bins,
);
if n_frames == 0 {
return Ok(None);
}
ensure_capacity(gpu_w, n_frames)?;
Ok(Some(ops.upload(mel)))
}
pub fn encoder_input_gpu<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
mel: &[f32],
n_frames: usize,
) -> Result<(Vec<f32>, usize)> {
let Some(mel_buf) = upload_mel(ops, gpu_w, mel, n_frames)? else {
return Ok((Vec::new(), 0));
};
let (x, t) = conv_stem_gpu(ops, gpu_w, &mel_buf, n_frames)?;
Ok((ops.download(&x, t * gpu_w.cfg.n_embd), t))
}
pub fn encode_audio_mel_gpu<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
mel: &[f32],
n_frames: usize,
) -> Result<(Vec<f32>, usize)> {
let Some(mel_buf) = upload_mel(ops, gpu_w, mel, n_frames)? else {
return Ok((Vec::new(), 0));
};
encode_audio_gpu(ops, gpu_w, &mel_buf, n_frames)
}
pub fn encode_audio_pcm_gpu<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
pcm: &[f32],
) -> Result<(Vec<f32>, usize)> {
let n_frames = crate::model::audio_preprocessor::n_frames_for(pcm.len());
if n_frames == 0 {
return Ok((Vec::new(), 0));
}
ensure_capacity(gpu_w, n_frames)?;
let Some((mel, n_frames)) = log_mel_spectrogram_gpu(ops, &gpu_w.mel_frontend, pcm)? else {
return Ok((Vec::new(), 0));
};
encode_audio_gpu(ops, gpu_w, &mel, n_frames)
}
fn encode_audio_gpu<O: AudioEncoderGpuOps>(
ops: &O,
gpu_w: &GpuAudioWeights<O>,
mel: &O::Buf,
n_frames: usize,
) -> Result<(Vec<f32>, usize)> {
let cfg = &gpu_w.cfg;
let (n_embd, n_ff, n_head, eps) = (cfg.n_embd, cfg.n_ff, cfg.n_head, cfg.eps);
let head_dim = n_embd / n_head;
let (mut x, t) = conv_stem_gpu(ops, gpu_w, mel, n_frames)?;
let seq_len = 2 * t - 1;
let pos_emb = ops.upload(&relative_pos_emb(t));
let x_len = t * n_embd;
for blk in &gpu_w.blocks {
macaron_ffn(
ops,
&x,
(&blk.ffn_norm_w, &blk.ffn_norm_b),
(&blk.ffn_up_w, &blk.ffn_up_b),
(&blk.ffn_down_w, &blk.ffn_down_b),
t,
n_embd,
n_ff,
eps,
);
let normed = ops.layernorm(&x, &blk.ln1_w, &blk.ln1_b, eps, t, n_embd);
let q = ops.linear(&normed, &blk.attn_q_w, t, n_embd, n_embd);
ops.bias_add(&q, &blk.attn_q_b, t, n_embd);
let k = ops.linear(&normed, &blk.attn_k_w, t, n_embd, n_embd);
ops.bias_add(&k, &blk.attn_k_b, t, n_embd);
let v = ops.linear(&normed, &blk.attn_v_w, t, n_embd, n_embd);
ops.bias_add(&v, &blk.attn_v_b, t, n_embd);
let p = ops.linear(&pos_emb, &blk.linear_pos_w, seq_len, n_embd, POS_EMB_DIM);
let attn = ops.xl_attention(
&q,
&k,
&v,
&p,
&blk.pos_bias_u,
&blk.pos_bias_v,
t,
n_head,
head_dim,
);
let proj = ops.linear(&attn, &blk.attn_o_w, t, n_embd, n_embd);
ops.bias_add(&proj, &blk.attn_o_b, t, n_embd);
ops.add(&x, &proj, x_len);
let normed = ops.layernorm(&x, &blk.norm_conv_w, &blk.norm_conv_b, eps, t, n_embd);
let pw1 = ops.linear(&normed, &blk.conv_pw1_w, t, 2 * n_embd, n_embd);
ops.bias_add(&pw1, &blk.conv_pw1_b, t, 2 * n_embd);
let glu = ops.glu_split(&pw1, t, n_embd);
let ch_major = ops.transpose_blocked(&glu, t, n_embd, 1);
let pad_total = blk.conv_dw_k - 1;
let pad_lo = pad_total / 2;
let conv_spec = Conv2dSpec::padded(
n_embd,
n_embd,
1,
t,
(1, blk.conv_dw_k),
(1, 1),
(0, 0),
(pad_lo, pad_total - pad_lo),
n_embd,
)
.ok_or_else(|| {
anyhow::anyhow!(
"audio conv module: degenerate depthwise conv (t {t}, k {})",
blk.conv_dw_k,
)
})?;
debug_assert_eq!(conv_spec.w_out, t, "conv module pad math drifted");
let conv = ops.conv2d(&ch_major, &blk.conv_dw_w, &blk.conv_dw_b, &conv_spec);
ops.chan_affine_silu(&conv, &blk.conv_norm_w, &blk.conv_norm_b, n_embd, t);
let time_major = ops.transpose_blocked(&conv, n_embd, t, 1);
let pw2 = ops.linear(&time_major, &blk.conv_pw2_w, t, n_embd, n_embd);
ops.bias_add(&pw2, &blk.conv_pw2_b, t, n_embd);
ops.add(&x, &pw2, x_len);
macaron_ffn(
ops,
&x,
(&blk.ffn_norm_1_w, &blk.ffn_norm_1_b),
(&blk.ffn_up_1_w, &blk.ffn_up_1_b),
(&blk.ffn_down_1_w, &blk.ffn_down_1_b),
t,
n_embd,
n_ff,
eps,
);
x = ops.layernorm(&x, &blk.ln2_w, &blk.ln2_b, eps, t, n_embd);
}
let n_ff_adapter = gpu_w.adapter_intermediate;
let llm_hidden = cfg.llm_hidden_size;
let normed = ops.layernorm(
&x,
&gpu_w.adapter_norm_w,
&gpu_w.adapter_norm_b,
eps,
t,
n_embd,
);
let mid = ops.linear(&normed, &gpu_w.adapter_up_w, t, n_ff_adapter, n_embd);
ops.bias_add(&mid, &gpu_w.adapter_up_b, t, n_ff_adapter);
ops.gelu_erf(&mid, t * n_ff_adapter);
let out = ops.linear(&mid, &gpu_w.adapter_down_w, t, llm_hidden, n_ff_adapter);
ops.bias_add(&out, &gpu_w.adapter_down_b, t, llm_hidden);
Ok((ops.download(&out, t * llm_hidden), t))
}
#[allow(clippy::too_many_arguments)]
fn macaron_ffn<O: AudioEncoderGpuOps>(
ops: &O,
x: &O::Buf,
norm: (&O::Buf, &O::Buf),
up: (&O::Weight, &O::Buf),
down: (&O::Weight, &O::Buf),
t: usize,
n_embd: usize,
n_ff: usize,
eps: f32,
) {
let normed = ops.layernorm(x, norm.0, norm.1, eps, t, n_embd);
let mid = ops.linear(&normed, up.0, t, n_ff, n_embd);
ops.bias_add(&mid, up.1, t, n_ff);
ops.silu(&mid, t * n_ff);
let out = ops.linear(&mid, down.0, t, n_embd, n_ff);
ops.bias_add(&out, down.1, t, n_embd);
ops.scaled_add(x, &out, t * n_embd, 0.5);
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
use crate::backend::metal::params::{
AudioXlAttnParams, Batch2dParams, BiasAddParams, Conv2dDirectParams, ElementwiseParams,
LayerNormBatchParams, MelNormParams, MelProjectParams, MetalParams, PowerSpecParams,
ScaleParams, StftFrameParams, TransposeBlockedParams,
};
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
use crate::backend::metal::{MetalLinear, MetalLinearWeight};
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
pub struct MetalAudioOps {
ctx: crate::backend::metal::MetalContext,
linear: MetalLinear,
p_bias: metal::ComputePipelineState,
p_layernorm: metal::ComputePipelineState,
p_relu: metal::ComputePipelineState,
p_silu: metal::ComputePipelineState,
p_gelu_erf: metal::ComputePipelineState,
p_add: metal::ComputePipelineState,
p_scaled_add: metal::ComputePipelineState,
p_conv2d: metal::ComputePipelineState,
p_transpose: metal::ComputePipelineState,
p_glu: metal::ComputePipelineState,
p_chan_affine: metal::ComputePipelineState,
p_attn: metal::ComputePipelineState,
p_stft: metal::ComputePipelineState,
p_power: metal::ComputePipelineState,
p_mel_project: metal::ComputePipelineState,
p_mel_norm: metal::ComputePipelineState,
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
impl MetalAudioOps {
pub fn new(ctx: crate::backend::metal::MetalContext) -> Result<Self> {
use crate::backend::metal::shaders;
Ok(Self {
linear: MetalLinear::new(&ctx)?,
p_bias: ctx.create_pipeline(shaders::BIAS_ADD, "bias_add")?,
p_layernorm: ctx.create_pipeline(shaders::LAYERNORM_BATCH, "layernorm_batch")?,
p_relu: ctx.create_pipeline(shaders::ACTIVATIONS, "relu_inplace")?,
p_silu: ctx.create_pipeline(shaders::ACTIVATIONS, "silu_inplace")?,
p_gelu_erf: ctx.create_pipeline(shaders::ACTIVATIONS, "gelu_erf_inplace")?,
p_add: ctx.create_pipeline(shaders::ELEMENTWISE_SLANG, "add_inplace")?,
p_scaled_add: ctx.create_pipeline(shaders::ELEMENTWISE_SLANG, "scaled_add_inplace")?,
p_conv2d: ctx.create_pipeline(shaders::CONV2D_DIRECT, "conv2d_direct")?,
p_transpose: ctx.create_pipeline(shaders::TRANSPOSE_BLOCKED, "transpose_blocked")?,
p_glu: ctx.create_pipeline(shaders::GLU_SPLIT, "glu_split")?,
p_chan_affine: ctx.create_pipeline(shaders::CHAN_AFFINE_SILU, "chan_affine_silu")?,
p_attn: ctx.create_pipeline(shaders::AUDIO_XL_ATTENTION, "audio_xl_attention")?,
p_stft: ctx.create_pipeline(shaders::STFT_FRAME, "stft_frame")?,
p_power: ctx.create_pipeline(shaders::POWER_SPEC, "power_spec")?,
p_mel_project: ctx.create_pipeline(shaders::MEL_PROJECT, "mel_project")?,
p_mel_norm: ctx.create_pipeline(shaders::MEL_NORM, "mel_norm")?,
ctx,
})
}
fn alloc(&self, len: usize) -> metal::Buffer {
self.ctx.create_buffer((len * 4) as u64)
}
fn run_flat<P: MetalParams>(
&self,
pipe: &metal::ComputePipelineState,
bufs: &[&metal::Buffer],
params: &P,
len: usize,
) {
self.ctx.run_kernel(
pipe,
bufs,
params,
metal::MTLSize::new((len as u64).div_ceil(256), 1, 1),
metal::MTLSize::new(256, 1, 1),
);
}
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
impl AudioEncoderGpuOps for MetalAudioOps {
type Buf = metal::Buffer;
type Weight = MetalLinearWeight;
fn upload(&self, data: &[f32]) -> Self::Buf {
self.ctx.upload_f32(data)
}
fn download(&self, buf: &Self::Buf, len: usize) -> Vec<f32> {
self.ctx.read_f32(buf, len)
}
fn upload_weight(&self, w: &MmapWeight) -> Self::Weight {
self.ctx.upload_linear_weight(w)
}
fn linear(
&self,
x: &Self::Buf,
w: &Self::Weight,
rows: usize,
out_dim: usize,
in_dim: usize,
) -> Self::Buf {
self.linear.forward(&self.ctx, x, w, rows, out_dim, in_dim)
}
fn bias_add(&self, x: &Self::Buf, bias: &Self::Buf, rows: usize, dim: usize) {
let total = rows * dim;
let params = BiasAddParams {
total: total as u32,
dim: dim as u32,
};
self.run_flat(&self.p_bias, &[x, bias], ¶ms, total);
}
fn layernorm(
&self,
src: &Self::Buf,
weight: &Self::Buf,
bias: &Self::Buf,
eps: f32,
rows: usize,
dim: usize,
) -> Self::Buf {
let dst = self.alloc(rows * dim);
let params = LayerNormBatchParams {
n: dim as u32,
eps_bits: eps.to_bits(),
src_stride: dim as u32,
dst_stride: dim as u32,
};
self.ctx.run_kernel(
&self.p_layernorm,
&[src, &dst, weight, bias],
¶ms,
metal::MTLSize::new(rows as u64, 1, 1),
metal::MTLSize::new(256, 1, 1),
);
dst
}
fn relu(&self, x: &Self::Buf, len: usize) {
self.run_flat(&self.p_relu, &[x], &ElementwiseParams::new(len as u32), len);
}
fn silu(&self, x: &Self::Buf, len: usize) {
self.run_flat(&self.p_silu, &[x], &ElementwiseParams::new(len as u32), len);
}
fn gelu_erf(&self, x: &Self::Buf, len: usize) {
self.run_flat(
&self.p_gelu_erf,
&[x],
&ElementwiseParams::new(len as u32),
len,
);
}
fn add(&self, dst: &Self::Buf, src: &Self::Buf, len: usize) {
self.run_flat(
&self.p_add,
&[dst, src],
&ElementwiseParams::new(len as u32),
len,
);
}
fn scaled_add(&self, dst: &Self::Buf, src: &Self::Buf, len: usize, scale: f32) {
let params = ScaleParams {
n: len as u32,
scale_bits: scale.to_bits(),
};
self.run_flat(&self.p_scaled_add, &[dst, src], ¶ms, len);
}
fn conv2d(
&self,
input: &Self::Buf,
weight: &Self::Buf,
bias: &Self::Buf,
spec: &Conv2dSpec,
) -> Self::Buf {
let total = spec.out_len();
let out = self.alloc(total);
let params = Conv2dDirectParams {
in_ch: spec.in_ch as u32,
out_ch: spec.out_ch as u32,
h_in: spec.h_in as u32,
w_in: spec.w_in as u32,
kh: spec.kh as u32,
kw: spec.kw as u32,
stride_h: spec.stride_h as u32,
stride_w: spec.stride_w as u32,
pad_h: spec.pad_h as u32,
pad_w: spec.pad_w as u32,
h_out: spec.h_out as u32,
w_out: spec.w_out as u32,
groups: spec.groups as u32,
_pad0: 0,
_pad1: 0,
_pad2: 0,
};
self.run_flat(&self.p_conv2d, &[input, weight, bias, &out], ¶ms, total);
out
}
fn transpose_blocked(&self, src: &Self::Buf, a: usize, b: usize, k: usize) -> Self::Buf {
let total = a * b * k;
let dst = self.alloc(total);
let params = TransposeBlockedParams {
a: a as u32,
b: b as u32,
k: k as u32,
_pad: 0,
};
self.run_flat(&self.p_transpose, &[src, &dst], ¶ms, total);
dst
}
fn glu_split(&self, src: &Self::Buf, rows: usize, n: usize) -> Self::Buf {
let total = rows * n;
let dst = self.alloc(total);
let params = Batch2dParams::new(rows as u32, n as u32);
self.run_flat(&self.p_glu, &[src, &dst], ¶ms, total);
dst
}
fn chan_affine_silu(
&self,
x: &Self::Buf,
w: &Self::Buf,
b: &Self::Buf,
channels: usize,
t: usize,
) {
let total = channels * t;
let params = Batch2dParams::new(channels as u32, t as u32);
self.run_flat(&self.p_chan_affine, &[x, w, b], ¶ms, total);
}
fn xl_attention(
&self,
q: &Self::Buf,
k: &Self::Buf,
v: &Self::Buf,
p: &Self::Buf,
bias_u: &Self::Buf,
bias_v: &Self::Buf,
tokens: usize,
n_head: usize,
head_dim: usize,
) -> Self::Buf {
let out = self.alloc(tokens * n_head * head_dim);
let params = AudioXlAttnParams {
tokens: tokens as u32,
n_head: n_head as u32,
head_dim: head_dim as u32,
scale_bits: (1.0f32 / (head_dim as f32).sqrt()).to_bits(),
};
self.ctx.run_kernel(
&self.p_attn,
&[q, k, v, p, bias_u, bias_v, &out],
¶ms,
metal::MTLSize::new(tokens as u64, n_head as u64, 1),
metal::MTLSize::new(256, 1, 1),
);
out
}
fn stft_frames(
&self,
pcm: &Self::Buf,
hann: &Self::Buf,
n_samples: usize,
n_frames: usize,
) -> Self::Buf {
use crate::model::audio_encoder::{HOP_LEN, N_FFT, PREEMPH};
let total = n_frames * N_FFT;
let frames = self.alloc(total);
let params = StftFrameParams {
n_frames: n_frames as u32,
n_fft: N_FFT as u32,
hop: HOP_LEN as u32,
center_pad: (N_FFT / 2) as u32,
n_samples: n_samples as u32,
preemph_bits: PREEMPH.to_bits(),
_pad0: 0,
_pad1: 0,
};
self.run_flat(&self.p_stft, &[pcm, hann, &frames], ¶ms, total);
frames
}
fn power_spec(&self, frames: &Self::Buf, twiddle: &Self::Buf, n_frames: usize) -> Self::Buf {
use crate::model::audio_encoder::N_FFT;
use crate::model::audio_preprocessor::N_FFT_BINS;
let total = n_frames * N_FFT_BINS;
let power = self.alloc(total);
let params = PowerSpecParams {
n_frames: n_frames as u32,
n_fft: N_FFT as u32,
n_bins: N_FFT_BINS as u32,
_pad: 0,
};
self.run_flat(&self.p_power, &[frames, twiddle, &power], ¶ms, total);
power
}
fn mel_project(
&self,
power: &Self::Buf,
filters: &Self::Buf,
n_mel: usize,
n_frames: usize,
) -> Self::Buf {
use crate::model::audio_encoder::LOG_MEL_EPS;
use crate::model::audio_preprocessor::N_FFT_BINS;
let total = n_mel * n_frames;
let mel = self.alloc(total);
let params = MelProjectParams {
n_mel: n_mel as u32,
n_frames: n_frames as u32,
n_bins: N_FFT_BINS as u32,
eps_bits: LOG_MEL_EPS.to_bits(),
};
self.run_flat(&self.p_mel_project, &[power, filters, &mel], ¶ms, total);
mel
}
fn mel_norm(
&self,
mel: &Self::Buf,
n_mel: usize,
n_frames: usize,
effective_n_len: usize,
) -> Self::Buf {
let dst = self.alloc(n_mel * n_frames);
let params = MelNormParams {
n_mel: n_mel as u32,
n_frames: n_frames as u32,
effective_n_len: effective_n_len as u32,
eps_bits: (crate::model::audio_encoder::NORM_VAR_EPS as f32).to_bits(),
};
self.ctx.run_kernel(
&self.p_mel_norm,
&[mel, &dst],
¶ms,
metal::MTLSize::new(n_mel as u64, 1, 1),
metal::MTLSize::new(256, 1, 1),
);
dst
}
}
pub trait AudioGpuEncode: Send + Sync {
fn encode_pcm(&self, pcm: &[f32]) -> Result<(Vec<f32>, usize)>;
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
struct MetalAudioEncoder {
ops: MetalAudioOps,
weights: GpuAudioWeights<MetalAudioOps>,
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
impl AudioGpuEncode for MetalAudioEncoder {
fn encode_pcm(&self, pcm: &[f32]) -> Result<(Vec<f32>, usize)> {
encode_audio_pcm_gpu(&self.ops, &self.weights, pcm)
}
}
pub fn build_gpu_audio_encoder(
weights: &AudioEncoderWeights,
backend: crate::engine::BackendPreference,
) -> Option<std::sync::Arc<dyn AudioGpuEncode>> {
use crate::engine::BackendPreference as BP;
match backend {
BP::Cpu | BP::Gpu => None,
BP::Metal | BP::Auto => try_metal_audio_encoder(weights),
}
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
fn try_metal_audio_encoder(
weights: &AudioEncoderWeights,
) -> Option<std::sync::Arc<dyn AudioGpuEncode>> {
let ctx = crate::backend::metal::MetalContext::new().ok()?;
let ops = MetalAudioOps::new(ctx).ok()?;
let gpu_w = match GpuAudioWeights::build(&ops, weights) {
Ok(w) => w,
Err(e) => {
tracing::warn!("audio encoder: Metal backend unavailable for this model: {e:#}");
return None;
}
};
tracing::info!("audio encoder: using native Metal backend");
Some(std::sync::Arc::new(MetalAudioEncoder {
ops,
weights: gpu_w,
}))
}
#[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
fn try_metal_audio_encoder(
_weights: &AudioEncoderWeights,
) -> Option<std::sync::Arc<dyn AudioGpuEncode>> {
None
}