use anyhow::{Context, Result};
use crate::gguf::GgufFile;
use std::sync::Arc;
use crate::model::weights::MmapWeight;
const KEY_N_LAYER: &str = "clip.audio.block_count";
const KEY_N_EMBD: &str = "clip.audio.embedding_length";
const KEY_N_FF: &str = "clip.audio.feed_forward_length";
const KEY_N_HEAD: &str = "clip.audio.attention.head_count";
const KEY_LN_EPS: &str = "clip.audio.attention.layer_norm_epsilon";
const KEY_N_MEL_BINS: &str = "clip.audio.num_mel_bins";
pub fn is_audio_encoder_gguf(gguf: &GgufFile) -> bool {
gguf.metadata.contains_key("clip.audio.block_count")
|| gguf.metadata.contains_key("clip.has_audio_encoder")
|| gguf.metadata.contains_key("audio.block_count")
|| gguf
.get_tensor("audio_model.encoder.layers.0.attn.q.weight")
.is_ok()
}
pub const SAMPLE_RATE: u32 = 16_000;
pub const N_FFT: usize = 512;
pub const WINDOW_LEN: usize = 400;
pub const HOP_LEN: usize = 160;
pub const PREEMPH: f32 = 0.97;
pub const LOG_MEL_EPS: f32 = 5.960_464_5e-8;
pub const NORM_VAR_EPS: f64 = 1e-5;
pub fn resample_linear(samples: &[f32], sr_in: u32, sr_out: u32) -> Vec<f32> {
if samples.is_empty()
|| !(1000..=192_000).contains(&sr_in)
|| !(1000..=192_000).contains(&sr_out)
{
return Vec::new();
}
if sr_in == sr_out {
return samples.to_vec();
}
let n_in = samples.len();
let ratio = sr_out as f64 / sr_in as f64;
let n_out = ((n_in as f64) * ratio).round().max(1.0) as usize;
let mut out = Vec::with_capacity(n_out);
let step = sr_in as f64 / sr_out as f64;
for i in 0..n_out {
let pos = i as f64 * step;
let idx = pos.floor() as usize;
let frac = (pos - idx as f64) as f32;
let a = samples[idx.min(n_in - 1)];
let b = samples[(idx + 1).min(n_in - 1)];
out.push(a + (b - a) * frac);
}
out
}
#[derive(Debug, Clone, PartialEq)]
pub struct AudioEncoderConfig {
pub n_layer: usize,
pub n_embd: usize,
pub n_ff: usize,
pub n_head: usize,
pub eps: f32,
pub n_mel_bins: usize,
pub llm_hidden_size: usize,
}
pub struct ConvLayerWeights {
pub name: String,
pub weight: Vec<f32>,
pub bias: Vec<f32>,
pub shape: Vec<usize>,
}
pub struct ConvStemWeights {
pub layers: Vec<ConvLayerWeights>,
pub pre_encode_out_w: MmapWeight,
pub pre_encode_out_b: Vec<f32>,
}
pub struct ConformerLayerWeights {
pub ffn_norm_w: Vec<f32>,
pub ffn_norm_b: Vec<f32>,
pub ffn_up_w: MmapWeight,
pub ffn_up_b: Vec<f32>,
pub ffn_down_w: MmapWeight,
pub ffn_down_b: Vec<f32>,
pub ln1_w: Vec<f32>,
pub ln1_b: Vec<f32>,
pub attn_q_w: MmapWeight,
pub attn_q_b: Vec<f32>,
pub attn_k_w: MmapWeight,
pub attn_k_b: Vec<f32>,
pub attn_v_w: MmapWeight,
pub attn_v_b: Vec<f32>,
pub attn_o_w: MmapWeight,
pub attn_o_b: Vec<f32>,
pub pos_bias_u: Vec<f32>,
pub pos_bias_v: Vec<f32>,
pub linear_pos_w: MmapWeight,
pub norm_conv_w: Vec<f32>,
pub norm_conv_b: Vec<f32>,
pub conv_pw1_w: MmapWeight,
pub conv_pw1_b: Vec<f32>,
pub conv_dw_w: Vec<f32>,
pub conv_dw_b: Vec<f32>,
pub conv_dw_shape: Vec<usize>,
pub conv_norm_w: Vec<f32>,
pub conv_norm_b: Vec<f32>,
pub conv_pw2_w: MmapWeight,
pub conv_pw2_b: Vec<f32>,
pub ffn_norm_1_w: Vec<f32>,
pub ffn_norm_1_b: Vec<f32>,
pub ffn_up_1_w: MmapWeight,
pub ffn_up_1_b: Vec<f32>,
pub ffn_down_1_w: MmapWeight,
pub ffn_down_1_b: Vec<f32>,
pub ln2_w: Vec<f32>,
pub ln2_b: Vec<f32>,
}
pub struct AudioMlpAdapterWeights {
pub norm_w: Vec<f32>,
pub norm_b: Vec<f32>,
pub up_w: MmapWeight,
pub up_b: Vec<f32>,
pub down_w: MmapWeight,
pub down_b: Vec<f32>,
}
pub struct AudioEncoderWeights {
pub config: AudioEncoderConfig,
pub conv_stem: ConvStemWeights,
pub layers: Vec<ConformerLayerWeights>,
pub mlp_adapter: AudioMlpAdapterWeights,
}
impl AudioEncoderWeights {
pub fn from_gguf(gguf: &Arc<GgufFile>) -> Result<Self> {
let n_layer = gguf
.get_u32(KEY_N_LAYER)
.with_context(|| format!("missing `{KEY_N_LAYER}`"))? as usize;
let n_embd = gguf
.get_u32(KEY_N_EMBD)
.with_context(|| format!("missing `{KEY_N_EMBD}`"))? as usize;
let n_head = gguf
.get_u32(KEY_N_HEAD)
.with_context(|| format!("missing `{KEY_N_HEAD}`"))? as usize;
let eps = gguf
.get_f32(KEY_LN_EPS)
.with_context(|| format!("missing `{KEY_LN_EPS}`"))?;
let n_mel_bins =
gguf.get_u32(KEY_N_MEL_BINS)
.with_context(|| format!("missing `{KEY_N_MEL_BINS}`"))? as usize;
let metadata_n_ff = gguf.get_u32(KEY_N_FF).map(|v| v as usize);
let stem_indices = [0u32, 2, 3, 5, 6];
let mut stem_layers = Vec::with_capacity(stem_indices.len());
for idx in stem_indices {
stem_layers.push(load_conv_layer(gguf, idx)?);
}
let pre_encode_out_w = MmapWeight::from_gguf(gguf, "a.pre_encode.out.weight")
.context("loading a.pre_encode.out.weight")?;
let pre_encode_out_b = load_vec_f32(gguf, "a.pre_encode.out.bias")?;
let conv_stem = ConvStemWeights {
layers: stem_layers,
pre_encode_out_w,
pre_encode_out_b,
};
let mut layers = Vec::with_capacity(n_layer);
for il in 0..n_layer {
layers.push(load_conformer_block(gguf, il)?);
}
anyhow::ensure!(
!layers.is_empty(),
"audio encoder must have at least one Conformer block"
);
let n_ff = layers[0].ffn_up_w.rows;
for (il, layer) in layers.iter().enumerate() {
anyhow::ensure!(
layer.ffn_up_w.rows == n_ff && layer.ffn_up_w.cols == n_embd,
"block {il} ffn_up shape ({}, {}) != ({n_ff}, {n_embd})",
layer.ffn_up_w.rows,
layer.ffn_up_w.cols
);
anyhow::ensure!(
layer.ffn_up_b.len() == n_ff,
"block {il} ffn_up.bias len ({}) != n_ff ({n_ff})",
layer.ffn_up_b.len()
);
anyhow::ensure!(
layer.ffn_down_w.rows == n_embd && layer.ffn_down_w.cols == n_ff,
"block {il} ffn_down shape ({}, {}) != ({n_embd}, {n_ff})",
layer.ffn_down_w.rows,
layer.ffn_down_w.cols
);
anyhow::ensure!(
layer.ffn_down_b.len() == n_embd,
"block {il} ffn_down.bias len ({}) != n_embd ({n_embd})",
layer.ffn_down_b.len()
);
anyhow::ensure!(
layer.ffn_up_1_w.rows == n_ff && layer.ffn_up_1_w.cols == n_embd,
"block {il} ffn_up_1 shape ({}, {}) != ({n_ff}, {n_embd})",
layer.ffn_up_1_w.rows,
layer.ffn_up_1_w.cols
);
anyhow::ensure!(
layer.ffn_up_1_b.len() == n_ff,
"block {il} ffn_up_1.bias len ({}) != n_ff ({n_ff})",
layer.ffn_up_1_b.len()
);
anyhow::ensure!(
layer.ffn_down_1_w.rows == n_embd && layer.ffn_down_1_w.cols == n_ff,
"block {il} ffn_down_1 shape ({}, {}) != ({n_embd}, {n_ff})",
layer.ffn_down_1_w.rows,
layer.ffn_down_1_w.cols
);
anyhow::ensure!(
layer.ffn_down_1_b.len() == n_embd,
"block {il} ffn_down_1.bias len ({}) != n_embd ({n_embd})",
layer.ffn_down_1_b.len()
);
anyhow::ensure!(
layer.ffn_norm_w.len() == n_embd && layer.ffn_norm_b.len() == n_embd,
"block {il} ffn_norm w/b lens ({}, {}) != n_embd ({n_embd})",
layer.ffn_norm_w.len(),
layer.ffn_norm_b.len()
);
anyhow::ensure!(
layer.ffn_norm_1_w.len() == n_embd && layer.ffn_norm_1_b.len() == n_embd,
"block {il} ffn_norm_1 w/b lens ({}, {}) != n_embd ({n_embd})",
layer.ffn_norm_1_w.len(),
layer.ffn_norm_1_b.len()
);
}
if let Some(meta) = metadata_n_ff
&& meta != n_ff
{
tracing::warn!(
target: "cera::audio_encoder",
metadata_n_ff = meta,
tensor_n_ff = n_ff,
"{KEY_N_FF} metadata disagrees with ffn_up tensor shape; trusting tensor"
);
}
let mlp_adapter = AudioMlpAdapterWeights {
norm_w: load_vec_f32(gguf, "mm.a.mlp.0.weight")?,
norm_b: load_vec_f32(gguf, "mm.a.mlp.0.bias")?,
up_w: MmapWeight::from_gguf(gguf, "mm.a.mlp.1.weight")
.context("loading mm.a.mlp.1.weight")?,
up_b: load_vec_f32(gguf, "mm.a.mlp.1.bias")?,
down_w: MmapWeight::from_gguf(gguf, "mm.a.mlp.3.weight")
.context("loading mm.a.mlp.3.weight")?,
down_b: load_vec_f32(gguf, "mm.a.mlp.3.bias")?,
};
let llm_hidden_size = mlp_adapter.down_w.rows;
let config = AudioEncoderConfig {
n_layer,
n_embd,
n_ff,
n_head,
eps,
n_mel_bins,
llm_hidden_size,
};
Ok(AudioEncoderWeights {
config,
conv_stem,
layers,
mlp_adapter,
})
}
}
fn load_conv_layer(gguf: &Arc<GgufFile>, idx: u32) -> Result<ConvLayerWeights> {
let weight_name = format!("a.conv1d.{idx}.weight");
let bias_name = format!("a.conv1d.{idx}.bias");
let weight_t = gguf
.get_tensor(&weight_name)
.with_context(|| format!("loading {weight_name}"))?;
let bias_t = gguf
.get_tensor(&bias_name)
.with_context(|| format!("loading {bias_name}"))?;
Ok(ConvLayerWeights {
name: weight_name,
shape: weight_t.shape().to_vec(),
weight: weight_t.to_f32_vec(),
bias: bias_t.to_f32_vec(),
})
}
fn load_conformer_block(gguf: &Arc<GgufFile>, il: usize) -> Result<ConformerLayerWeights> {
let pfx = format!("a.blk.{il}");
let wt_f32 = |name: &str| -> Result<MmapWeight> {
MmapWeight::from_gguf(gguf, name).with_context(|| format!("loading {name}"))
};
let vec_f32 = |name: &str| load_vec_f32(gguf, name);
let conv_dw_w_t = gguf
.get_tensor(&format!("{pfx}.conv_dw.weight"))
.with_context(|| format!("loading {pfx}.conv_dw.weight"))?;
let conv_dw_shape = conv_dw_w_t.shape().to_vec();
let conv_dw_w = conv_dw_w_t.to_f32_vec();
let conv_dw_b = vec_f32(&format!("{pfx}.conv_dw.bias"))?;
Ok(ConformerLayerWeights {
ffn_norm_w: vec_f32(&format!("{pfx}.ffn_norm.weight"))?,
ffn_norm_b: vec_f32(&format!("{pfx}.ffn_norm.bias"))?,
ffn_up_w: wt_f32(&format!("{pfx}.ffn_up.weight"))?,
ffn_up_b: vec_f32(&format!("{pfx}.ffn_up.bias"))?,
ffn_down_w: wt_f32(&format!("{pfx}.ffn_down.weight"))?,
ffn_down_b: vec_f32(&format!("{pfx}.ffn_down.bias"))?,
ln1_w: vec_f32(&format!("{pfx}.ln1.weight"))?,
ln1_b: vec_f32(&format!("{pfx}.ln1.bias"))?,
attn_q_w: wt_f32(&format!("{pfx}.attn_q.weight"))?,
attn_q_b: vec_f32(&format!("{pfx}.attn_q.bias"))?,
attn_k_w: wt_f32(&format!("{pfx}.attn_k.weight"))?,
attn_k_b: vec_f32(&format!("{pfx}.attn_k.bias"))?,
attn_v_w: wt_f32(&format!("{pfx}.attn_v.weight"))?,
attn_v_b: vec_f32(&format!("{pfx}.attn_v.bias"))?,
attn_o_w: wt_f32(&format!("{pfx}.attn_out.weight"))?,
attn_o_b: vec_f32(&format!("{pfx}.attn_out.bias"))?,
pos_bias_u: vec_f32(&format!("{pfx}.pos_bias_u"))?,
pos_bias_v: vec_f32(&format!("{pfx}.pos_bias_v"))?,
linear_pos_w: wt_f32(&format!("{pfx}.linear_pos.weight"))?,
norm_conv_w: vec_f32(&format!("{pfx}.norm_conv.weight"))?,
norm_conv_b: vec_f32(&format!("{pfx}.norm_conv.bias"))?,
conv_pw1_w: wt_f32(&format!("{pfx}.conv_pw1.weight"))?,
conv_pw1_b: vec_f32(&format!("{pfx}.conv_pw1.bias"))?,
conv_dw_w,
conv_dw_b,
conv_dw_shape,
conv_norm_w: vec_f32(&format!("{pfx}.conv_norm.weight"))?,
conv_norm_b: vec_f32(&format!("{pfx}.conv_norm.bias"))?,
conv_pw2_w: wt_f32(&format!("{pfx}.conv_pw2.weight"))?,
conv_pw2_b: vec_f32(&format!("{pfx}.conv_pw2.bias"))?,
ffn_norm_1_w: vec_f32(&format!("{pfx}.ffn_norm_1.weight"))?,
ffn_norm_1_b: vec_f32(&format!("{pfx}.ffn_norm_1.bias"))?,
ffn_up_1_w: wt_f32(&format!("{pfx}.ffn_up_1.weight"))?,
ffn_up_1_b: vec_f32(&format!("{pfx}.ffn_up_1.bias"))?,
ffn_down_1_w: wt_f32(&format!("{pfx}.ffn_down_1.weight"))?,
ffn_down_1_b: vec_f32(&format!("{pfx}.ffn_down_1.bias"))?,
ln2_w: vec_f32(&format!("{pfx}.ln2.weight"))?,
ln2_b: vec_f32(&format!("{pfx}.ln2.bias"))?,
})
}
fn load_vec_f32(gguf: &Arc<GgufFile>, name: &str) -> Result<Vec<f32>> {
let tensor = gguf
.get_tensor(name)
.with_context(|| format!("loading {name}"))?;
Ok(tensor.to_f32_vec())
}
pub const POS_EMB_DIM: usize = 512;
#[allow(clippy::chunks_exact_to_as_chunks)]
pub fn relative_pos_emb(n_frames: usize) -> Vec<f32> {
assert!(n_frames > 0, "n_frames must be > 0");
let two_n = n_frames
.checked_mul(2)
.expect("relative_pos_emb: 2 * n_frames overflowed usize");
let seq_len = two_n
.checked_sub(1)
.expect("relative_pos_emb: 2 * n_frames - 1 underflowed");
let total = POS_EMB_DIM
.checked_mul(seq_len)
.expect("relative_pos_emb: POS_EMB_DIM * seq_len overflowed usize");
let mut pos_emb = vec![0.0f32; total];
let inv_freq = inv_freq_cached();
let n_frames_f = n_frames as f64;
for (pos, row) in pos_emb.chunks_exact_mut(POS_EMB_DIM).enumerate() {
let rel_pos = n_frames_f - pos as f64 - 1.0;
for (i, pair) in row.chunks_exact_mut(2).enumerate() {
let (sin, cos) = ((rel_pos * inv_freq[i]) as f32).sin_cos();
pair[0] = sin;
pair[1] = cos;
}
}
pos_emb
}
fn inv_freq_cached() -> &'static [f64] {
static CACHE: std::sync::OnceLock<Vec<f64>> = std::sync::OnceLock::new();
CACHE.get_or_init(|| {
let log_10000 = (10000.0_f64).ln();
let half_dim = POS_EMB_DIM / 2;
(0..half_dim)
.map(|i| (-(log_10000 / POS_EMB_DIM as f64) * (2.0 * i as f64)).exp())
.collect()
})
}
#[allow(clippy::too_many_arguments)]
pub fn conformer_ffn_forward(
x: &mut [f32],
norm_w: &[f32],
norm_b: &[f32],
up_w: &MmapWeight,
up_b: &[f32],
down_w: &MmapWeight,
down_b: &[f32],
n_embd: usize,
n_ff: usize,
t: usize,
eps: f32,
scratch_pre_norm: &mut [f32],
scratch_ff: &mut [f32],
) {
debug_assert_eq!(x.len(), t * n_embd);
debug_assert_eq!(norm_w.len(), n_embd);
debug_assert_eq!(norm_b.len(), n_embd);
debug_assert_eq!(up_b.len(), n_ff);
debug_assert_eq!(down_b.len(), n_embd);
debug_assert_eq!(scratch_pre_norm.len(), n_embd);
debug_assert_eq!(scratch_ff.len(), n_ff);
debug_assert_eq!(up_w.rows, n_ff);
debug_assert_eq!(up_w.cols, n_embd);
debug_assert_eq!(down_w.rows, n_embd);
debug_assert_eq!(down_w.cols, n_ff);
for row in x.chunks_mut(n_embd) {
scratch_pre_norm.copy_from_slice(row);
crate::backend::cpu::layer_norm_inplace(scratch_pre_norm, norm_w, norm_b, eps);
up_w.gemv(scratch_pre_norm, scratch_ff);
crate::backend::cpu::add_inplace(scratch_ff, up_b);
crate::backend::cpu::silu_inplace(scratch_ff);
down_w.gemv(scratch_ff, scratch_pre_norm);
crate::backend::cpu::add_inplace(scratch_pre_norm, down_b);
for (xv, &dv) in row.iter_mut().zip(scratch_pre_norm.iter()) {
*xv += 0.5 * dv;
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn conformer_conv_module_forward(
x: &mut [f32],
norm_conv_w: &[f32],
norm_conv_b: &[f32],
pw1_w: &MmapWeight,
pw1_b: &[f32],
conv_dw_w: &[f32],
conv_dw_b: &[f32],
conv_norm_w: &[f32],
conv_norm_b: &[f32],
pw2_w: &MmapWeight,
pw2_b: &[f32],
n_embd: usize,
t: usize,
kernel_size: usize,
eps: f32,
) {
assert!(kernel_size > 0, "kernel_size must be > 0");
if t == 0 {
return;
}
let n_2embd = 2 * n_embd;
debug_assert_eq!(x.len(), t * n_embd);
debug_assert_eq!(norm_conv_w.len(), n_embd);
debug_assert_eq!(norm_conv_b.len(), n_embd);
debug_assert_eq!(pw1_w.rows, n_2embd);
debug_assert_eq!(pw1_w.cols, n_embd);
debug_assert_eq!(pw1_b.len(), n_2embd);
debug_assert_eq!(conv_dw_w.len(), n_embd * kernel_size);
debug_assert_eq!(conv_dw_b.len(), n_embd);
debug_assert_eq!(conv_norm_w.len(), n_embd);
debug_assert_eq!(conv_norm_b.len(), n_embd);
debug_assert_eq!(pw2_w.rows, n_embd);
debug_assert_eq!(pw2_w.cols, n_embd);
debug_assert_eq!(pw2_b.len(), n_embd);
let mut pre_norm = vec![0.0f32; n_embd];
let mut glu_time_major = vec![0.0f32; t * n_embd];
let mut pw1_out = vec![0.0f32; n_2embd];
for ti in 0..t {
let in_row = &x[ti * n_embd..(ti + 1) * n_embd];
pre_norm.copy_from_slice(in_row);
crate::backend::cpu::layer_norm_inplace(&mut pre_norm, norm_conv_w, norm_conv_b, eps);
pw1_w.gemv(&pre_norm, &mut pw1_out);
crate::backend::cpu::add_inplace(&mut pw1_out, pw1_b);
let out_row = &mut glu_time_major[ti * n_embd..(ti + 1) * n_embd];
crate::backend::cpu::glu_split(&pw1_out, out_row);
}
let pad_total = kernel_size - 1;
let pad_left = pad_total / 2;
let _pad_right = pad_total - pad_left;
let mut padded = vec![0.0f32; n_embd * (t + pad_total)];
for c in 0..n_embd {
for ti in 0..t {
padded[c * (t + pad_total) + pad_left + ti] = glu_time_major[ti * n_embd + c];
}
}
let mut conv_out = vec![0.0f32; n_embd * t];
let t_out = crate::backend::cpu::conv1d(
&padded,
conv_dw_w,
Some(conv_dw_b),
&mut conv_out,
n_embd, n_embd, t + pad_total, kernel_size,
1, 0, n_embd, );
debug_assert_eq!(
t_out, t,
"symmetric pad math drifted: t_out={t_out} != t={t}"
);
for c in 0..n_embd {
let row = &mut conv_out[c * t..(c + 1) * t];
let w = conv_norm_w[c];
let b = conv_norm_b[c];
for v in row.iter_mut() {
*v = *v * w + b;
}
crate::backend::cpu::silu_inplace(row);
}
let mut pw2_in = vec![0.0f32; n_embd];
let mut pw2_out = vec![0.0f32; n_embd];
for ti in 0..t {
for c in 0..n_embd {
pw2_in[c] = conv_out[c * t + ti];
}
pw2_w.gemv(&pw2_in, &mut pw2_out);
crate::backend::cpu::add_inplace(&mut pw2_out, pw2_b);
let res = &mut x[ti * n_embd..(ti + 1) * n_embd];
crate::backend::cpu::add_inplace(res, &pw2_out);
}
}
#[allow(clippy::too_many_arguments)]
pub fn conformer_self_attention_forward(
x: &mut [f32],
pos_emb: &[f32],
ln1_w: &[f32],
ln1_b: &[f32],
attn_q_w: &MmapWeight,
attn_q_b: &[f32],
attn_k_w: &MmapWeight,
attn_k_b: &[f32],
attn_v_w: &MmapWeight,
attn_v_b: &[f32],
attn_o_w: &MmapWeight,
attn_o_b: &[f32],
pos_bias_u: &[f32],
pos_bias_v: &[f32],
linear_pos_w: &MmapWeight,
n_embd: usize,
n_head: usize,
t: usize,
eps: f32,
) {
assert!(n_head > 0, "n_head must be > 0");
assert!(
n_embd.is_multiple_of(n_head),
"n_embd ({n_embd}) must be divisible by n_head ({n_head})"
);
if t == 0 {
return;
}
let d_head = n_embd / n_head;
let seq_len = t
.checked_mul(2)
.and_then(|v| v.checked_sub(1))
.expect("conformer_self_attention_forward: 2 * t - 1 overflowed usize");
let t_n_embd = t
.checked_mul(n_embd)
.expect("conformer_self_attention_forward: t * n_embd overflowed usize");
let seq_n_embd = seq_len
.checked_mul(n_embd)
.expect("conformer_self_attention_forward: seq_len * n_embd overflowed usize");
let scores_len = n_head
.checked_mul(t)
.and_then(|v| v.checked_mul(t))
.expect("conformer_self_attention_forward: n_head * t * t overflowed usize");
debug_assert_eq!(x.len(), t_n_embd);
debug_assert_eq!(pos_emb.len(), seq_len * POS_EMB_DIM);
debug_assert_eq!(ln1_w.len(), n_embd);
debug_assert_eq!(ln1_b.len(), n_embd);
debug_assert_eq!(attn_q_w.rows, n_embd);
debug_assert_eq!(attn_q_w.cols, n_embd);
debug_assert_eq!(attn_q_b.len(), n_embd);
debug_assert_eq!(attn_k_w.rows, n_embd);
debug_assert_eq!(attn_k_w.cols, n_embd);
debug_assert_eq!(attn_k_b.len(), n_embd);
debug_assert_eq!(attn_v_w.rows, n_embd);
debug_assert_eq!(attn_v_w.cols, n_embd);
debug_assert_eq!(attn_v_b.len(), n_embd);
debug_assert_eq!(attn_o_w.rows, n_embd);
debug_assert_eq!(attn_o_w.cols, n_embd);
debug_assert_eq!(attn_o_b.len(), n_embd);
debug_assert_eq!(pos_bias_u.len(), n_embd);
debug_assert_eq!(pos_bias_v.len(), n_embd);
debug_assert_eq!(linear_pos_w.rows, n_embd);
debug_assert_eq!(linear_pos_w.cols, POS_EMB_DIM);
let mut pre_norm = vec![0.0f32; n_embd];
let mut q_proj = vec![0.0f32; t_n_embd];
let mut k_proj = vec![0.0f32; t_n_embd];
let mut v_proj = vec![0.0f32; t_n_embd];
for ti in 0..t {
let in_row = &x[ti * n_embd..(ti + 1) * n_embd];
pre_norm.copy_from_slice(in_row);
crate::backend::cpu::layer_norm_inplace(&mut pre_norm, ln1_w, ln1_b, eps);
let q_row = &mut q_proj[ti * n_embd..(ti + 1) * n_embd];
attn_q_w.gemv(&pre_norm, q_row);
crate::backend::cpu::add_inplace(q_row, attn_q_b);
let k_row = &mut k_proj[ti * n_embd..(ti + 1) * n_embd];
attn_k_w.gemv(&pre_norm, k_row);
crate::backend::cpu::add_inplace(k_row, attn_k_b);
let v_row = &mut v_proj[ti * n_embd..(ti + 1) * n_embd];
attn_v_w.gemv(&pre_norm, v_row);
crate::backend::cpu::add_inplace(v_row, attn_v_b);
}
let mut p_proj = vec![0.0f32; seq_n_embd];
for pi in 0..seq_len {
let pe_row = &pos_emb[pi * POS_EMB_DIM..(pi + 1) * POS_EMB_DIM];
let p_row = &mut p_proj[pi * n_embd..(pi + 1) * n_embd];
linear_pos_w.gemv(pe_row, p_row);
}
let scale = 1.0f64 / (d_head as f64).sqrt();
let t_minus_1 = t - 1;
let mut scores = vec![0.0f32; scores_len];
let mut q_plus_u = vec![0.0f32; d_head];
let mut q_plus_v = vec![0.0f32; d_head];
for h in 0..n_head {
let u_h = &pos_bias_u[h * d_head..(h + 1) * d_head];
let v_h = &pos_bias_v[h * d_head..(h + 1) * d_head];
for q in 0..t {
let q_h = &q_proj[q * n_embd + h * d_head..q * n_embd + (h + 1) * d_head];
for d in 0..d_head {
q_plus_u[d] = q_h[d] + u_h[d];
q_plus_v[d] = q_h[d] + v_h[d];
}
for k in 0..t {
let k_h = &k_proj[k * n_embd + h * d_head..k * n_embd + (h + 1) * d_head];
let pos_idx = t_minus_1 + k - q;
let p_h =
&p_proj[pos_idx * n_embd + h * d_head..pos_idx * n_embd + (h + 1) * d_head];
let mut ac = 0.0f64;
let mut bd = 0.0f64;
for d in 0..d_head {
ac += q_plus_u[d] as f64 * k_h[d] as f64;
bd += q_plus_v[d] as f64 * p_h[d] as f64;
}
scores[h * t * t + q * t + k] = ((ac + bd) * scale) as f32;
}
}
}
for h in 0..n_head {
for q in 0..t {
let row = &mut scores[h * t * t + q * t..h * t * t + (q + 1) * t];
crate::backend::cpu::softmax_inplace(row);
}
}
let attn = scores;
let mut attn_v = vec![0.0f32; t_n_embd];
let mut acc = vec![0.0f64; d_head];
for h in 0..n_head {
for q in 0..t {
let attn_row = &attn[h * t * t + q * t..h * t * t + (q + 1) * t];
acc.fill(0.0);
for k in 0..t {
let attn_k = attn_row[k] as f64;
let v_row = &v_proj[k * n_embd + h * d_head..k * n_embd + (h + 1) * d_head];
for d in 0..d_head {
acc[d] += attn_k * v_row[d] as f64;
}
}
let out_slot = &mut attn_v[q * n_embd + h * d_head..q * n_embd + (h + 1) * d_head];
for d in 0..d_head {
out_slot[d] = acc[d] as f32;
}
}
}
let mut out_row = vec![0.0f32; n_embd];
for ti in 0..t {
let av_row = &attn_v[ti * n_embd..(ti + 1) * n_embd];
attn_o_w.gemv(av_row, &mut out_row);
crate::backend::cpu::add_inplace(&mut out_row, attn_o_b);
let res = &mut x[ti * n_embd..(ti + 1) * n_embd];
crate::backend::cpu::add_inplace(res, &out_row);
}
}
pub fn conv_stem_forward(
mel: &[f32],
n_frames: usize,
weights: &ConvStemWeights,
config: &AudioEncoderConfig,
) -> (Vec<f32>, usize) {
let n_mel_bins = config.n_mel_bins;
let n_embd = config.n_embd;
assert_eq!(
mel.len(),
n_frames * n_mel_bins,
"conv_stem_forward: mel.len() = {} != n_frames * n_mel_bins = {} * {}",
mel.len(),
n_frames,
n_mel_bins,
);
assert_eq!(
weights.layers.len(),
5,
"conv_stem_forward: expected exactly 5 stem conv layers (positions 0,2,3,5,6); got {}",
weights.layers.len()
);
if n_frames == 0 {
return (Vec::new(), 0);
}
let layer_modes: [(usize, usize, usize); 5] = [
(0, 2, 1), (1, 2, 1), (0, 1, 0), (1, 2, 1), (0, 1, 0), ];
let relu_after: [bool; 5] = [true, false, true, false, true];
let mut cur_data: Vec<f32> = mel.to_vec();
let mut cur_ch: usize = 1;
let mut cur_h: usize = n_frames;
let mut cur_w: usize = n_mel_bins;
for (pos, layer) in weights.layers.iter().enumerate() {
assert_eq!(
layer.shape.len(),
4,
"conv stem layer {pos}: expected 4-dim weight shape, got {:?}",
layer.shape
);
let kw = layer.shape[0];
let kh = layer.shape[1];
let in_per_group = layer.shape[2];
let out_ch = layer.shape[3];
let (groups_kind, stride, pad) = layer_modes[pos];
let groups = if groups_kind == 0 { 1 } else { cur_ch };
assert_eq!(
in_per_group * groups,
cur_ch,
"conv stem layer {pos}: in_per_group ({in_per_group}) * groups ({groups}) != cur_ch ({cur_ch})"
);
let two_pad = pad
.checked_mul(2)
.expect("conv_stem_forward: 2 * pad overflowed usize");
let padded_h = cur_h
.checked_add(two_pad)
.expect("conv_stem_forward: cur_h + 2 * pad overflowed usize");
let padded_w = cur_w
.checked_add(two_pad)
.expect("conv_stem_forward: cur_w + 2 * pad overflowed usize");
assert!(
padded_h >= kh,
"conv stem layer {pos}: kh ({kh}) > padded_h ({padded_h})"
);
assert!(
padded_w >= kw,
"conv stem layer {pos}: kw ({kw}) > padded_w ({padded_w})"
);
let new_h = (padded_h - kh) / stride + 1;
let new_w = (padded_w - kw) / stride + 1;
let next_len = out_ch
.checked_mul(new_h)
.and_then(|v| v.checked_mul(new_w))
.expect("conv_stem_forward: out_ch * new_h * new_w overflowed usize");
let mut next = vec![0.0f32; next_len];
crate::backend::cpu::conv2d(
&cur_data,
&layer.weight,
Some(&layer.bias),
&mut next,
cur_ch,
out_ch,
cur_h,
cur_w,
kh,
kw,
stride,
stride,
pad,
pad,
groups,
);
if relu_after[pos] {
crate::backend::cpu::relu_inplace(&mut next);
}
cur_data = next;
cur_ch = out_ch;
cur_h = new_h;
cur_w = new_w;
}
let t_out = cur_h;
let f_out = cur_w;
let plane = cur_ch * f_out;
debug_assert_eq!(weights.pre_encode_out_w.cols, plane);
debug_assert_eq!(weights.pre_encode_out_w.rows, n_embd);
debug_assert_eq!(weights.pre_encode_out_b.len(), n_embd);
let mut flat_per_step = vec![0.0f32; plane];
let mut encoder_in = vec![0.0f32; t_out * n_embd];
for ti in 0..t_out {
for c in 0..cur_ch {
let src =
&cur_data[c * t_out * f_out + ti * f_out..c * t_out * f_out + (ti + 1) * f_out];
let dst = &mut flat_per_step[c * f_out..(c + 1) * f_out];
dst.copy_from_slice(src);
}
let out_row = &mut encoder_in[ti * n_embd..(ti + 1) * n_embd];
weights.pre_encode_out_w.gemv(&flat_per_step, out_row);
crate::backend::cpu::add_inplace(out_row, &weights.pre_encode_out_b);
}
(encoder_in, t_out)
}
pub fn audio_encoder_forward(
mel: &[f32],
n_frames: usize,
weights: &AudioEncoderWeights,
) -> (Vec<f32>, usize) {
let cfg = &weights.config;
let n_embd = cfg.n_embd;
let n_ff = cfg.n_ff;
let n_head = cfg.n_head;
let n_layer = cfg.n_layer;
let eps = cfg.eps;
let llm_hidden_size = cfg.llm_hidden_size;
let (mut x, t_out) = conv_stem_forward(mel, n_frames, &weights.conv_stem, cfg);
if t_out == 0 {
return (Vec::new(), 0);
}
let pos_emb = relative_pos_emb(t_out);
let mut scratch_pre_norm = vec![0.0f32; n_embd];
let mut scratch_ff = vec![0.0f32; n_ff];
assert_eq!(
weights.layers.len(),
n_layer,
"audio_encoder_forward: config.n_layer ({}) != weights.layers.len() ({})",
n_layer,
weights.layers.len()
);
for il in 0..n_layer {
let layer = &weights.layers[il];
conformer_ffn_forward(
&mut x,
&layer.ffn_norm_w,
&layer.ffn_norm_b,
&layer.ffn_up_w,
&layer.ffn_up_b,
&layer.ffn_down_w,
&layer.ffn_down_b,
n_embd,
n_ff,
t_out,
eps,
&mut scratch_pre_norm,
&mut scratch_ff,
);
conformer_self_attention_forward(
&mut x,
&pos_emb,
&layer.ln1_w,
&layer.ln1_b,
&layer.attn_q_w,
&layer.attn_q_b,
&layer.attn_k_w,
&layer.attn_k_b,
&layer.attn_v_w,
&layer.attn_v_b,
&layer.attn_o_w,
&layer.attn_o_b,
&layer.pos_bias_u,
&layer.pos_bias_v,
&layer.linear_pos_w,
n_embd,
n_head,
t_out,
eps,
);
let dw_rank = layer.conv_dw_shape.len();
assert!(
dw_rank == 2 || dw_rank == 3,
"audio_encoder_forward: block {il}: expected 2- or 3-dim conv_dw shape, got {:?}",
layer.conv_dw_shape
);
let kernel_size = layer.conv_dw_shape[0];
assert_eq!(
kernel_size * n_embd,
layer.conv_dw_w.len(),
"audio_encoder_forward: block {il}: kernel_size ({kernel_size}) * n_embd ({n_embd}) != conv_dw_w.len() ({})",
layer.conv_dw_w.len()
);
conformer_conv_module_forward(
&mut x,
&layer.norm_conv_w,
&layer.norm_conv_b,
&layer.conv_pw1_w,
&layer.conv_pw1_b,
&layer.conv_dw_w,
&layer.conv_dw_b,
&layer.conv_norm_w,
&layer.conv_norm_b,
&layer.conv_pw2_w,
&layer.conv_pw2_b,
n_embd,
t_out,
kernel_size,
eps,
);
conformer_ffn_forward(
&mut x,
&layer.ffn_norm_1_w,
&layer.ffn_norm_1_b,
&layer.ffn_up_1_w,
&layer.ffn_up_1_b,
&layer.ffn_down_1_w,
&layer.ffn_down_1_b,
n_embd,
n_ff,
t_out,
eps,
&mut scratch_pre_norm,
&mut scratch_ff,
);
for row in x.chunks_exact_mut(n_embd) {
crate::backend::cpu::layer_norm_inplace(row, &layer.ln2_w, &layer.ln2_b, eps);
}
}
let n_ff_adapter = weights.mlp_adapter.up_w.rows;
assert_eq!(
weights.mlp_adapter.up_w.cols, n_embd,
"audio_encoder_forward: mlp_adapter.up_w.cols ({}) != n_embd ({n_embd})",
weights.mlp_adapter.up_w.cols
);
assert_eq!(
weights.mlp_adapter.down_w.rows, llm_hidden_size,
"audio_encoder_forward: mlp_adapter.down_w.rows ({}) != llm_hidden_size ({llm_hidden_size})",
weights.mlp_adapter.down_w.rows
);
assert_eq!(
weights.mlp_adapter.down_w.cols, n_ff_adapter,
"audio_encoder_forward: mlp_adapter.down_w.cols ({}) != up_w.rows ({n_ff_adapter})",
weights.mlp_adapter.down_w.cols
);
assert_eq!(
weights.mlp_adapter.up_b.len(),
n_ff_adapter,
"audio_encoder_forward: mlp_adapter.up_b.len ({}) != n_ff_adapter ({n_ff_adapter})",
weights.mlp_adapter.up_b.len()
);
assert_eq!(
weights.mlp_adapter.down_b.len(),
llm_hidden_size,
"audio_encoder_forward: mlp_adapter.down_b.len ({}) != llm_hidden_size ({llm_hidden_size})",
weights.mlp_adapter.down_b.len()
);
assert_eq!(
weights.mlp_adapter.norm_w.len(),
n_embd,
"audio_encoder_forward: mlp_adapter.norm_w.len ({}) != n_embd ({n_embd})",
weights.mlp_adapter.norm_w.len()
);
assert_eq!(
weights.mlp_adapter.norm_b.len(),
n_embd,
"audio_encoder_forward: mlp_adapter.norm_b.len ({}) != n_embd ({n_embd})",
weights.mlp_adapter.norm_b.len()
);
let mut adapter_mid = vec![0.0f32; n_ff_adapter];
let total_out_len = t_out
.checked_mul(llm_hidden_size)
.expect("audio_encoder_forward: t_out * llm_hidden_size overflowed usize");
let mut encoder_out = vec![0.0f32; total_out_len];
for ti in 0..t_out {
let in_row = &x[ti * n_embd..(ti + 1) * n_embd];
scratch_pre_norm.copy_from_slice(in_row);
crate::backend::cpu::layer_norm_inplace(
&mut scratch_pre_norm,
&weights.mlp_adapter.norm_w,
&weights.mlp_adapter.norm_b,
eps,
);
weights
.mlp_adapter
.up_w
.gemv(&scratch_pre_norm, &mut adapter_mid);
crate::backend::cpu::add_inplace(&mut adapter_mid, &weights.mlp_adapter.up_b);
crate::backend::cpu::gelu_erf_inplace(&mut adapter_mid);
let out_row = &mut encoder_out[ti * llm_hidden_size..(ti + 1) * llm_hidden_size];
weights.mlp_adapter.down_w.gemv(&adapter_mid, out_row);
crate::backend::cpu::add_inplace(out_row, &weights.mlp_adapter.down_b);
}
(encoder_out, t_out)
}
pub fn encode_audio_pcm(pcm: &[f32], weights: &AudioEncoderWeights) -> (Vec<f32>, usize) {
let (mel, n_frames) =
crate::model::audio_preprocessor::log_mel_spectrogram(pcm, weights.config.n_mel_bins);
if n_frames == 0 {
return (Vec::new(), 0);
}
audio_encoder_forward(&mel, n_frames, weights)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::gguf::GgufFile;
use std::sync::Arc;
fn empty_gguf_bytes() -> Vec<u8> {
let mut data = Vec::new();
data.extend_from_slice(&0x46554747u32.to_le_bytes()); data.extend_from_slice(&3u32.to_le_bytes()); data.extend_from_slice(&0u64.to_le_bytes()); data.extend_from_slice(&0u64.to_le_bytes()); data
}
#[test]
fn from_gguf_errors_on_missing_metadata() {
let bytes: Arc<[u8]> = Arc::from(empty_gguf_bytes().into_boxed_slice());
let gguf = Arc::new(GgufFile::from_bytes(bytes).expect("parse minimal gguf"));
match AudioEncoderWeights::from_gguf(&gguf) {
Ok(_) => panic!("expected missing-metadata error, got Ok"),
Err(e) => {
let msg = format!("{e:#}");
assert!(
msg.contains("clip.audio.block_count"),
"expected missing-key error, got: {msg}"
);
}
}
}
#[cfg(feature = "mmap")]
#[test]
#[ignore = "needs ~/.leap/models/LFM2.5-Audio-1.5B-Q4_0/mmproj-...gguf locally"]
fn from_gguf_real_lfm2a_audio_mmproj() {
let Ok(home) = std::env::var("HOME") else {
eprintln!("skip: HOME not set");
return;
};
let path = std::path::PathBuf::from(home)
.join(".leap")
.join("models")
.join("LFM2.5-Audio-1.5B-Q4_0")
.join("mmproj-LFM2.5-Audio-1.5B-Q4_0.gguf");
if !path.exists() {
eprintln!("skip: mmproj GGUF not at {}", path.display());
return;
}
let gguf = GgufFile::open_arc(&path).expect("open mmproj gguf");
let w = AudioEncoderWeights::from_gguf(&gguf).expect("load mmproj gguf");
assert_eq!(w.config.n_layer, 17, "n_layer");
assert_eq!(w.config.n_embd, 512, "n_embd");
assert_eq!(w.config.n_head, 8, "n_head");
assert_eq!(
w.config.n_ff, 2048,
"n_ff must be derived from ffn_up tensor (2048), not from \
clip.audio.feed_forward_length metadata (512)"
);
for (il, layer) in w.layers.iter().enumerate() {
assert_eq!(layer.ffn_up_w.rows, 2048, "block {il} ffn_up");
assert_eq!(layer.ffn_up_1_w.rows, 2048, "block {il} ffn_up_1");
}
}
#[test]
fn conformer_conv_module_forward_identity_weights() {
let n_embd = 4;
let t = 5;
let kernel_size = 3;
let eps = 1e-5;
let original_x: Vec<f32> = (0..t * n_embd).map(|i| (i as f32) * 0.1 + 0.05).collect();
let mut x = original_x.clone();
let norm_conv_w = vec![1.0; n_embd];
let norm_conv_b = vec![0.0; n_embd];
let mut pw1_w_data = vec![0.0f32; 2 * n_embd * n_embd];
for i in 0..n_embd {
pw1_w_data[i * n_embd + i] = 1.0;
}
let pw1_w = MmapWeight::from_owned_f32(pw1_w_data, 2 * n_embd, n_embd);
let pw1_b = vec![0.0; 2 * n_embd];
let mut conv_dw_w = vec![0.0f32; n_embd * kernel_size];
for c in 0..n_embd {
conv_dw_w[c * kernel_size + kernel_size / 2] = 1.0;
}
let conv_dw_b = vec![0.0; n_embd];
let conv_norm_w = vec![1.0; n_embd];
let conv_norm_b = vec![0.0; n_embd];
let mut pw2_w_data = vec![0.0f32; n_embd * n_embd];
for i in 0..n_embd {
pw2_w_data[i * n_embd + i] = 1.0;
}
let pw2_w = MmapWeight::from_owned_f32(pw2_w_data, n_embd, n_embd);
let pw2_b = vec![0.0; n_embd];
conformer_conv_module_forward(
&mut x,
&norm_conv_w,
&norm_conv_b,
&pw1_w,
&pw1_b,
&conv_dw_w,
&conv_dw_b,
&conv_norm_w,
&conv_norm_b,
&pw2_w,
&pw2_b,
n_embd,
t,
kernel_size,
eps,
);
for ti in 0..t {
let orig_t = &original_x[ti * n_embd..(ti + 1) * n_embd];
let mut ln_t = orig_t.to_vec();
crate::backend::cpu::layer_norm_inplace(&mut ln_t, &norm_conv_w, &norm_conv_b, eps);
for c in 0..n_embd {
let half = 0.5 * ln_t[c];
let silu = half / (1.0 + (-half).exp());
let expected = orig_t[c] + silu;
let actual = x[ti * n_embd + c];
assert!(
(actual - expected).abs() < 5e-3,
"t={ti}, c={c}: got {actual}, expected {expected}"
);
}
}
}
#[test]
fn conformer_ffn_forward_matches_scalar_reference() {
let n_embd = 4;
let n_ff = 4;
let t = 2;
let mut x: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let original_x = x.clone();
let norm_w = vec![1.0; n_embd];
let norm_b = vec![0.0; n_embd];
let identity = vec![
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, ];
let up_w = MmapWeight::from_owned_f32(identity.clone(), n_ff, n_embd);
let up_b = vec![0.0; n_ff];
let down_w = MmapWeight::from_owned_f32(identity, n_embd, n_ff);
let down_b = vec![0.0; n_embd];
let mut scratch_pre_norm = vec![0.0; n_embd];
let mut scratch_ff = vec![0.0; n_ff];
conformer_ffn_forward(
&mut x,
&norm_w,
&norm_b,
&up_w,
&up_b,
&down_w,
&down_b,
n_embd,
n_ff,
t,
1e-5,
&mut scratch_pre_norm,
&mut scratch_ff,
);
for ti in 0..t {
let orig = &original_x[ti * n_embd..(ti + 1) * n_embd];
let mut ln = orig.to_vec();
crate::backend::cpu::layer_norm_inplace(&mut ln, &norm_w, &norm_b, 1e-5);
let mut silu = ln;
crate::backend::cpu::silu_inplace(&mut silu);
for c in 0..n_embd {
let expected = orig[c] + 0.5 * silu[c];
let actual = x[ti * n_embd + c];
assert!(
(actual - expected).abs() < 1e-5,
"t={ti}, c={c}: got {actual}, expected {expected}"
);
}
}
}
#[test]
fn relative_pos_emb_n_frames_1_is_zero_sin_one_cos() {
let pe = relative_pos_emb(1);
assert_eq!(pe.len(), POS_EMB_DIM);
for i in 0..POS_EMB_DIM / 2 {
assert!(pe[2 * i].abs() < 1e-6, "sin slot {} = {}", 2 * i, pe[2 * i]);
assert!(
(pe[2 * i + 1] - 1.0).abs() < 1e-6,
"cos slot {} = {}",
2 * i + 1,
pe[2 * i + 1]
);
}
}
#[test]
fn relative_pos_emb_n_frames_2_first_freq_known_values() {
let pe = relative_pos_emb(2);
assert_eq!(pe.len(), 3 * POS_EMB_DIM);
let sin_1: f32 = 1.0_f32.sin();
let cos_1: f32 = 1.0_f32.cos();
assert!((pe[0] - sin_1).abs() < 1e-6, "row 0 sin[0] = {}", pe[0]);
assert!((pe[1] - cos_1).abs() < 1e-6, "row 0 cos[0] = {}", pe[1]);
assert!(
pe[POS_EMB_DIM].abs() < 1e-6,
"row 1 sin[0] = {}",
pe[POS_EMB_DIM]
);
assert!(
(pe[POS_EMB_DIM + 1] - 1.0).abs() < 1e-6,
"row 1 cos[0] = {}",
pe[POS_EMB_DIM + 1]
);
assert!(
(pe[2 * POS_EMB_DIM] + sin_1).abs() < 1e-6,
"row 2 sin[0] = {}",
pe[2 * POS_EMB_DIM]
);
assert!(
(pe[2 * POS_EMB_DIM + 1] - cos_1).abs() < 1e-6,
"row 2 cos[0] = {}",
pe[2 * POS_EMB_DIM + 1]
);
}
#[test]
fn relative_pos_emb_length() {
for &n in &[1usize, 2, 5, 100] {
let pe = relative_pos_emb(n);
assert_eq!(pe.len(), (2 * n - 1) * POS_EMB_DIM, "n_frames={n}");
}
}
#[test]
fn conformer_self_attention_forward_matches_scalar_reference() {
let n_embd = 6;
let n_head = 2;
let d_head = n_embd / n_head; let t = 4;
let seq_len = 2 * t - 1; let eps = 1e-5_f32;
let pat = |seed: usize, i: usize| -> f32 {
let v = ((i.wrapping_mul(7).wrapping_add(seed) * 31) % 23) as f32;
-0.4 + v * 0.06
};
let x_orig: Vec<f32> = (0..t * n_embd).map(|i| pat(11, i)).collect();
let mut x_under_test = x_orig.clone();
let pos_emb = relative_pos_emb(t);
let ln1_w: Vec<f32> = (0..n_embd).map(|i| 0.9 + 0.05 * i as f32).collect();
let ln1_b: Vec<f32> = (0..n_embd).map(|i| 0.01 * i as f32).collect();
let mk_w = |seed: usize| -> MmapWeight {
let data: Vec<f32> = (0..n_embd * n_embd).map(|i| pat(seed, i)).collect();
MmapWeight::from_owned_f32(data, n_embd, n_embd)
};
let mk_b = |seed: usize| -> Vec<f32> { (0..n_embd).map(|i| pat(seed, i) * 0.5).collect() };
let attn_q_w = mk_w(101);
let attn_q_b = mk_b(102);
let attn_k_w = mk_w(103);
let attn_k_b = mk_b(104);
let attn_v_w = mk_w(105);
let attn_v_b = mk_b(106);
let attn_o_w = mk_w(107);
let attn_o_b = mk_b(108);
let pos_bias_u: Vec<f32> = (0..n_embd).map(|i| pat(201, i)).collect();
let pos_bias_v: Vec<f32> = (0..n_embd).map(|i| pat(202, i)).collect();
let linear_pos_w = {
let mut data = vec![0.0f32; n_embd * POS_EMB_DIM];
for r in 0..n_embd {
for c in 0..12 {
data[r * POS_EMB_DIM + c] = pat(301 + r, c);
}
}
MmapWeight::from_owned_f32(data, n_embd, POS_EMB_DIM)
};
conformer_self_attention_forward(
&mut x_under_test,
&pos_emb,
&ln1_w,
&ln1_b,
&attn_q_w,
&attn_q_b,
&attn_k_w,
&attn_k_b,
&attn_v_w,
&attn_v_b,
&attn_o_w,
&attn_o_b,
&pos_bias_u,
&pos_bias_v,
&linear_pos_w,
n_embd,
n_head,
t,
eps,
);
let layer_norm = |row: &[f32], w: &[f32], b: &[f32]| -> Vec<f32> {
let n = row.len();
let mean = row.iter().map(|&v| v as f64).sum::<f64>() / n as f64;
let var = row.iter().map(|&v| (v as f64 - mean).powi(2)).sum::<f64>() / n as f64;
let inv_std = 1.0 / (var + eps as f64).sqrt();
row.iter()
.enumerate()
.map(|(i, &v)| ((v as f64 - mean) * inv_std * w[i] as f64 + b[i] as f64) as f32)
.collect()
};
let matvec = |w: &MmapWeight, x: &[f32], b: &[f32]| -> Vec<f32> {
let w_data = w.as_f32();
(0..w.rows)
.map(|r| {
let row = &w_data[r * w.cols..(r + 1) * w.cols];
let mut s = b[r] as f64;
for c in 0..w.cols {
s += row[c] as f64 * x[c] as f64;
}
s as f32
})
.collect()
};
let mut q = vec![0.0f32; t * n_embd];
let mut k = vec![0.0f32; t * n_embd];
let mut v = vec![0.0f32; t * n_embd];
for ti in 0..t {
let row = &x_orig[ti * n_embd..(ti + 1) * n_embd];
let ln = layer_norm(row, &ln1_w, &ln1_b);
q[ti * n_embd..(ti + 1) * n_embd].copy_from_slice(&matvec(&attn_q_w, &ln, &attn_q_b));
k[ti * n_embd..(ti + 1) * n_embd].copy_from_slice(&matvec(&attn_k_w, &ln, &attn_k_b));
v[ti * n_embd..(ti + 1) * n_embd].copy_from_slice(&matvec(&attn_v_w, &ln, &attn_v_b));
}
let zero_b = vec![0.0f32; n_embd];
let mut p_proj = vec![0.0f32; seq_len * n_embd];
for pi in 0..seq_len {
let pe = &pos_emb[pi * POS_EMB_DIM..(pi + 1) * POS_EMB_DIM];
p_proj[pi * n_embd..(pi + 1) * n_embd].copy_from_slice(&matvec(
&linear_pos_w,
pe,
&zero_b,
));
}
let scale = 1.0f32 / (d_head as f32).sqrt();
let mut scores = vec![0.0f32; n_head * t * t];
for h in 0..n_head {
for qi in 0..t {
for ki in 0..t {
let pos_idx = (t - 1) - qi + ki;
let mut ac = 0.0f64;
let mut bd = 0.0f64;
for d in 0..d_head {
let qd = q[qi * n_embd + h * d_head + d];
let kd = k[ki * n_embd + h * d_head + d];
let pd = p_proj[pos_idx * n_embd + h * d_head + d];
let u = pos_bias_u[h * d_head + d];
let vv = pos_bias_v[h * d_head + d];
ac += ((qd + u) * kd) as f64;
bd += ((qd + vv) * pd) as f64;
}
scores[h * t * t + qi * t + ki] = (ac + bd) as f32 * scale;
}
}
}
for h in 0..n_head {
for qi in 0..t {
let row = &mut scores[h * t * t + qi * t..h * t * t + (qi + 1) * t];
let m = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f64;
for v in row.iter_mut() {
*v = (*v - m).exp();
sum += *v as f64;
}
for v in row.iter_mut() {
*v = (*v as f64 / sum) as f32;
}
}
}
let mut attn_v = vec![0.0f32; t * n_embd];
for h in 0..n_head {
for qi in 0..t {
for d in 0..d_head {
let mut s = 0.0f64;
for ki in 0..t {
s += (scores[h * t * t + qi * t + ki] * v[ki * n_embd + h * d_head + d])
as f64;
}
attn_v[qi * n_embd + h * d_head + d] = s as f32;
}
}
}
let mut x_ref = x_orig.clone();
for ti in 0..t {
let av = &attn_v[ti * n_embd..(ti + 1) * n_embd];
let out = matvec(&attn_o_w, av, &attn_o_b);
for c in 0..n_embd {
x_ref[ti * n_embd + c] += out[c];
}
}
for i in 0..x_ref.len() {
let diff = (x_under_test[i] - x_ref[i]).abs();
assert!(
diff < 1e-4,
"mismatch at i={i}: under_test={} ref={} diff={}",
x_under_test[i],
x_ref[i],
diff
);
}
}
#[test]
fn conformer_self_attention_forward_empty_sequence_is_noop() {
let n_embd = 4;
let n_head = 2;
let mut x: Vec<f32> = vec![];
let pos_emb: Vec<f32> = vec![];
let zero_w = MmapWeight::from_owned_f32(vec![0.0; n_embd * n_embd], n_embd, n_embd);
let zero_b = vec![0.0; n_embd];
let zero_pos = vec![0.0; n_embd];
let zero_lp =
MmapWeight::from_owned_f32(vec![0.0; n_embd * POS_EMB_DIM], n_embd, POS_EMB_DIM);
let ones = vec![1.0; n_embd];
conformer_self_attention_forward(
&mut x, &pos_emb, &ones, &zero_b, &zero_w, &zero_b, &zero_w, &zero_b, &zero_w, &zero_b,
&zero_w, &zero_b, &zero_pos, &zero_pos, &zero_lp, n_embd, n_head, 0, 1e-5,
);
assert!(x.is_empty(), "t=0 should not modify or grow x");
}
fn zero_kernel_layer(
name: &str,
kw: usize,
kh: usize,
in_per_group: usize,
out_ch: usize,
bias: Vec<f32>,
) -> ConvLayerWeights {
ConvLayerWeights {
name: name.into(),
weight: vec![0.0f32; out_ch * in_per_group * kh * kw],
bias,
shape: vec![kw, kh, in_per_group, out_ch],
}
}
#[test]
fn conv_stem_forward_smoke() {
let out_ch = 4;
let n_mel_bins = 8;
let n_frames = 24;
let expected_t_out = 3;
let expected_f_out = 1;
let plane = out_ch * expected_f_out; let n_embd = plane;
let b5: Vec<f32> = (0..out_ch).map(|c| 0.1 * (c as f32 + 1.0)).collect();
let mut id_pw = vec![0.0f32; out_ch * out_ch];
for c in 0..out_ch {
id_pw[c * out_ch + c] = 1.0;
}
let stem = ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: id_pw.clone(),
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, out_ch, b5.clone()),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: id_pw,
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
],
pre_encode_out_w: {
let mut data = vec![0.0f32; n_embd * plane];
for i in 0..n_embd {
data[i * plane + i] = 1.0;
}
MmapWeight::from_owned_f32(data, n_embd, plane)
},
pre_encode_out_b: vec![0.0; n_embd],
};
let cfg = AudioEncoderConfig {
n_layer: 0,
n_embd,
n_ff: 0,
n_head: 0,
eps: 1e-5,
n_mel_bins,
llm_hidden_size: 0,
};
let mel: Vec<f32> = (0..n_frames * n_mel_bins)
.map(|i| ((i % 17) as f32) * 0.05 - 0.3)
.collect();
let (encoder_in, t_out) = conv_stem_forward(&mel, n_frames, &stem, &cfg);
assert_eq!(t_out, expected_t_out);
assert_eq!(encoder_in.len(), t_out * n_embd);
for ti in 0..t_out {
for c in 0..out_ch {
let expected = b5[c];
let actual = encoder_in[ti * n_embd + c];
assert!(
(actual - expected).abs() < 1e-6,
"encoder_in[ti={ti}, c={c}] = {actual}, expected {expected}"
);
}
}
}
#[test]
fn conv_stem_forward_empty_input_is_empty_output() {
let cfg = AudioEncoderConfig {
n_layer: 0,
n_embd: 2,
n_ff: 0,
n_head: 0,
eps: 1e-5,
n_mel_bins: 4,
llm_hidden_size: 0,
};
let stem = ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, 2, vec![0.0; 2]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, 2, vec![0.0; 2]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; 2],
shape: vec![1, 1, 2, 2],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, 2, vec![0.0; 2]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; 2],
shape: vec![1, 1, 2, 2],
},
],
pre_encode_out_w: MmapWeight::from_owned_f32(vec![0.0; 2 * 2], 2, 2),
pre_encode_out_b: vec![0.0; 2],
};
let (out, t_out) = conv_stem_forward(&[], 0, &stem, &cfg);
assert_eq!(t_out, 0);
assert!(out.is_empty());
}
fn neutral_conformer_block(n_embd: usize, n_ff: usize) -> ConformerLayerWeights {
let zero_w = |rows, cols| MmapWeight::from_owned_f32(vec![0.0; rows * cols], rows, cols);
ConformerLayerWeights {
ffn_norm_w: vec![1.0; n_embd],
ffn_norm_b: vec![0.0; n_embd],
ffn_up_w: zero_w(n_ff, n_embd),
ffn_up_b: vec![0.0; n_ff],
ffn_down_w: zero_w(n_embd, n_ff),
ffn_down_b: vec![0.0; n_embd],
ln1_w: vec![1.0; n_embd],
ln1_b: vec![0.0; n_embd],
attn_q_w: zero_w(n_embd, n_embd),
attn_q_b: vec![0.0; n_embd],
attn_k_w: zero_w(n_embd, n_embd),
attn_k_b: vec![0.0; n_embd],
attn_v_w: zero_w(n_embd, n_embd),
attn_v_b: vec![0.0; n_embd],
attn_o_w: zero_w(n_embd, n_embd),
attn_o_b: vec![0.0; n_embd],
pos_bias_u: vec![0.0; n_embd],
pos_bias_v: vec![0.0; n_embd],
linear_pos_w: zero_w(n_embd, POS_EMB_DIM),
norm_conv_w: vec![1.0; n_embd],
norm_conv_b: vec![0.0; n_embd],
conv_pw1_w: zero_w(2 * n_embd, n_embd),
conv_pw1_b: vec![0.0; 2 * n_embd],
conv_dw_w: vec![0.0; n_embd * 3],
conv_dw_b: vec![0.0; n_embd],
conv_dw_shape: vec![3, n_embd],
conv_norm_w: vec![1.0; n_embd],
conv_norm_b: vec![0.0; n_embd],
conv_pw2_w: zero_w(n_embd, n_embd),
conv_pw2_b: vec![0.0; n_embd],
ffn_norm_1_w: vec![1.0; n_embd],
ffn_norm_1_b: vec![0.0; n_embd],
ffn_up_1_w: zero_w(n_ff, n_embd),
ffn_up_1_b: vec![0.0; n_ff],
ffn_down_1_w: zero_w(n_embd, n_ff),
ffn_down_1_b: vec![0.0; n_embd],
ln2_w: vec![1.0; n_embd],
ln2_b: vec![0.0; n_embd],
}
}
#[test]
fn audio_encoder_forward_smoke() {
let n_embd = 4;
let n_ff = 8;
let n_head = 1;
let n_layer = 1;
let n_mel_bins = 8;
let n_frames = 24;
let llm_hidden_size = 6;
let n_ff_adapter = 8;
let expected_t_out = 3;
let out_ch = n_embd; let mut id_pw = vec![0.0f32; out_ch * out_ch];
for c in 0..out_ch {
id_pw[c * out_ch + c] = 1.0;
}
let conv_stem = ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: id_pw.clone(),
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, out_ch, vec![0.1; out_ch]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: id_pw,
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
],
pre_encode_out_w: {
let mut data = vec![0.0f32; n_embd * n_embd];
for i in 0..n_embd {
data[i * n_embd + i] = 1.0;
}
MmapWeight::from_owned_f32(data, n_embd, n_embd)
},
pre_encode_out_b: vec![0.0; n_embd],
};
let mut adapter_up = vec![0.0f32; n_ff_adapter * n_embd];
for i in 0..n_embd.min(n_ff_adapter) {
adapter_up[i * n_embd + i] = 1.0;
}
let mut adapter_down = vec![0.0f32; llm_hidden_size * n_ff_adapter];
for i in 0..llm_hidden_size.min(n_ff_adapter) {
adapter_down[i * n_ff_adapter + i] = 1.0;
}
let mlp_adapter = AudioMlpAdapterWeights {
norm_w: vec![1.0; n_embd],
norm_b: vec![0.0; n_embd],
up_w: MmapWeight::from_owned_f32(adapter_up, n_ff_adapter, n_embd),
up_b: vec![0.0; n_ff_adapter],
down_w: MmapWeight::from_owned_f32(adapter_down, llm_hidden_size, n_ff_adapter),
down_b: vec![0.0; llm_hidden_size],
};
let weights = AudioEncoderWeights {
config: AudioEncoderConfig {
n_layer,
n_embd,
n_ff,
n_head,
eps: 1e-5,
n_mel_bins,
llm_hidden_size,
},
conv_stem,
layers: vec![neutral_conformer_block(n_embd, n_ff)],
mlp_adapter,
};
let mel: Vec<f32> = (0..n_frames * n_mel_bins)
.map(|i| ((i % 17) as f32) * 0.05 - 0.3)
.collect();
let (encoder_out, t_out) = audio_encoder_forward(&mel, n_frames, &weights);
assert_eq!(t_out, expected_t_out);
assert_eq!(encoder_out.len(), t_out * llm_hidden_size);
for (i, &v) in encoder_out.iter().enumerate() {
assert!(v.is_finite(), "encoder_out[{i}] = {v} (not finite)");
}
}
#[test]
fn audio_encoder_forward_multi_layer_smoke() {
let n_embd = 4;
let n_ff = 8;
let n_head = 1;
let n_layer = 3;
let n_mel_bins = 8;
let n_frames = 24;
let llm_hidden_size = 6;
let n_ff_adapter = 8;
let expected_t_out = 3;
let out_ch = n_embd;
let mut id_pw = vec![0.0f32; out_ch * out_ch];
for c in 0..out_ch {
id_pw[c * out_ch + c] = 1.0;
}
let conv_stem = ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: id_pw.clone(),
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, out_ch, vec![0.1; out_ch]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: id_pw,
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
],
pre_encode_out_w: {
let mut data = vec![0.0f32; n_embd * n_embd];
for i in 0..n_embd {
data[i * n_embd + i] = 1.0;
}
MmapWeight::from_owned_f32(data, n_embd, n_embd)
},
pre_encode_out_b: vec![0.0; n_embd],
};
let mut adapter_up = vec![0.0f32; n_ff_adapter * n_embd];
for i in 0..n_embd.min(n_ff_adapter) {
adapter_up[i * n_embd + i] = 1.0;
}
let mut adapter_down = vec![0.0f32; llm_hidden_size * n_ff_adapter];
for i in 0..llm_hidden_size.min(n_ff_adapter) {
adapter_down[i * n_ff_adapter + i] = 1.0;
}
let mlp_adapter = AudioMlpAdapterWeights {
norm_w: vec![1.0; n_embd],
norm_b: vec![0.0; n_embd],
up_w: MmapWeight::from_owned_f32(adapter_up, n_ff_adapter, n_embd),
up_b: vec![0.0; n_ff_adapter],
down_w: MmapWeight::from_owned_f32(adapter_down, llm_hidden_size, n_ff_adapter),
down_b: vec![0.0; llm_hidden_size],
};
let weights = AudioEncoderWeights {
config: AudioEncoderConfig {
n_layer,
n_embd,
n_ff,
n_head,
eps: 1e-5,
n_mel_bins,
llm_hidden_size,
},
conv_stem,
layers: (0..n_layer)
.map(|_| neutral_conformer_block(n_embd, n_ff))
.collect(),
mlp_adapter,
};
let mel: Vec<f32> = (0..n_frames * n_mel_bins)
.map(|i| ((i % 17) as f32) * 0.05 - 0.3)
.collect();
let (encoder_out, t_out) = audio_encoder_forward(&mel, n_frames, &weights);
assert_eq!(t_out, expected_t_out);
assert_eq!(encoder_out.len(), t_out * llm_hidden_size);
for (i, &v) in encoder_out.iter().enumerate() {
assert!(v.is_finite(), "encoder_out[{i}] = {v} (not finite)");
}
}
#[test]
fn audio_encoder_forward_empty_input_is_empty_output() {
let n_embd = 2;
let weights = AudioEncoderWeights {
config: AudioEncoderConfig {
n_layer: 0,
n_embd,
n_ff: 2,
n_head: 1,
eps: 1e-5,
n_mel_bins: 4,
llm_hidden_size: 2,
},
conv_stem: ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; n_embd],
shape: vec![1, 1, n_embd, n_embd],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; n_embd],
shape: vec![1, 1, n_embd, n_embd],
},
],
pre_encode_out_w: MmapWeight::from_owned_f32(
vec![0.0; n_embd * n_embd],
n_embd,
n_embd,
),
pre_encode_out_b: vec![0.0; n_embd],
},
layers: vec![],
mlp_adapter: AudioMlpAdapterWeights {
norm_w: vec![1.0; n_embd],
norm_b: vec![0.0; n_embd],
up_w: MmapWeight::from_owned_f32(vec![0.0; n_embd * n_embd], n_embd, n_embd),
up_b: vec![0.0; n_embd],
down_w: MmapWeight::from_owned_f32(vec![0.0; n_embd * n_embd], n_embd, n_embd),
down_b: vec![0.0; n_embd],
},
};
let (out, t_out) = audio_encoder_forward(&[], 0, &weights);
assert_eq!(t_out, 0);
assert!(out.is_empty());
}
#[test]
fn encode_audio_pcm_smoke() {
let n_embd = 4;
let n_ff = 8;
let n_head = 1;
let n_mel_bins = 8; let llm_hidden_size = 6;
let n_ff_adapter = 8;
let out_ch = n_embd;
let mut id_pw = vec![0.0f32; out_ch * out_ch];
for c in 0..out_ch {
id_pw[c * out_ch + c] = 1.0;
}
let conv_stem = ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: id_pw.clone(),
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, out_ch, vec![0.0; out_ch]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: id_pw,
bias: vec![0.0; out_ch],
shape: vec![1, 1, out_ch, out_ch],
},
],
pre_encode_out_w: {
let mut data = vec![0.0f32; n_embd * n_embd];
for i in 0..n_embd {
data[i * n_embd + i] = 1.0;
}
MmapWeight::from_owned_f32(data, n_embd, n_embd)
},
pre_encode_out_b: vec![0.0; n_embd],
};
let mut adapter_up = vec![0.0f32; n_ff_adapter * n_embd];
for i in 0..n_embd.min(n_ff_adapter) {
adapter_up[i * n_embd + i] = 1.0;
}
let mut adapter_down = vec![0.0f32; llm_hidden_size * n_ff_adapter];
for i in 0..llm_hidden_size.min(n_ff_adapter) {
adapter_down[i * n_ff_adapter + i] = 1.0;
}
let mlp_adapter = AudioMlpAdapterWeights {
norm_w: vec![1.0; n_embd],
norm_b: vec![0.0; n_embd],
up_w: MmapWeight::from_owned_f32(adapter_up, n_ff_adapter, n_embd),
up_b: vec![0.0; n_ff_adapter],
down_w: MmapWeight::from_owned_f32(adapter_down, llm_hidden_size, n_ff_adapter),
down_b: vec![0.0; llm_hidden_size],
};
let weights = AudioEncoderWeights {
config: AudioEncoderConfig {
n_layer: 1,
n_embd,
n_ff,
n_head,
eps: 1e-5,
n_mel_bins,
llm_hidden_size,
},
conv_stem,
layers: vec![neutral_conformer_block(n_embd, n_ff)],
mlp_adapter,
};
let n_samples = SAMPLE_RATE as usize / 2;
let pcm: Vec<f32> = (0..n_samples)
.map(|i| (2.0 * std::f32::consts::PI * 1000.0 * i as f32 / SAMPLE_RATE as f32).sin())
.collect();
let (out, t_out) = encode_audio_pcm(&pcm, &weights);
assert!(t_out > 0, "expected at least one output frame");
assert_eq!(out.len(), t_out * llm_hidden_size);
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "out[{i}] = {v} (not finite)");
}
}
#[test]
fn encode_audio_pcm_empty_input_is_empty_output() {
let n_embd = 2;
let weights = AudioEncoderWeights {
config: AudioEncoderConfig {
n_layer: 0,
n_embd,
n_ff: 2,
n_head: 1,
eps: 1e-5,
n_mel_bins: 4,
llm_hidden_size: 2,
},
conv_stem: ConvStemWeights {
layers: vec![
zero_kernel_layer("a.conv1d.0.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
zero_kernel_layer("a.conv1d.2.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
ConvLayerWeights {
name: "a.conv1d.3.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; n_embd],
shape: vec![1, 1, n_embd, n_embd],
},
zero_kernel_layer("a.conv1d.5.weight", 3, 3, 1, n_embd, vec![0.0; n_embd]),
ConvLayerWeights {
name: "a.conv1d.6.weight".into(),
weight: vec![1.0, 0.0, 0.0, 1.0],
bias: vec![0.0; n_embd],
shape: vec![1, 1, n_embd, n_embd],
},
],
pre_encode_out_w: MmapWeight::from_owned_f32(
vec![0.0; n_embd * n_embd],
n_embd,
n_embd,
),
pre_encode_out_b: vec![0.0; n_embd],
},
layers: vec![],
mlp_adapter: AudioMlpAdapterWeights {
norm_w: vec![1.0; n_embd],
norm_b: vec![0.0; n_embd],
up_w: MmapWeight::from_owned_f32(vec![0.0; n_embd * n_embd], n_embd, n_embd),
up_b: vec![0.0; n_embd],
down_w: MmapWeight::from_owned_f32(vec![0.0; n_embd * n_embd], n_embd, n_embd),
down_b: vec![0.0; n_embd],
},
};
let (out, t_out) = encode_audio_pcm(&[], &weights);
assert_eq!(t_out, 0);
assert!(out.is_empty());
}
}