use std::sync::Arc;
use anyhow::{Context, Result};
use crate::backend::cpu;
use crate::gguf::GgufFile;
use crate::model::weights::MmapWeight;
pub trait AudioGpu: Send + Sync {
fn sample_audio_frame(&self, embedding: &[f32], temperature: f32, top_k: usize) -> [i32; 8];
fn sample_audio_frame_async<'a>(
&'a self,
embedding: &'a [f32],
temperature: f32,
top_k: usize,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<[i32; 8]>> + Send + 'a>> {
Box::pin(async move { Ok(self.sample_audio_frame(embedding, temperature, top_k)) })
}
#[cfg(feature = "gpu")]
fn sample_audio_frame_from_gpu_hidden_async<'a>(
&'a self,
_hidden_buf: &'a wgpu::Buffer,
_temperature: f32,
_top_k: usize,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<[i32; 8]>> + Send + 'a>> {
Box::pin(async move {
anyhow::bail!("sample_audio_frame_from_gpu_hidden_async not implemented")
})
}
fn detokenize_to_spectrum(&self, cpu_weights: &DetokenizerWeights, codes: &[i32]) -> Vec<f32>;
fn detokenize_to_spectrum_async<'a>(
&'a self,
cpu_weights: &'a DetokenizerWeights,
codes: &'a [i32],
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<Vec<f32>>> + Send + 'a>> {
Box::pin(async move { Ok(self.detokenize_to_spectrum(cpu_weights, codes)) })
}
fn reset_depthformer(&self);
fn reset_detokenizer(&self);
fn supports_depthformer(&self) -> bool;
fn istft_to_pcm(&self, spectrum: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> {
istft_to_pcm(spectrum, n_fft, hop_length)
}
fn istft_to_pcm_async<'a>(
&'a self,
spectrum: &'a [f32],
n_fft: usize,
hop_length: usize,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<Vec<f32>>> + Send + 'a>> {
Box::pin(async move { Ok(self.istft_to_pcm(spectrum, n_fft, hop_length)) })
}
fn try_acquire_session(&self) -> bool {
true
}
fn release_session(&self) {}
}
pub fn build_gpu_audio_decoder(
gguf: &Arc<GgufFile>,
backend: crate::engine::BackendPreference,
) -> Option<Arc<dyn AudioGpu>> {
use crate::engine::BackendPreference as BP;
match backend {
BP::Cpu => None,
BP::Metal => try_metal_audio_decoder(gguf),
BP::Gpu => try_wgpu_audio_decoder(gguf),
BP::Auto => try_metal_audio_decoder(gguf).or_else(|| try_wgpu_audio_decoder(gguf)),
}
}
#[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
fn try_metal_audio_decoder(gguf: &Arc<GgufFile>) -> Option<Arc<dyn AudioGpu>> {
let dummy_path = std::path::Path::new("");
match crate::model::metal_audio_decoder::MetalAudioDecoder::from_gguf(gguf, dummy_path) {
Ok(d) => {
tracing::info!("audio decoder: using native Metal backend");
Some(Arc::new(d))
}
Err(e) => {
tracing::warn!("audio decoder: Metal backend unavailable: {e:#}");
None
}
}
}
#[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
fn try_metal_audio_decoder(_gguf: &Arc<GgufFile>) -> Option<Arc<dyn AudioGpu>> {
None
}
#[cfg(feature = "gpu")]
fn try_wgpu_audio_decoder(gguf: &Arc<GgufFile>) -> Option<Arc<dyn AudioGpu>> {
let dummy_path = std::path::Path::new("");
match crate::model::wgpu_audio_decoder::WgpuAudioDecoder::from_gguf(gguf, dummy_path) {
Ok(d) => {
tracing::info!("audio decoder: using WGPU backend");
Some(Arc::new(d))
}
Err(e) => {
tracing::warn!("audio decoder: WGPU backend unavailable: {e:#}");
None
}
}
}
#[cfg(not(feature = "gpu"))]
fn try_wgpu_audio_decoder(_gguf: &Arc<GgufFile>) -> Option<Arc<dyn AudioGpu>> {
None
}
#[derive(Debug, Clone)]
pub struct DepthformerConfig {
pub n_layer: usize,
pub n_embd: usize,
pub n_head: usize,
pub n_head_kv: usize,
pub n_embd_head: usize,
pub ffn_dim: usize,
pub rms_norm_eps: f32,
pub rope_freq_base: f32,
pub max_seq_len: usize,
}
#[derive(Debug, Clone)]
pub struct DecoderConfig {
pub n_codebook: usize,
pub n_vocab: usize,
pub n_embd: usize, pub rms_norm_eps: f32,
}
pub struct DepthformerLayerWeights {
pub operator_norm: Vec<f32>,
pub wqkv: MmapWeight,
pub q_norm: Vec<f32>,
pub k_norm: Vec<f32>,
pub wo: MmapWeight,
pub ffn_norm: Vec<f32>,
pub w1: MmapWeight, pub w2: MmapWeight, pub w3: MmapWeight, }
pub struct CodebookWeights {
pub embedding: MmapWeight,
pub norm: Vec<f32>,
pub to_logits: MmapWeight,
}
pub struct AudioDecoderWeights {
pub depthformer_config: DepthformerConfig,
pub decoder_config: DecoderConfig,
pub depthformer_layers: Vec<DepthformerLayerWeights>,
pub depth_linear_w: MmapWeight,
pub depth_linear_b: Vec<f32>,
pub depth_embeddings: Vec<CodebookWeights>,
pub audio_embedding: CodebookWeights,
}
impl AudioDecoderWeights {
pub fn from_gguf(gguf: &Arc<GgufFile>) -> Result<Self> {
let n_layer = gguf
.get_u32("depthformer_n_layer")
.context("missing depthformer_n_layer")? as usize;
let n_embd = gguf
.get_u32("depthformer_n_embd")
.context("missing depthformer_n_embd")? as usize;
let qkv = MmapWeight::from_gguf(gguf, "depthformer.layers.0.operator.qkv_proj.weight")?;
let q_norm_w =
gguf.get_tensor("depthformer.layers.0.operator.attention.q_layernorm.weight")?;
let n_embd_head = q_norm_w.shape()[0]; let qkv_out = qkv.rows; let n_head = 32; let n_head_kv = 8;
assert_eq!(
qkv_out,
(n_head + 2 * n_head_kv) * n_embd_head,
"qkv_proj shape mismatch"
);
let w1_0 = MmapWeight::from_gguf(gguf, "depthformer.layers.0.feed_forward.w1.weight")?;
let ffn_dim = w1_0.rows;
let depthformer_config = DepthformerConfig {
n_layer,
n_embd,
n_head,
n_head_kv,
n_embd_head,
ffn_dim,
rms_norm_eps: 1e-5,
rope_freq_base: 1_000_000.0,
max_seq_len: 8,
};
let mut depthformer_layers = Vec::with_capacity(n_layer);
for i in 0..n_layer {
let prefix = format!("depthformer.layers.{i}");
depthformer_layers.push(DepthformerLayerWeights {
operator_norm: gguf
.get_tensor(&format!("{prefix}.operator_norm.weight"))?
.to_f32_vec(),
wqkv: MmapWeight::from_gguf(gguf, &format!("{prefix}.operator.qkv_proj.weight"))?,
q_norm: gguf
.get_tensor(&format!("{prefix}.operator.attention.q_layernorm.weight"))?
.to_f32_vec(),
k_norm: gguf
.get_tensor(&format!("{prefix}.operator.attention.k_layernorm.weight"))?
.to_f32_vec(),
wo: MmapWeight::from_gguf(gguf, &format!("{prefix}.operator.out_proj.weight"))?,
ffn_norm: gguf
.get_tensor(&format!("{prefix}.ffn_norm.weight"))?
.to_f32_vec(),
w1: MmapWeight::from_gguf(gguf, &format!("{prefix}.feed_forward.w1.weight"))?,
w2: MmapWeight::from_gguf(gguf, &format!("{prefix}.feed_forward.w2.weight"))?,
w3: MmapWeight::from_gguf(gguf, &format!("{prefix}.feed_forward.w3.weight"))?,
});
}
let depth_linear_w = MmapWeight::from_gguf(gguf, "depth_linear.weight")?;
let depth_linear_b = gguf.get_tensor("depth_linear.bias")?.to_f32_vec();
let n_codebook = 8;
let mut depth_embeddings = Vec::with_capacity(n_codebook);
for i in 0..n_codebook {
let prefix = format!("depth_embeddings.{i}");
depth_embeddings.push(CodebookWeights {
embedding: MmapWeight::from_gguf(gguf, &format!("{prefix}.embedding.weight"))?,
norm: gguf
.get_tensor(&format!("{prefix}.embedding_norm.weight"))?
.to_f32_vec(),
to_logits: MmapWeight::from_gguf(gguf, &format!("{prefix}.to_logits.weight"))?,
});
}
let audio_embedding = CodebookWeights {
embedding: MmapWeight::from_gguf(gguf, "audio_embedding.embedding.weight")?,
norm: gguf
.get_tensor("audio_embedding.embedding_norm.weight")?
.to_f32_vec(),
to_logits: MmapWeight::from_gguf(gguf, "audio_embedding.to_logits.weight")?,
};
let n_embd_llm = depth_linear_w.cols; let n_vocab = depth_embeddings[0].to_logits.rows;
let decoder_config = DecoderConfig {
n_codebook,
n_vocab,
n_embd: n_embd_llm,
rms_norm_eps: 1e-5,
};
Ok(Self {
depthformer_config,
decoder_config,
depthformer_layers,
depth_linear_w,
depth_linear_b,
depth_embeddings,
audio_embedding,
})
}
}
struct LayerKvCache {
k: Vec<f32>,
v: Vec<f32>,
}
pub struct DepthformerState {
kv: Vec<LayerKvCache>,
n_past: usize,
cur: Vec<f32>,
residual: Vec<f32>,
qkv: Vec<f32>,
attn_out: Vec<f32>,
proj: Vec<f32>,
gate: Vec<f32>,
up: Vec<f32>,
scores: Vec<f32>,
q8_scales: Vec<f32>,
q8_quants: Vec<i8>,
depthformer_in: Vec<f32>,
emb_row: Vec<f32>,
normed: Vec<f32>,
logits: Vec<f32>,
indices: Vec<usize>,
}
impl DepthformerState {
pub fn new(cfg: &DepthformerConfig) -> Self {
let cache_size = cfg.max_seq_len * cfg.n_head_kv * cfg.n_embd_head;
let kv = (0..cfg.n_layer)
.map(|_| LayerKvCache {
k: vec![0.0; cache_size],
v: vec![0.0; cache_size],
})
.collect();
let n_embd = cfg.n_embd;
let hd = cfg.n_embd_head;
let qkv_dim = cfg.n_head * hd + 2 * (cfg.n_head_kv * hd);
let nb = n_embd / 32;
Self {
kv,
n_past: 0,
cur: vec![0.0; n_embd],
residual: vec![0.0; n_embd],
qkv: vec![0.0; qkv_dim],
attn_out: vec![0.0; cfg.n_head * hd],
proj: vec![0.0; n_embd],
gate: vec![0.0; cfg.ffn_dim],
up: vec![0.0; cfg.ffn_dim],
scores: vec![0.0; cfg.max_seq_len],
q8_scales: vec![0.0; nb],
q8_quants: vec![0; n_embd],
depthformer_in: vec![0.0; n_embd],
emb_row: vec![0.0; n_embd],
normed: vec![0.0; n_embd],
logits: vec![0.0; 2049],
indices: Vec::with_capacity(2049),
}
}
pub fn reset(&mut self) {
for layer in &mut self.kv {
layer.k.fill(0.0);
layer.v.fill(0.0);
}
self.n_past = 0;
}
}
fn apply_rope_interleaved(x: &mut [f32], pos: usize, head_dim: usize, freq_base: f32) {
let theta_scale = freq_base.powf(-2.0 / head_dim as f32);
let mut theta = pos as f32;
for i in 0..head_dim / 2 {
let (sin_t, cos_t) = theta.sin_cos();
let x0 = x[2 * i];
let x1 = x[2 * i + 1];
x[2 * i] = x0 * cos_t - x1 * sin_t;
x[2 * i + 1] = x0 * sin_t + x1 * cos_t;
theta *= theta_scale;
}
}
fn apply_rope_neox(x: &mut [f32], pos: usize, head_dim: usize, freq_base: f32) {
let theta_scale = freq_base.powf(-2.0 / head_dim as f32);
let mut theta = pos as f32;
for i in 0..head_dim / 2 {
let (sin_t, cos_t) = theta.sin_cos();
let x0 = x[i];
let x1 = x[i + head_dim / 2];
x[i] = x0 * cos_t - x1 * sin_t;
x[i + head_dim / 2] = x0 * sin_t + x1 * cos_t;
theta *= theta_scale;
}
}
pub fn depthformer_forward(
weights: &AudioDecoderWeights,
state: &mut DepthformerState,
input: &[f32],
) -> Vec<f32> {
depthformer_forward_into(weights, state, input);
state.cur.clone()
}
pub fn depthformer_forward_into(
weights: &AudioDecoderWeights,
state: &mut DepthformerState,
input: &[f32],
) {
state.cur.copy_from_slice(input);
depthformer_forward_step(weights, state);
}
pub fn depthformer_forward_inplace(weights: &AudioDecoderWeights, state: &mut DepthformerState) {
state.cur.copy_from_slice(&state.depthformer_in);
depthformer_forward_step(weights, state);
}
fn depthformer_forward_step(weights: &AudioDecoderWeights, state: &mut DepthformerState) {
let cfg = &weights.depthformer_config;
let n_embd = cfg.n_embd;
let n_head = cfg.n_head;
let n_kv = cfg.n_head_kv;
let hd = cfg.n_embd_head;
let pos = state.n_past;
let kv_dim = n_kv * hd;
let q_dim = n_head * hd;
let k_dim = n_kv * hd;
let group_size = n_head / n_kv;
let scale = 1.0 / (hd as f32).sqrt();
for (il, lw) in weights.depthformer_layers.iter().enumerate() {
state.residual.copy_from_slice(&state.cur);
cpu::rmsnorm(&mut state.cur, &lw.operator_norm, cfg.rms_norm_eps);
let q8_scratch = if lw.wqkv.dtype != crate::tensor::DType::F32
|| lw.w1.dtype != crate::tensor::DType::F32
{
crate::backend::cpu::quantize_f32_to_q8_0_into(
&state.cur,
&mut state.q8_scales,
&mut state.q8_quants,
);
Some((&mut state.q8_scales, &mut state.q8_quants))
} else {
None
};
state.qkv.iter_mut().for_each(|v| *v = 0.0);
crate::backend::cpu::gemv_dispatch(
lw.wqkv.dtype,
lw.wqkv.data(),
&state.cur,
&mut state.qkv,
lw.wqkv.rows,
lw.wqkv.cols,
q8_scratch,
);
for h in 0..n_head {
let s = &mut state.qkv[h * hd..(h + 1) * hd];
cpu::rmsnorm(s, &lw.q_norm, cfg.rms_norm_eps);
}
for h in 0..n_kv {
let s = &mut state.qkv[q_dim + h * hd..q_dim + (h + 1) * hd];
cpu::rmsnorm(s, &lw.k_norm, cfg.rms_norm_eps);
}
for h in 0..n_head {
apply_rope_interleaved(
&mut state.qkv[h * hd..(h + 1) * hd],
pos,
hd,
cfg.rope_freq_base,
);
}
for h in 0..n_kv {
apply_rope_interleaved(
&mut state.qkv[q_dim + h * hd..q_dim + (h + 1) * hd],
pos,
hd,
cfg.rope_freq_base,
);
}
let kv = &mut state.kv[il];
for h in 0..n_kv {
let cache_off = pos * kv_dim + h * hd;
kv.k[cache_off..cache_off + hd]
.copy_from_slice(&state.qkv[q_dim + h * hd..q_dim + (h + 1) * hd]);
kv.v[cache_off..cache_off + hd]
.copy_from_slice(&state.qkv[q_dim + k_dim + h * hd..q_dim + k_dim + (h + 1) * hd]);
}
let seq_len = pos + 1;
state.attn_out.iter_mut().for_each(|v| *v = 0.0);
for h in 0..n_head {
let kv_h = h / group_size;
let q_head = &state.qkv[h * hd..(h + 1) * hd];
let kv_h_offset = kv_h * hd;
let sc = &mut state.scores[..seq_len];
cpu::attn_scores(q_head, &kv.k, sc, kv_dim, kv_h_offset, hd, scale, seq_len);
cpu::softmax_inplace(sc);
let out = &mut state.attn_out[h * hd..(h + 1) * hd];
cpu::attn_values(sc, &kv.v, out, kv_dim, kv_h_offset, hd, seq_len);
}
state.proj.iter_mut().for_each(|v| *v = 0.0);
lw.wo.gemv(&state.attn_out, &mut state.proj);
for (c, (r, p)) in state
.cur
.iter_mut()
.zip(state.residual.iter().zip(&state.proj))
{
*c = r + p;
}
state.residual.copy_from_slice(&state.cur);
cpu::rmsnorm(&mut state.cur, &lw.ffn_norm, cfg.rms_norm_eps);
let ffn_q8 = if lw.w1.dtype != crate::tensor::DType::F32 {
crate::backend::cpu::quantize_f32_to_q8_0_into(
&state.cur,
&mut state.q8_scales,
&mut state.q8_quants,
);
Some((&mut state.q8_scales, &mut state.q8_quants))
} else {
None
};
state.gate.iter_mut().for_each(|v| *v = 0.0);
state.up.iter_mut().for_each(|v| *v = 0.0);
crate::backend::cpu::gemv_dispatch(
lw.w1.dtype,
lw.w1.data(),
&state.cur,
&mut state.gate,
lw.w1.rows,
lw.w1.cols,
ffn_q8,
);
crate::backend::cpu::gemv_dispatch(
lw.w3.dtype,
lw.w3.data(),
&state.cur,
&mut state.up,
lw.w3.rows,
lw.w3.cols,
None,
);
cpu::silu_mul_inplace(&mut state.gate, &state.up);
state.proj.iter_mut().for_each(|v| *v = 0.0);
lw.w2.gemv(&state.gate, &mut state.proj);
for (c, (r, d)) in state
.cur
.iter_mut()
.zip(state.residual.iter().zip(&state.proj))
{
*c = r + d;
}
}
state.n_past += 1;
}
pub fn sample_audio_frame(
weights: &AudioDecoderWeights,
state: &mut DepthformerState,
embedding: &[f32],
temperature: f32,
top_k: usize,
) -> [i32; 8] {
let cfg = &weights.decoder_config;
let mut token = [0i32; 8];
let mut prev_token: i32 = -1;
state.reset();
let n_embd_d = weights.depth_embeddings[0].embedding.cols;
for j in 0..cfg.n_codebook {
let cb = &weights.depth_embeddings[j];
debug_assert_eq!(
cb.embedding.cols, n_embd_d,
"codebook {j} has mismatched depthformer dim"
);
state.depthformer_in.fill(0.0);
let row_start = j * n_embd_d;
weights
.depth_linear_w
.gemv_rows(embedding, &mut state.depthformer_in, row_start, n_embd_d);
for (r, out) in state.depthformer_in.iter_mut().enumerate() {
*out += weights.depth_linear_b[row_start + r];
}
if j > 0 && prev_token >= 0 {
let prev_cb = &weights.depth_embeddings[j - 1];
let tok = prev_token as usize;
if tok < prev_cb.embedding.rows {
prev_cb.embedding.dequantize_row(tok, &mut state.emb_row);
for (d, e) in state.depthformer_in.iter_mut().zip(&state.emb_row) {
*d += e;
}
}
}
depthformer_forward_inplace(weights, state);
state.normed.copy_from_slice(&state.cur);
cpu::rmsnorm(&mut state.normed, &cb.norm, cfg.rms_norm_eps);
state.logits.resize(cfg.n_vocab, 0.0);
state.logits.fill(0.0);
cb.to_logits.gemv(&state.normed, &mut state.logits);
let sampled = if state.logits.is_empty() {
0
} else if !temperature.is_finite() || temperature <= 0.0 || top_k <= 1 {
crate::sampler::argmax(&state.logits) as i32
} else {
let inv_temp = 1.0 / temperature;
for l in &mut state.logits {
*l *= inv_temp;
}
cpu::softmax_inplace(&mut state.logits);
let n_logits = state.logits.len();
let k = top_k.min(n_logits).max(1);
state.indices.clear();
state.indices.extend(0..n_logits);
if k < n_logits {
state.indices.select_nth_unstable_by(k - 1, |&a, &b| {
state.logits[b].total_cmp(&state.logits[a])
});
}
let top_indices = &state.indices[..k];
let sum: f32 = top_indices.iter().map(|&i| state.logits[i]).sum();
let mut r = rand::random::<f32>() * sum;
let mut picked = top_indices[0];
for &i in top_indices {
r -= state.logits[i];
if r <= 0.0 {
picked = i;
break;
}
}
picked as i32
};
token[j] = sampled;
prev_token = sampled;
}
token
}
pub fn embed_audio_token(weights: &AudioDecoderWeights, codes: &[i32; 8]) -> Vec<f32> {
let emb = &weights.audio_embedding.embedding;
let n_codebook = weights.decoder_config.n_codebook;
let n_vocab = 2049; let emb_dim = emb.cols;
let mut result = vec![0.0f32; emb_dim];
let mut row = vec![0f32; emb_dim];
for (j, &code) in codes.iter().enumerate() {
if code >= 0 {
let offset_idx = j * n_vocab + code as usize;
if offset_idx < emb.rows {
emb.dequantize_row(offset_idx, &mut row);
for (r, e) in result.iter_mut().zip(&row) {
*r += e;
}
}
}
}
result
}
#[derive(Debug, Clone)]
pub struct DetokenizerConfig {
pub n_layer: usize,
pub n_embd: usize,
pub n_head: usize,
pub n_head_kv: usize,
pub n_embd_head: usize,
pub ffn_dim: usize,
pub d_conv: usize, pub rms_norm_eps: f32,
pub rope_freq_base: f32,
pub swa_window_size: usize,
pub n_codes: usize,
pub n_fft: usize,
pub hop_length: usize,
pub sample_rate: usize,
pub layer_is_conv: Vec<bool>,
}
pub struct DetokLayerWeights {
pub operator_norm: Vec<f32>,
pub ffn_norm: Vec<f32>,
pub ffn_w1: MmapWeight, pub ffn_w2: MmapWeight, pub ffn_w3: MmapWeight, pub conv_in_proj: Option<MmapWeight>,
pub conv_out_proj: Option<MmapWeight>,
pub conv_weight: Option<Vec<f32>>, pub wq: Option<MmapWeight>,
pub wk: Option<MmapWeight>,
pub wv: Option<MmapWeight>,
pub wo: Option<MmapWeight>,
pub q_norm: Option<Vec<f32>>,
pub k_norm: Option<Vec<f32>>,
}
pub struct DetokenizerWeights {
pub config: DetokenizerConfig,
pub output_norm: Vec<f32>,
pub emb_weight: MmapWeight, pub lin_w: MmapWeight, pub lin_b: Vec<f32>,
pub layers: Vec<DetokLayerWeights>,
}
fn get_detok_tensor(gguf: &GgufFile, name1: &str, name2: &str) -> Result<crate::tensor::Tensor> {
gguf.get_tensor(name1)
.or_else(|_| gguf.get_tensor(name2))
.with_context(|| format!("tensor not found: neither `{name1}` nor `{name2}`"))
}
fn get_detok_mmap_weight(gguf: &Arc<GgufFile>, name1: &str, name2: &str) -> Result<MmapWeight> {
MmapWeight::from_gguf(gguf, name1)
.or_else(|_| MmapWeight::from_gguf(gguf, name2))
.with_context(|| format!("tensor weight not found: neither `{name1}` nor `{name2}`"))
}
impl DetokenizerWeights {
pub fn from_gguf(gguf: &Arc<GgufFile>) -> Result<Self> {
let layer_is_conv = vec![true, true, false, true, false, true, false, true];
let n_layer = layer_is_conv.len();
let conv_in = get_detok_mmap_weight(
gguf,
"lfm.layers.0.conv.in_proj.weight",
"blk.0.shortconv.in_proj.weight",
)?;
let n_embd = conv_in.cols;
let q_norm_w = get_detok_tensor(
gguf,
"lfm.layers.2.self_attn.q_layernorm.weight",
"blk.2.attn_q_norm.weight",
)?;
let n_embd_head = q_norm_w.shape()[0];
anyhow::ensure!(
n_embd_head > 0,
"detokenizer n_embd_head must be > 0 (q_layernorm shape was empty)"
);
let q_w = get_detok_mmap_weight(
gguf,
"lfm.layers.2.self_attn.q_proj.weight",
"blk.2.attn_q.weight",
)?;
let n_head = q_w.rows / n_embd_head;
let k_w = get_detok_mmap_weight(
gguf,
"lfm.layers.2.self_attn.k_proj.weight",
"blk.2.attn_k.weight",
)?;
let n_head_kv = k_w.rows / n_embd_head;
let ffn_w1_0 = get_detok_mmap_weight(
gguf,
"lfm.layers.0.feed_forward.w1.weight",
"blk.0.ffn_gate.weight",
)?;
let ffn_dim = ffn_w1_0.rows;
let config = DetokenizerConfig {
n_layer,
n_embd,
n_head,
n_head_kv,
n_embd_head,
ffn_dim,
d_conv: 2, rms_norm_eps: 1e-5,
rope_freq_base: 1_000_000.0,
swa_window_size: 30,
n_codes: 8,
n_fft: 1280,
hop_length: 320,
sample_rate: 24000,
layer_is_conv,
};
let output_norm =
get_detok_tensor(gguf, "lfm.embedding_norm.weight", "token_embd_norm.weight")?
.to_f32_vec();
let emb_weight = get_detok_mmap_weight(gguf, "emb.emb.weight", "token_embd.weight")?;
let lin_w = get_detok_mmap_weight(gguf, "lin.weight", "dense_2.weight")?;
let lin_b = get_detok_tensor(gguf, "lin.bias", "dense_2.bias")?.to_f32_vec();
anyhow::ensure!(
emb_weight.cols == n_embd,
"detokenizer emb_weight cols ({}) mismatch n_embd ({})",
emb_weight.cols,
n_embd
);
anyhow::ensure!(
emb_weight.rows == config.n_codes * 2048,
"detokenizer emb_weight rows ({}) mismatch expected {} ({} codebooks * 2048)",
emb_weight.rows,
config.n_codes * 2048,
config.n_codes
);
anyhow::ensure!(
lin_w.cols == n_embd,
"detokenizer lin_w cols ({}) mismatch n_embd ({})",
lin_w.cols,
n_embd
);
let expected_lin_out = (config.n_fft / 2 + 1) * 2;
anyhow::ensure!(
lin_w.rows == expected_lin_out,
"detokenizer lin_w rows ({}) mismatch expected {expected_lin_out}",
lin_w.rows
);
anyhow::ensure!(
lin_b.len() == expected_lin_out,
"detokenizer lin_b len ({}) mismatch expected {expected_lin_out}",
lin_b.len()
);
let mut layers = Vec::with_capacity(n_layer);
for i in 0..n_layer {
let prefix_lfm = format!("lfm.layers.{i}");
let prefix_blk = format!("blk.{i}");
let is_conv = config.layer_is_conv[i];
layers.push(DetokLayerWeights {
operator_norm: get_detok_tensor(
gguf,
&format!("{prefix_lfm}.operator_norm.weight"),
&format!("{prefix_blk}.attn_norm.weight"),
)?
.to_f32_vec(),
ffn_norm: get_detok_tensor(
gguf,
&format!("{prefix_lfm}.ffn_norm.weight"),
&format!("{prefix_blk}.ffn_norm.weight"),
)?
.to_f32_vec(),
ffn_w1: get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.feed_forward.w1.weight"),
&format!("{prefix_blk}.ffn_gate.weight"),
)?,
ffn_w2: get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.feed_forward.w2.weight"),
&format!("{prefix_blk}.ffn_down.weight"),
)?,
ffn_w3: get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.feed_forward.w3.weight"),
&format!("{prefix_blk}.ffn_up.weight"),
)?,
conv_in_proj: if is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.conv.in_proj.weight"),
&format!("{prefix_blk}.shortconv.in_proj.weight"),
)?)
} else {
None
},
conv_out_proj: if is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.conv.out_proj.weight"),
&format!("{prefix_blk}.shortconv.out_proj.weight"),
)?)
} else {
None
},
conv_weight: if is_conv {
Some(
get_detok_tensor(
gguf,
&format!("{prefix_lfm}.conv.conv.weight"),
&format!("{prefix_blk}.shortconv.conv.weight"),
)?
.to_f32_vec(),
)
} else {
None
},
wq: if !is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.self_attn.q_proj.weight"),
&format!("{prefix_blk}.attn_q.weight"),
)?)
} else {
None
},
wk: if !is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.self_attn.k_proj.weight"),
&format!("{prefix_blk}.attn_k.weight"),
)?)
} else {
None
},
wv: if !is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.self_attn.v_proj.weight"),
&format!("{prefix_blk}.attn_v.weight"),
)?)
} else {
None
},
wo: if !is_conv {
Some(get_detok_mmap_weight(
gguf,
&format!("{prefix_lfm}.self_attn.out_proj.weight"),
&format!("{prefix_blk}.attn_output.weight"),
)?)
} else {
None
},
q_norm: if !is_conv {
Some(
get_detok_tensor(
gguf,
&format!("{prefix_lfm}.self_attn.q_layernorm.weight"),
&format!("{prefix_blk}.attn_q_norm.weight"),
)?
.to_f32_vec(),
)
} else {
None
},
k_norm: if !is_conv {
Some(
get_detok_tensor(
gguf,
&format!("{prefix_lfm}.self_attn.k_layernorm.weight"),
&format!("{prefix_blk}.attn_k_norm.weight"),
)?
.to_f32_vec(),
)
} else {
None
},
});
}
Ok(Self {
config,
output_norm,
emb_weight,
lin_w,
lin_b,
layers,
})
}
}
pub struct DetokenizerState {
conv_bufs: Vec<Vec<f32>>,
attn_kv: Vec<Option<(Vec<f32>, Vec<f32>)>>,
n_past: usize,
}
impl DetokenizerState {
pub fn new(cfg: &DetokenizerConfig) -> Self {
let mut conv_bufs = Vec::new();
let mut attn_kv = Vec::new();
let kv_size = cfg.swa_window_size * cfg.n_head_kv * cfg.n_embd_head;
for &is_conv in &cfg.layer_is_conv {
if is_conv {
conv_bufs.push(vec![0.0; cfg.d_conv * cfg.n_embd]);
attn_kv.push(None);
} else {
conv_bufs.push(vec![]); attn_kv.push(Some((vec![0.0; kv_size], vec![0.0; kv_size])));
}
}
Self {
conv_bufs,
attn_kv,
n_past: 0,
}
}
pub fn reset(&mut self) {
for buf in &mut self.conv_bufs {
buf.fill(0.0);
}
for (k, v) in self.attn_kv.iter_mut().flatten() {
k.fill(0.0);
v.fill(0.0);
}
self.n_past = 0;
}
}
pub fn detok_embed_codes(weights: &DetokenizerWeights, codes: &[i32]) -> Vec<f32> {
let n_codes = weights.config.n_codes;
let emb = &weights.emb_weight;
let emb_dim = emb.cols;
let n_vocab_per_cb = emb.rows / n_codes;
let mut result = vec![0.0f32; emb_dim];
let mut row = vec![0f32; emb_dim];
for (j, &code) in codes.iter().enumerate() {
if code >= 0 {
let idx = j * n_vocab_per_cb + code as usize;
if idx < emb.rows {
emb.dequantize_row(idx, &mut row);
for (r, e) in result.iter_mut().zip(&row) {
*r += e;
}
}
}
}
let scale = 1.0 / n_codes as f32;
for r in &mut result {
*r *= scale;
}
result
}
pub fn upsample(input: &[f32], n_embd: usize, n_up: usize) -> Vec<f32> {
let n_in = input.len() / n_embd;
let n_out = n_in * n_up;
let mut output = vec![0.0; n_out * n_embd];
for i in 0..n_out {
let src_f = i as f32 / n_up as f32;
let src_i = src_f as usize;
let frac = src_f - src_i as f32;
let i0 = src_i.min(n_in - 1);
let i1 = (src_i + 1).min(n_in - 1);
for d in 0..n_embd {
output[i * n_embd + d] =
input[i0 * n_embd + d] * (1.0 - frac) + input[i1 * n_embd + d] * frac;
}
}
output
}
fn detok_conv_block(
lw: &DetokLayerWeights,
conv_buf: &mut [f32],
cur: &[f32],
n_embd: usize,
d_conv: usize,
) -> Vec<f32> {
let (Some(in_proj), Some(conv_w), Some(out_proj)) =
(&lw.conv_in_proj, &lw.conv_weight, &lw.conv_out_proj)
else {
return cur.to_vec();
};
let chunk_size = n_embd;
let mut bcx = vec![0.0; 3 * chunk_size];
in_proj.gemv(cur, &mut bcx);
let b = &bcx[..chunk_size];
let c = &bcx[chunk_size..2 * chunk_size];
let x = &bcx[2 * chunk_size..];
let bx: Vec<f32> = b.iter().zip(x).map(|(bi, xi)| bi * xi).collect();
let kernel_size = d_conv + 1;
let mut conv_out = vec![0.0; chunk_size];
for ch in 0..chunk_size {
let mut sum = 0.0;
for k in 0..d_conv {
sum += conv_buf[k * n_embd + ch] * conv_w[ch * kernel_size + k];
}
sum += bx[ch] * conv_w[ch * kernel_size + d_conv];
conv_out[ch] = sum;
}
if d_conv > 1 {
let row_size = n_embd;
for k in 0..d_conv - 1 {
let src = (k + 1) * row_size;
let dst = k * row_size;
conv_buf.copy_within(src..src + row_size, dst);
}
}
conv_buf[(d_conv - 1) * n_embd..d_conv * n_embd].copy_from_slice(&bx);
let y: Vec<f32> = c.iter().zip(&conv_out).map(|(ci, co)| ci * co).collect();
let mut out = vec![0.0; n_embd];
out_proj.gemv(&y, &mut out);
out
}
#[allow(dead_code)]
fn detok_attn_block(
lw: &DetokLayerWeights,
kv: &mut (Vec<f32>, Vec<f32>),
cur: &[f32],
pos: usize,
cfg: &DetokenizerConfig,
) -> Vec<f32> {
let (Some(wq), Some(wk), Some(wv), Some(wo)) = (&lw.wq, &lw.wk, &lw.wv, &lw.wo) else {
return cur.to_vec();
};
let n_embd = cfg.n_embd;
let n_head = cfg.n_head;
let n_kv = cfg.n_head_kv;
let hd = cfg.n_embd_head;
let mut q = vec![0.0; n_head * hd];
let mut k = vec![0.0; n_kv * hd];
let mut v = vec![0.0; n_kv * hd];
wq.gemv(cur, &mut q);
wk.gemv(cur, &mut k);
wv.gemv(cur, &mut v);
if let (Some(q_norm), Some(k_norm)) = (&lw.q_norm, &lw.k_norm) {
for h in 0..n_head {
cpu::rmsnorm(&mut q[h * hd..(h + 1) * hd], q_norm, cfg.rms_norm_eps);
apply_rope_neox(&mut q[h * hd..(h + 1) * hd], pos, hd, cfg.rope_freq_base);
}
for h in 0..n_kv {
cpu::rmsnorm(&mut k[h * hd..(h + 1) * hd], k_norm, cfg.rms_norm_eps);
apply_rope_neox(&mut k[h * hd..(h + 1) * hd], pos, hd, cfg.rope_freq_base);
}
}
let (k_cache, v_cache) = kv;
let kv_stride = n_kv * hd;
let write_pos = pos % cfg.swa_window_size;
k_cache[write_pos * kv_stride..(write_pos + 1) * kv_stride].copy_from_slice(&k);
v_cache[write_pos * kv_stride..(write_pos + 1) * kv_stride].copy_from_slice(&v);
let kv_start = (pos + 1).saturating_sub(cfg.swa_window_size);
let kv_len = pos + 1 - kv_start;
let group_size = n_head / n_kv;
let scale = 1.0 / (hd as f32).sqrt();
let mut attn_out = vec![0.0; n_head * hd];
for h in 0..n_head {
let kv_h = h / group_size;
let q_head = &q[h * hd..(h + 1) * hd];
let mut scores = vec![0.0f32; kv_len];
for t in 0..kv_len {
let cache_pos = (kv_start + t) % cfg.swa_window_size;
let k_off = cache_pos * kv_stride + kv_h * hd;
let k_t = &k_cache[k_off..k_off + hd];
scores[t] = q_head.iter().zip(k_t).map(|(a, b)| a * b).sum::<f32>() * scale;
}
cpu::softmax_inplace(&mut scores);
let out = &mut attn_out[h * hd..(h + 1) * hd];
for t in 0..kv_len {
let cache_pos = (kv_start + t) % cfg.swa_window_size;
let v_off = cache_pos * kv_stride + kv_h * hd;
let v_t = &v_cache[v_off..v_off + hd];
let s = scores[t];
for d in 0..hd {
out[d] += s * v_t[d];
}
}
}
let mut proj = vec![0.0; n_embd];
wo.gemv(&attn_out, &mut proj);
proj
}
pub fn detokenize_to_spectrum(
weights: &DetokenizerWeights,
_detok_weights: &AudioDecoderWeights,
state: &mut DetokenizerState,
codes: &[i32],
) -> Vec<f32> {
let cfg = &weights.config;
let n_embd = cfg.n_embd;
let embedding = detok_embed_codes(weights, codes);
let tokens = upsample(&embedding, n_embd, 6);
let n_tokens = tokens.len() / n_embd;
let mut hidden = tokens.clone();
for (il, lw) in weights.layers.iter().enumerate() {
let mut new_hidden = Vec::with_capacity(n_tokens * n_embd);
if cfg.layer_is_conv[il] {
for t in 0..n_tokens {
let mut cur = hidden[t * n_embd..(t + 1) * n_embd].to_vec();
let residual = cur.clone();
cpu::rmsnorm(&mut cur, &lw.operator_norm, cfg.rms_norm_eps);
let out = detok_conv_block(lw, &mut state.conv_bufs[il], &cur, n_embd, cfg.d_conv);
cur = residual.iter().zip(&out).map(|(r, o)| r + o).collect();
let ffn_res = cur.clone();
cpu::rmsnorm(&mut cur, &lw.ffn_norm, cfg.rms_norm_eps);
let mut gate = vec![0.0; cfg.ffn_dim];
let mut up = vec![0.0; cfg.ffn_dim];
lw.ffn_w1.gemv(&cur, &mut gate);
lw.ffn_w3.gemv(&cur, &mut up);
cpu::silu_mul_inplace(&mut gate, &up);
let mut down = vec![0.0; n_embd];
lw.ffn_w2.gemv(&gate, &mut down);
cur = ffn_res.iter().zip(&down).map(|(r, d)| r + d).collect();
new_hidden.extend_from_slice(&cur);
}
} else {
let (Some(wq), Some(wk), Some(wv), Some(wo), Some(q_norm), Some(k_norm), Some(kv)) = (
&lw.wq,
&lw.wk,
&lw.wv,
&lw.wo,
&lw.q_norm,
&lw.k_norm,
state.attn_kv.get_mut(il).and_then(|opt| opt.as_mut()),
) else {
new_hidden.extend_from_slice(&hidden);
continue;
};
let n_head = cfg.n_head;
let n_kv = cfg.n_head_kv;
let hd = cfg.n_embd_head;
let group_size = n_head / n_kv;
let scale = 1.0 / (hd as f32).sqrt();
let kv_stride = n_kv * hd;
let mut all_q = vec![0.0f32; n_tokens * n_head * hd];
for t in 0..n_tokens {
let pos = state.n_past + t;
let mut cur = hidden[t * n_embd..(t + 1) * n_embd].to_vec();
cpu::rmsnorm(&mut cur, &lw.operator_norm, cfg.rms_norm_eps);
let mut q = vec![0.0; n_head * hd];
let mut k = vec![0.0; n_kv * hd];
let mut v = vec![0.0; n_kv * hd];
wq.gemv(&cur, &mut q);
wk.gemv(&cur, &mut k);
wv.gemv(&cur, &mut v);
for h in 0..n_head {
cpu::rmsnorm(&mut q[h * hd..(h + 1) * hd], q_norm, cfg.rms_norm_eps);
apply_rope_neox(&mut q[h * hd..(h + 1) * hd], pos, hd, cfg.rope_freq_base);
}
for h in 0..n_kv {
cpu::rmsnorm(&mut k[h * hd..(h + 1) * hd], k_norm, cfg.rms_norm_eps);
apply_rope_neox(&mut k[h * hd..(h + 1) * hd], pos, hd, cfg.rope_freq_base);
}
let wp = pos % cfg.swa_window_size;
kv.0[wp * kv_stride..(wp + 1) * kv_stride].copy_from_slice(&k);
kv.1[wp * kv_stride..(wp + 1) * kv_stride].copy_from_slice(&v);
all_q[t * n_head * hd..(t + 1) * n_head * hd].copy_from_slice(&q);
}
for t in 0..n_tokens {
let q = &all_q[t * n_head * hd..(t + 1) * n_head * hd];
let kv_end = state.n_past + n_tokens;
let kv_start = kv_end.saturating_sub(cfg.swa_window_size);
let kv_len = kv_end - kv_start;
let mut attn_out = vec![0.0; n_head * hd];
for h in 0..n_head {
let kv_h = h / group_size;
let qh = &q[h * hd..(h + 1) * hd];
let mut scores = vec![0.0f32; kv_len];
for tt in 0..kv_len {
let cp = (kv_start + tt) % cfg.swa_window_size;
let ko = cp * kv_stride + kv_h * hd;
scores[tt] = qh
.iter()
.zip(&kv.0[ko..ko + hd])
.map(|(a, b)| a * b)
.sum::<f32>()
* scale;
}
cpu::softmax_inplace(&mut scores);
let out = &mut attn_out[h * hd..(h + 1) * hd];
for tt in 0..kv_len {
let cp = (kv_start + tt) % cfg.swa_window_size;
let vo = cp * kv_stride + kv_h * hd;
let s = scores[tt];
for d in 0..hd {
out[d] += s * kv.1[vo + d];
}
}
}
let mut proj = vec![0.0; n_embd];
wo.gemv(&attn_out, &mut proj);
let residual = &hidden[t * n_embd..(t + 1) * n_embd];
let mut cur: Vec<f32> = residual.iter().zip(&proj).map(|(r, p)| r + p).collect();
let ffn_res = cur.clone();
cpu::rmsnorm(&mut cur, &lw.ffn_norm, cfg.rms_norm_eps);
let mut gate = vec![0.0; cfg.ffn_dim];
let mut up = vec![0.0; cfg.ffn_dim];
lw.ffn_w1.gemv(&cur, &mut gate);
lw.ffn_w3.gemv(&cur, &mut up);
cpu::silu_mul_inplace(&mut gate, &up);
let mut down = vec![0.0; n_embd];
lw.ffn_w2.gemv(&gate, &mut down);
cur = ffn_res.iter().zip(&down).map(|(r, d)| r + d).collect();
new_hidden.extend_from_slice(&cur);
}
}
hidden = new_hidden;
}
state.n_past += n_tokens;
let mut outputs = Vec::with_capacity(n_tokens * n_embd);
for t in 0..n_tokens {
let mut cur = hidden[t * n_embd..(t + 1) * n_embd].to_vec();
cpu::rmsnorm(&mut cur, &weights.output_norm, cfg.rms_norm_eps);
outputs.extend_from_slice(&cur);
}
let lin_out_dim = weights.lin_w.rows; let mut spectrum = Vec::with_capacity(n_tokens * lin_out_dim);
for t in 0..n_tokens {
let hidden = &outputs[t * n_embd..(t + 1) * n_embd];
let mut frame = vec![0.0; lin_out_dim];
weights.lin_w.gemv(hidden, &mut frame);
for (f, b) in frame.iter_mut().zip(&weights.lin_b) {
*f += b;
}
spectrum.extend_from_slice(&frame);
}
spectrum
}
pub fn istft_to_pcm(spectrum: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> {
let n_fft_bins = n_fft / 2 + 1;
let frame_size = n_fft_bins * 2;
let n_frames = spectrum.len() / frame_size;
if n_frames == 0 {
return vec![];
}
let mut planner = rustfft::FftPlanner::new();
let ifft = planner.plan_fft_inverse(n_fft);
let hann = build_hann(n_fft);
let mut overlap_buf = vec![0.0f32; n_fft];
let mut window_sum = vec![0.0f32; n_fft];
let mut output = Vec::with_capacity(n_frames * hop_length);
use rustfft::num_complex::Complex32;
let mut fft_buf: Vec<Complex32> = vec![Complex32::new(0.0, 0.0); n_fft];
let mut time_domain = vec![0.0f32; n_fft];
for i in 0..n_frames {
for c in fft_buf.iter_mut() {
*c = Complex32::new(0.0, 0.0);
}
for j in 0..n_fft_bins {
let log_abs = spectrum[i * frame_size + j];
let angle = spectrum[i * frame_size + n_fft_bins + j];
let mag = log_abs.exp();
fft_buf[j] = Complex32::new(mag * angle.cos(), mag * angle.sin());
}
for j in 1..n_fft_bins - 1 {
fft_buf[n_fft - j] = Complex32::new(fft_buf[j].re, -fft_buf[j].im);
}
ifft_frame(&mut fft_buf, &mut time_domain, n_fft, ifft.as_ref());
for j in 0..n_fft {
overlap_buf[j] += time_domain[j] * hann[j];
window_sum[j] += hann[j] * hann[j];
}
for k in 0..hop_length {
let sample = if window_sum[k] > 1e-8 {
overlap_buf[k] / window_sum[k]
} else {
overlap_buf[k]
};
output.push(sample);
}
overlap_buf.copy_within(hop_length..n_fft, 0);
overlap_buf[n_fft - hop_length..].fill(0.0);
window_sum.copy_within(hop_length..n_fft, 0);
window_sum[n_fft - hop_length..].fill(0.0);
}
let padding = n_fft.saturating_sub(hop_length) / 2;
if output.len() > padding {
output.drain(..padding);
}
output
}
pub fn build_hann(n_fft: usize) -> Vec<f32> {
(0..n_fft)
.map(|i| 0.5 * (1.0 - (2.0 * std::f32::consts::PI * i as f32 / n_fft as f32).cos()))
.collect()
}
pub fn build_idft_basis(n_fft: usize) -> Vec<f32> {
let bins = n_fft / 2 + 1;
let n = n_fft as f32;
let two_pi = 2.0 * std::f32::consts::PI;
let mut basis = vec![0.0f32; n_fft * 2 * bins];
for t in 0..n_fft {
let row = t * 2 * bins;
for k in 0..bins {
let ang = two_pi * k as f32 * t as f32 / n;
let (re_c, im_c) = if k == 0 {
(1.0 / n, 0.0)
} else if k == bins - 1 {
(ang.cos() / n, 0.0)
} else {
(2.0 * ang.cos() / n, -2.0 * ang.sin() / n)
};
basis[row + k] = re_c;
basis[row + bins + k] = im_c;
}
}
basis
}
fn ifft_frame(
fft_buf: &mut [rustfft::num_complex::Complex32],
out: &mut [f32],
n_fft: usize,
ifft: &dyn rustfft::Fft<f32>,
) {
ifft.process(fft_buf);
let inv_n = 1.0 / n_fft as f32;
for (o, c) in out.iter_mut().zip(fft_buf.iter()) {
*o = c.re * inv_n;
}
}
pub struct IstftStreamer {
n_fft: usize,
hop_length: usize,
n_fft_bins: usize,
frame_size: usize,
hann: Vec<f32>,
overlap_buf: Vec<f32>,
window_sum: Vec<f32>,
padding_remaining: usize,
fft_buf: Vec<rustfft::num_complex::Complex32>,
time_domain: Vec<f32>,
ifft: std::sync::Arc<dyn rustfft::Fft<f32>>,
}
impl IstftStreamer {
pub fn new(n_fft: usize, hop_length: usize) -> Self {
let mut planner = rustfft::FftPlanner::new();
let ifft = planner.plan_fft_inverse(n_fft);
let n_fft_bins = n_fft / 2 + 1;
let frame_size = n_fft_bins * 2;
let hann = build_hann(n_fft);
let padding = n_fft.saturating_sub(hop_length) / 2;
Self {
n_fft,
hop_length,
n_fft_bins,
frame_size,
hann,
overlap_buf: vec![0.0f32; n_fft],
window_sum: vec![0.0f32; n_fft],
padding_remaining: padding,
fft_buf: vec![rustfft::num_complex::Complex32::new(0.0, 0.0); n_fft],
time_domain: vec![0.0f32; n_fft],
ifft,
}
}
pub fn feed_frames(&mut self, spectrum: &[f32]) -> Vec<f32> {
let n_frames = spectrum.len() / self.frame_size;
if n_frames == 0 {
return vec![];
}
let mut output = Vec::with_capacity(n_frames * self.hop_length);
for i in 0..n_frames {
for c in self.fft_buf.iter_mut() {
*c = rustfft::num_complex::Complex32::new(0.0, 0.0);
}
for j in 0..self.n_fft_bins {
let log_abs = spectrum[i * self.frame_size + j];
let angle = spectrum[i * self.frame_size + self.n_fft_bins + j];
let mag = log_abs.exp();
self.fft_buf[j] =
rustfft::num_complex::Complex32::new(mag * angle.cos(), mag * angle.sin());
}
for j in 1..self.n_fft_bins - 1 {
self.fft_buf[self.n_fft - j] =
rustfft::num_complex::Complex32::new(self.fft_buf[j].re, -self.fft_buf[j].im);
}
ifft_frame(
&mut self.fft_buf,
&mut self.time_domain,
self.n_fft,
self.ifft.as_ref(),
);
for j in 0..self.n_fft {
self.overlap_buf[j] += self.time_domain[j] * self.hann[j];
self.window_sum[j] += self.hann[j] * self.hann[j];
}
for k in 0..self.hop_length {
let sample = if self.window_sum[k] > 1e-8 {
self.overlap_buf[k] / self.window_sum[k]
} else {
self.overlap_buf[k]
};
if self.padding_remaining > 0 {
self.padding_remaining -= 1;
} else {
output.push(sample);
}
}
self.overlap_buf.copy_within(self.hop_length..self.n_fft, 0);
self.overlap_buf[self.n_fft - self.hop_length..].fill(0.0);
self.window_sum.copy_within(self.hop_length..self.n_fft, 0);
self.window_sum[self.n_fft - self.hop_length..].fill(0.0);
}
output
}
pub fn flush(&mut self) -> Vec<f32> {
self.overlap_buf.fill(0.0);
self.window_sum.fill(0.0);
vec![]
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_istft_streamer_exact_parity_with_batch() {
let n_fft = 1280;
let hop_length = 320;
let n_fft_bins = n_fft / 2 + 1;
let frame_size = n_fft_bins * 2;
let n_frames = 18;
let mut spectrum = vec![0.0f32; n_frames * frame_size];
for i in 0..spectrum.len() {
spectrum[i] = ((i as f32) * 0.01).sin();
}
let batch_pcm = istft_to_pcm(&spectrum, n_fft, hop_length);
let mut streamer = IstftStreamer::new(n_fft, hop_length);
let mut streamed_pcm = Vec::new();
for chunk in spectrum.chunks(6 * frame_size) {
let pcm_chunk = streamer.feed_frames(chunk);
streamed_pcm.extend_from_slice(&pcm_chunk);
}
assert_eq!(batch_pcm.len(), streamed_pcm.len());
for (idx, (&b, &s)) in batch_pcm.iter().zip(streamed_pcm.iter()).enumerate() {
assert!(
(b - s).abs() < 1e-6,
"Sample mismatch at {idx}: batch={b}, stream={s}"
);
}
}
#[test]
fn test_build_gpu_audio_decoder_empty_gguf_graceful_none() {
let mut data = Vec::new();
data.extend_from_slice(b"GGUF");
data.extend_from_slice(&3u32.to_le_bytes());
data.extend_from_slice(&0u64.to_le_bytes());
data.extend_from_slice(&0u64.to_le_bytes());
let bytes: Arc<[u8]> = Arc::from(data.into_boxed_slice());
let gguf = Arc::new(GgufFile::from_bytes(bytes).expect("parse minimal gguf"));
assert!(build_gpu_audio_decoder(&gguf, crate::engine::BackendPreference::Cpu).is_none());
assert!(build_gpu_audio_decoder(&gguf, crate::engine::BackendPreference::Metal).is_none());
assert!(build_gpu_audio_decoder(&gguf, crate::engine::BackendPreference::Gpu).is_none());
assert!(build_gpu_audio_decoder(&gguf, crate::engine::BackendPreference::Auto).is_none());
}
struct MockAudioGpu {
active: std::sync::atomic::AtomicBool,
}
impl AudioGpu for MockAudioGpu {
fn supports_depthformer(&self) -> bool {
false
}
fn sample_audio_frame(
&self,
_embedding: &[f32],
_temperature: f32,
_top_k: usize,
) -> [i32; 8] {
[0; 8]
}
fn detokenize_to_spectrum(
&self,
_cpu_weights: &DetokenizerWeights,
_codes: &[i32],
) -> Vec<f32> {
vec![]
}
fn reset_depthformer(&self) {}
fn reset_detokenizer(&self) {}
fn try_acquire_session(&self) -> bool {
self.active
.compare_exchange(
false,
true,
std::sync::atomic::Ordering::Acquire,
std::sync::atomic::Ordering::Relaxed,
)
.is_ok()
}
fn release_session(&self) {
self.active
.store(false, std::sync::atomic::Ordering::Release);
}
}
#[test]
fn test_audio_gpu_session_lease_lifecycle() {
let gpu = MockAudioGpu {
active: std::sync::atomic::AtomicBool::new(false),
};
assert!(gpu.try_acquire_session());
assert!(!gpu.try_acquire_session());
gpu.release_session();
assert!(gpu.try_acquire_session());
gpu.release_session();
}
}