use anyhow::{Context, Result};
use rayon::prelude::*;
use std::collections::BTreeMap;
use std::path::Path;
pub mod cfg {
pub const T3_HIDDEN: usize = 1024;
pub const T3_INTERMEDIATE: usize = 4096;
pub const T3_LAYERS: usize = 30;
pub const T3_HEADS: usize = 16;
pub const T3_HEAD_DIM: usize = T3_HIDDEN / T3_HEADS;
pub const T3_TEXT_VOCAB_MTL: usize = 2454;
pub const T3_TEXT_VOCAB_EN: usize = 704;
pub const T3_SPEECH_VOCAB: usize = 8194;
pub const T3_MAX_TEXT_TOKENS: usize = 2048;
pub const T3_MAX_SPEECH_TOKENS: usize = 4096;
pub const START_TEXT_TOKEN: u32 = 255;
pub const STOP_TEXT_TOKEN: u32 = 0;
pub const START_SPEECH_TOKEN: u32 = 6561;
pub const STOP_SPEECH_TOKEN: u32 = 6562;
pub const T3_LEARNED_POS: bool = true;
pub const VE_EMBED_DIM: usize = 256;
pub const PERCEIVER_QUERIES: usize = 32;
pub const PERCEIVER_DIM: usize = T3_HIDDEN;
pub const PERCEIVER_HEADS: usize = 4;
pub const DEFAULT_EMOTION_ADV: f32 = 0.5;
pub const ENC_COND_SECS: usize = 6;
pub const DEC_COND_SECS: usize = 10;
pub const SPEECH_COND_PROMPT_LEN: usize = 150;
pub const CFG_WEIGHT: f32 = 0.5;
pub const SAMPLE_TEMPERATURE: f32 = 0.8;
pub const SAMPLE_TOP_P: f32 = 1.0;
pub const SAMPLE_MIN_P: f32 = 0.05;
pub const SAMPLE_REPETITION_PENALTY: f32 = 1.2;
pub const S3_SR: u32 = 16_000;
pub const S3_SPEECH_VOCAB: usize = 6561;
pub const S3_TOKEN_RATE: u32 = 25; pub const S3_TOKEN_HOP: usize = 640; pub const S3_N_FFT: usize = 400;
pub const S3_HOP: usize = 160;
pub const S3_MEL_PER_TOKEN: usize = 4;
pub const S3GEN_SR: u32 = 24_000;
pub const S3GEN_MELS: usize = 80;
pub const CONFORMER_DIM: usize = 512;
pub const CONFORMER_HEADS: usize = 8;
pub const CONFORMER_BLOCKS: usize = 6;
pub const CFM_TIMESTEPS: usize = 10;
pub const CFM_INFERENCE_CFG_RATE: f32 = 0.7;
pub const S3GEN_SPK_EMBED_DIM: usize = 192;
pub const S3GEN_TOKEN_MEL_RATIO: usize = 2;
pub const DEC_IN_CHANNELS: usize = 320; pub const DEC_CHANNELS: usize = 256;
pub const DEC_TIME_DIM: usize = 1024; pub const DEC_HEADS: usize = 8;
pub const DEC_HEAD_DIM: usize = 64; pub const DEC_N_BLOCKS: usize = 4; pub const DEC_MID_BLOCKS: usize = 12;
pub const CONFORMER_UNITS: usize = 2048; pub const UP_CONFORMER_BLOCKS: usize = 4; pub const CAMP_EMBED: usize = 192;
pub const CAMP_GROWTH: usize = 32; pub const CAMP_INIT: usize = 128; pub const CAMP_BN: usize = 128; pub const CAMP_BLOCKS: [(usize, usize); 3] = [(12, 1), (24, 2), (16, 2)];
pub const HIFT_NFFT: usize = 16;
pub const HIFT_HOP: usize = 4;
pub const HIFT_BASE: usize = 512;
pub const HIFT_UP_RATES: [usize; 3] = [8, 5, 3];
pub const HIFT_UP_KERNELS: [usize; 3] = [16, 11, 7];
pub const HIFT_HARMONICS: usize = 9;
pub const HIFT_UPSAMPLE: [usize; 3] = [8, 5, 3];
pub const S3GEN_SIL: u32 = 4299;
pub const VE_NUM_MELS: usize = 40;
pub const VE_HIDDEN: usize = 256;
pub const VE_SR: u32 = 16_000;
pub const VE_PARTIAL_FRAMES: usize = 160;
pub const VE_N_FFT: usize = 400;
pub const VE_HOP: usize = 160;
pub const VE_WIN: usize = 400;
pub const VE_FMAX: f32 = 8000.0;
pub const FILE_VE: &str = "ve.safetensors";
pub const FILE_T3_MTL: &str = "t3_mtl23ls_v3.safetensors";
pub const FILE_T3_EN: &str = "t3_cfg.safetensors";
pub const FILE_S3GEN_MTL: &str = "s3gen_v3.safetensors";
pub const FILE_S3GEN_EN: &str = "s3gen.safetensors";
pub const FILE_TOKENIZER_MTL: &str = "mtl_tokenizer.json";
pub const FILE_TOKENIZER_EN: &str = "tokenizer.json";
pub const FILE_CONDS: &str = "conds.pt";
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TensorInfo {
pub name: String,
pub dtype: String,
pub shape: Vec<usize>,
}
impl TensorInfo {
pub fn numel(&self) -> usize {
self.shape.iter().product::<usize>().max(1)
}
}
pub fn read_manifest(path: &Path) -> Result<Vec<TensorInfo>> {
use std::io::Read;
let mut f = std::fs::File::open(path)
.with_context(|| format!("open safetensors {}", path.display()))?;
let mut len8 = [0u8; 8];
f.read_exact(&mut len8)
.with_context(|| format!("read header length from {}", path.display()))?;
let hlen = u64::from_le_bytes(len8);
anyhow::ensure!(
hlen > 0 && hlen < 512 * 1024 * 1024,
"implausible safetensors header length {hlen} in {}",
path.display()
);
let mut hdr = vec![0u8; hlen as usize];
f.read_exact(&mut hdr)
.with_context(|| format!("read {hlen}-byte header from {}", path.display()))?;
let v: serde_json::Value = serde_json::from_slice(&hdr)
.with_context(|| format!("parse header of {}", path.display()))?;
let obj = v
.as_object()
.context("safetensors header is not a JSON object")?;
let mut out = Vec::with_capacity(obj.len());
for (name, meta) in obj {
if name == "__metadata__" {
continue;
}
let dtype = meta
.get("dtype")
.and_then(|d| d.as_str())
.unwrap_or("?")
.to_string();
let shape: Vec<usize> = meta
.get("shape")
.and_then(|s| s.as_array())
.map(|a| {
a.iter()
.filter_map(|x| x.as_u64().map(|x| x as usize))
.collect()
})
.unwrap_or_default();
out.push(TensorInfo {
name: name.clone(),
dtype,
shape,
});
}
out.sort_by(|a, b| a.name.cmp(&b.name));
Ok(out)
}
pub fn verify_shapes(path: &Path, expected: &[(&str, &[usize])]) -> Result<()> {
let manifest = read_manifest(path)?;
let by_name: BTreeMap<&str, &TensorInfo> =
manifest.iter().map(|t| (t.name.as_str(), t)).collect();
let mut problems = Vec::new();
for (name, want) in expected {
match by_name.get(name) {
None => problems.push(format!(" MISSING {name} (want {want:?})")),
Some(t) => {
let ok = t.shape.len() == want.len()
&& t.shape
.iter()
.zip(want.iter())
.all(|(g, w)| *w == 0 || g == w);
if !ok {
problems.push(format!(
" SHAPE {name}: got {:?}, want {want:?}",
t.shape
));
}
}
}
}
anyhow::ensure!(
problems.is_empty(),
"checkpoint {} failed shape verification:\n{}",
path.display(),
problems.join("\n")
);
Ok(())
}
struct StTensor {
dtype: String,
shape: Vec<usize>,
range: (u64, u64),
}
pub struct StReader {
file: std::fs::File,
map: std::collections::HashMap<String, StTensor>,
}
impl StReader {
pub fn open(path: &Path) -> Result<Self> {
use std::io::Read;
let mut file =
std::fs::File::open(path).with_context(|| format!("open {}", path.display()))?;
let mut len8 = [0u8; 8];
file.read_exact(&mut len8)?;
let hlen = u64::from_le_bytes(len8);
anyhow::ensure!(
hlen > 0 && hlen < 512 * 1024 * 1024,
"bad header len {hlen}"
);
let mut hdr = vec![0u8; hlen as usize];
file.read_exact(&mut hdr)?;
let v: serde_json::Value = serde_json::from_slice(&hdr)?;
let obj = v.as_object().context("safetensors header")?;
let base = 8 + hlen;
let mut map = std::collections::HashMap::new();
for (name, meta) in obj {
if name == "__metadata__" {
continue;
}
let dtype = meta
.get("dtype")
.and_then(|d| d.as_str())
.unwrap_or("?")
.to_string();
let shape: Vec<usize> = meta
.get("shape")
.and_then(|s| s.as_array())
.map(|a| {
a.iter()
.filter_map(|x| x.as_u64().map(|x| x as usize))
.collect()
})
.unwrap_or_default();
let offs = meta
.get("data_offsets")
.and_then(|o| o.as_array())
.context("data_offsets")?;
let s0 = offs[0].as_u64().context("off0")?;
let s1 = offs[1].as_u64().context("off1")?;
map.insert(
name.clone(),
StTensor {
dtype,
shape,
range: (base + s0, base + s1),
},
);
}
Ok(Self { file, map })
}
fn loc(&self, name: &str) -> Result<&StTensor> {
self.map
.get(name)
.with_context(|| format!("missing tensor {name}"))
}
pub fn shape(&self, name: &str) -> Result<&[usize]> {
Ok(&self.loc(name)?.shape)
}
pub fn f32(&self, name: &str) -> Result<Vec<f32>> {
let t = self.loc(name)?;
let (a, b) = t.range;
let mut raw = vec![0u8; (b - a) as usize];
pread(&self.file, &mut raw, a).with_context(|| format!("pread {name}"))?;
Ok(match t.dtype.as_str() {
"F32" => bytemuck::cast_slice::<u8, f32>(&raw).to_vec(),
"F16" => raw
.chunks_exact(2)
.map(|c| half::f16::from_le_bytes([c[0], c[1]]).to_f32())
.collect(),
"BF16" => raw
.chunks_exact(2)
.map(|c| half::bf16::from_le_bytes([c[0], c[1]]).to_f32())
.collect(),
other => anyhow::bail!("tensor {name}: unsupported dtype {other}"),
})
}
fn mat(&self, name: &str, rows: usize, cols: usize) -> Result<Vec<f32>> {
let s = self.shape(name)?;
anyhow::ensure!(s == [rows, cols], "{name}: shape {s:?} != [{rows}, {cols}]");
self.f32(name)
}
}
fn pread(file: &std::fs::File, buf: &mut [u8], offset: u64) -> std::io::Result<()> {
#[cfg(unix)]
{
use std::os::unix::fs::FileExt;
file.read_exact_at(buf, offset)
}
#[cfg(not(unix))]
{
use std::io::{Read, Seek, SeekFrom};
let mut f = file.try_clone()?;
f.seek(SeekFrom::Start(offset))?;
f.read_exact(buf)
}
}
const MATVEC_PAR_THRESHOLD: usize = 1 << 20;
fn matvec(w: &[f32], x: &[f32], out: usize, cols: usize) -> Vec<f32> {
debug_assert_eq!(w.len(), out * cols);
debug_assert_eq!(x.len(), cols);
let dot = |o: usize| -> f32 {
let row = &w[o * cols..o * cols + cols];
row.iter().zip(x).map(|(a, b)| a * b).sum()
};
if out * cols >= MATVEC_PAR_THRESHOLD {
(0..out).into_par_iter().map(dot).collect()
} else {
(0..out).map(dot).collect()
}
}
fn silu(v: f32) -> f32 {
v / (1.0 + (-v).exp())
}
fn rms_norm(x: &[f32], w: &[f32], eps: f32) -> Vec<f32> {
let ms = x.iter().map(|v| v * v).sum::<f32>() / x.len() as f32;
let inv = 1.0 / (ms + eps).sqrt();
x.iter().zip(w).map(|(v, g)| v * inv * g).collect()
}
fn softmax_inplace(v: &mut [f32]) {
let m = v.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut s = 0.0;
for x in v.iter_mut() {
*x = (*x - m).exp();
s += *x;
}
for x in v.iter_mut() {
*x /= s;
}
}
struct T3Layer {
input_ln: Vec<f32>,
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
o: Vec<f32>,
post_ln: Vec<f32>,
gate: Vec<f32>,
up: Vec<f32>,
down: Vec<f32>,
}
pub struct T3 {
layers: Vec<T3Layer>,
final_norm: Vec<f32>,
text_emb: Vec<f32>, speech_emb: Vec<f32>, text_head: Vec<f32>, speech_head: Vec<f32>, text_pos: Vec<f32>, speech_pos: Vec<f32>, inv_freq: Vec<f32>,
text_vocab: usize,
spkr_w: Vec<f32>, spkr_b: Vec<f32>, emo_w: Vec<f32>, perceiver: Perceiver,
}
impl T3 {
pub fn load(path: &Path) -> Result<Self> {
let st = StReader::open(path).with_context(|| format!("open T3 {}", path.display()))?;
let (h, i, tv, sv) = (
cfg::T3_HIDDEN,
cfg::T3_INTERMEDIATE,
cfg::T3_TEXT_VOCAB_MTL,
cfg::T3_SPEECH_VOCAB,
);
let mut layers = Vec::with_capacity(cfg::T3_LAYERS);
for l in 0..cfg::T3_LAYERS {
let p = format!("tfmr.layers.{l}");
layers.push(T3Layer {
input_ln: st.f32(&format!("{p}.input_layernorm.weight"))?,
q: st.mat(&format!("{p}.self_attn.q_proj.weight"), h, h)?,
k: st.mat(&format!("{p}.self_attn.k_proj.weight"), h, h)?,
v: st.mat(&format!("{p}.self_attn.v_proj.weight"), h, h)?,
o: st.mat(&format!("{p}.self_attn.o_proj.weight"), h, h)?,
post_ln: st.f32(&format!("{p}.post_attention_layernorm.weight"))?,
gate: st.mat(&format!("{p}.mlp.gate_proj.weight"), i, h)?,
up: st.mat(&format!("{p}.mlp.up_proj.weight"), i, h)?,
down: st.mat(&format!("{p}.mlp.down_proj.weight"), h, i)?,
});
}
Ok(Self {
layers,
final_norm: st.f32("tfmr.norm.weight")?,
text_emb: st.mat("text_emb.weight", tv, h)?,
speech_emb: st.mat("speech_emb.weight", sv, h)?,
text_head: st.mat("text_head.weight", tv, h)?,
speech_head: st.mat("speech_head.weight", sv, h)?,
text_pos: st.f32("text_pos_emb.emb.weight")?,
speech_pos: st.f32("speech_pos_emb.emb.weight")?,
inv_freq: llama3_inv_freq(cfg::T3_HEAD_DIM),
text_vocab: tv,
spkr_w: st.mat("cond_enc.spkr_enc.weight", h, cfg::VE_EMBED_DIM)?,
spkr_b: st.f32("cond_enc.spkr_enc.bias")?,
emo_w: st.f32("cond_enc.emotion_adv_fc.weight")?,
perceiver: Perceiver::load(&st, "cond_enc.perceiver")?,
})
}
fn embed_row(table: &[f32], idx: usize) -> &[f32] {
&table[idx * cfg::T3_HIDDEN..(idx + 1) * cfg::T3_HIDDEN]
}
pub fn build_inputs_embeds(
&self,
cond: &[Vec<f32>],
text_tokens: &[u32],
speech_tokens: &[u32],
) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let mut seq =
Vec::with_capacity((cond.len() + text_tokens.len() + speech_tokens.len()) * h);
for c in cond {
debug_assert_eq!(c.len(), h);
seq.extend_from_slice(c);
}
for (pos, &t) in text_tokens.iter().enumerate() {
let e = Self::embed_row(&self.text_emb, t as usize);
let p = &self.text_pos[pos * h..pos * h + h];
seq.extend(e.iter().zip(p).map(|(a, b)| a + b));
}
for (pos, &t) in speech_tokens.iter().enumerate() {
let e = Self::embed_row(&self.speech_emb, t as usize);
let p = &self.speech_pos[pos * h..pos * h + h];
seq.extend(e.iter().zip(p).map(|(a, b)| a + b));
}
seq
}
pub fn forward_hidden(&self, inputs_embeds: &[f32]) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let seq = inputs_embeds.len() / h;
let (nh, hd) = (cfg::T3_HEADS, cfg::T3_HEAD_DIM);
let scale = 1.0 / (hd as f32).sqrt();
let (cos, sin) = self.rope_tables(seq);
let mut x = inputs_embeds.to_vec();
for layer in &self.layers {
let normed: Vec<f32> = (0..seq)
.flat_map(|s| rms_norm(&x[s * h..s * h + h], &layer.input_ln, 1e-5))
.collect();
let mut q = vec![0f32; seq * h];
let mut k = vec![0f32; seq * h];
let mut v = vec![0f32; seq * h];
for s in 0..seq {
let xn = &normed[s * h..s * h + h];
let qs = matvec(&layer.q, xn, h, h);
let ks = matvec(&layer.k, xn, h, h);
let vs = matvec(&layer.v, xn, h, h);
q[s * h..s * h + h].copy_from_slice(&qs);
k[s * h..s * h + h].copy_from_slice(&ks);
v[s * h..s * h + h].copy_from_slice(&vs);
apply_rope(&mut q[s * h..s * h + h], &cos[s], &sin[s], nh, hd);
apply_rope(&mut k[s * h..s * h + h], &cos[s], &sin[s], nh, hd);
}
let mut attn = vec![0f32; seq * h];
for head in 0..nh {
let off = head * hd;
for s in 0..seq {
let qh = &q[s * h + off..s * h + off + hd];
let mut scores = vec![0f32; s + 1];
for (t, sc) in scores.iter_mut().enumerate() {
let kh = &k[t * h + off..t * h + off + hd];
*sc = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_inplace(&mut scores);
let out = &mut attn[s * h + off..s * h + off + hd];
for (t, &w) in scores.iter().enumerate() {
let vh = &v[t * h + off..t * h + off + hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
}
for s in 0..seq {
let o = matvec(&layer.o, &attn[s * h..s * h + h], h, h);
for d in 0..h {
x[s * h + d] += o[d];
}
}
for s in 0..seq {
let xn = rms_norm(&x[s * h..s * h + h], &layer.post_ln, 1e-5);
let g = matvec(&layer.gate, &xn, cfg::T3_INTERMEDIATE, h);
let u = matvec(&layer.up, &xn, cfg::T3_INTERMEDIATE, h);
let act: Vec<f32> = g.iter().zip(&u).map(|(gv, uv)| silu(*gv) * uv).collect();
let d = matvec(&layer.down, &act, h, cfg::T3_INTERMEDIATE);
for j in 0..h {
x[s * h + j] += d[j];
}
}
}
(0..seq)
.flat_map(|s| rms_norm(&x[s * h..s * h + h], &self.final_norm, 1e-5))
.collect()
}
pub fn speech_logits_last(&self, hidden: &[f32]) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let last = &hidden[hidden.len() - h..];
matvec(&self.speech_head, last, cfg::T3_SPEECH_VOCAB, h)
}
pub fn text_logits_at(&self, hidden: &[f32], pos: usize) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
matvec(
&self.text_head,
&hidden[pos * h..pos * h + h],
self.text_vocab,
h,
)
}
pub fn build_cond(
&self,
ve_embed: &[f32],
prompt_speech_tokens: &[u32],
emotion_adv: f32,
) -> Vec<Vec<f32>> {
let h = cfg::T3_HIDDEN;
let spkr = add_bias(
matvec(&self.spkr_w, ve_embed, h, cfg::VE_EMBED_DIM),
&self.spkr_b,
);
let n = prompt_speech_tokens.len();
let mut prompt = Vec::with_capacity(n * h);
for (pos, &t) in prompt_speech_tokens.iter().enumerate() {
let e = Self::embed_row(&self.speech_emb, t as usize);
let p = &self.speech_pos[pos * h..pos * h + h];
prompt.extend(e.iter().zip(p).map(|(a, b)| a + b));
}
let resampled = self.perceiver.forward(&prompt, n); let emo: Vec<f32> = self.emo_w.iter().map(|w| w * emotion_adv).collect();
let mut rows = Vec::with_capacity(2 + cfg::PERCEIVER_QUERIES);
rows.push(spkr);
for i in 0..cfg::PERCEIVER_QUERIES {
rows.push(resampled[i * h..i * h + h].to_vec());
}
rows.push(emo);
rows
}
fn build_inputs_embeds_cfg(
&self,
cond: &[Vec<f32>],
text_tokens: &[u32],
speech_tokens: &[u32],
zero_text: bool,
) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let mut seq =
Vec::with_capacity((cond.len() + text_tokens.len() + speech_tokens.len()) * h);
for c in cond {
seq.extend_from_slice(c);
}
for (pos, &t) in text_tokens.iter().enumerate() {
let p = &self.text_pos[pos * h..pos * h + h];
if zero_text {
seq.extend_from_slice(p); } else {
let e = Self::embed_row(&self.text_emb, t as usize);
seq.extend(e.iter().zip(p).map(|(a, b)| a + b));
}
}
for (pos, &t) in speech_tokens.iter().enumerate() {
let e = Self::embed_row(&self.speech_emb, t as usize);
let p = &self.speech_pos[pos * h..pos * h + h];
seq.extend(e.iter().zip(p).map(|(a, b)| a + b));
}
seq
}
pub fn generate(
&self,
cond: &[Vec<f32>],
text_tokens: &[u32],
max_new_tokens: usize,
seed: u64,
) -> SpeechGen {
let mut speech = vec![cfg::START_SPEECH_TOKEN];
let mut out = Vec::new();
let mut rng = SplitMix64::new(seed);
let mut stopped = false;
for _ in 0..max_new_tokens {
let cl = self.speech_logits_last(&self.forward_hidden(&self.build_inputs_embeds_cfg(
cond,
text_tokens,
&speech,
false,
)));
let ul = self.speech_logits_last(&self.forward_hidden(&self.build_inputs_embeds_cfg(
cond,
text_tokens,
&speech,
true,
)));
let mut logits: Vec<f32> = cl
.iter()
.zip(&ul)
.map(|(c, u)| c + cfg::CFG_WEIGHT * (c - u))
.collect();
for (i, v) in logits.iter_mut().enumerate() {
if i >= cfg::S3_SPEECH_VOCAB && i != cfg::STOP_SPEECH_TOKEN as usize {
*v = f32::NEG_INFINITY;
}
}
apply_repetition_penalty(&mut logits, &speech, cfg::SAMPLE_REPETITION_PENALTY);
if cfg::SAMPLE_TEMPERATURE != 1.0 {
for v in logits.iter_mut() {
*v /= cfg::SAMPLE_TEMPERATURE;
}
}
apply_min_p(&mut logits, cfg::SAMPLE_MIN_P);
apply_top_p(&mut logits, cfg::SAMPLE_TOP_P);
let next = sample_multinomial(&logits, &mut rng);
if next == cfg::STOP_SPEECH_TOKEN {
stopped = true;
break;
}
out.push(next);
speech.push(next);
}
SpeechGen {
tokens: out,
stopped,
}
}
fn forward_cached(
&self,
embeds: &[f32],
pos0: usize,
cache: &mut [(Vec<f32>, Vec<f32>)],
) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let n = embeds.len() / h;
let (nh, hd) = (cfg::T3_HEADS, cfg::T3_HEAD_DIM);
let scale = 1.0 / (hd as f32).sqrt();
let cs: Vec<(Vec<f32>, Vec<f32>)> = (0..n)
.map(|i| {
let p = (pos0 + i) as f32;
let (mut cos, mut sin) = (vec![0f32; hd], vec![0f32; hd]);
for (j, &f) in self.inv_freq.iter().enumerate() {
let (s, c) = (p * f).sin_cos();
cos[j] = c;
cos[j + hd / 2] = c;
sin[j] = s;
sin[j + hd / 2] = s;
}
(cos, sin)
})
.collect();
let mut x = embeds.to_vec();
for (li, layer) in self.layers.iter().enumerate() {
let mut q = vec![0f32; n * h];
for s in 0..n {
let xn = rms_norm(&x[s * h..s * h + h], &layer.input_ln, 1e-5);
let mut qs = matvec(&layer.q, &xn, h, h);
let mut ks = matvec(&layer.k, &xn, h, h);
let vs = matvec(&layer.v, &xn, h, h);
apply_rope(&mut qs, &cs[s].0, &cs[s].1, nh, hd);
apply_rope(&mut ks, &cs[s].0, &cs[s].1, nh, hd);
q[s * h..s * h + h].copy_from_slice(&qs);
cache[li].0.extend_from_slice(&ks);
cache[li].1.extend_from_slice(&vs);
}
let mut attn = vec![0f32; n * h];
for head in 0..nh {
let off = head * hd;
for s in 0..n {
let abs = pos0 + s;
let qh = &q[s * h + off..s * h + off + hd];
let mut scores = vec![0f32; abs + 1];
for (t, sc) in scores.iter_mut().enumerate() {
let kh = &cache[li].0[t * h + off..t * h + off + hd];
*sc = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_inplace(&mut scores);
let out = &mut attn[s * h + off..s * h + off + hd];
for (t, &w) in scores.iter().enumerate() {
let vh = &cache[li].1[t * h + off..t * h + off + hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
}
for s in 0..n {
let o = matvec(&layer.o, &attn[s * h..s * h + h], h, h);
for d in 0..h {
x[s * h + d] += o[d];
}
}
for s in 0..n {
let xn = rms_norm(&x[s * h..s * h + h], &layer.post_ln, 1e-5);
let g = matvec(&layer.gate, &xn, cfg::T3_INTERMEDIATE, h);
let u = matvec(&layer.up, &xn, cfg::T3_INTERMEDIATE, h);
let act: Vec<f32> = g.iter().zip(&u).map(|(gv, uv)| silu(*gv) * uv).collect();
let dn = matvec(&layer.down, &act, h, cfg::T3_INTERMEDIATE);
for j in 0..h {
x[s * h + j] += dn[j];
}
}
}
(0..n)
.flat_map(|s| rms_norm(&x[s * h..s * h + h], &self.final_norm, 1e-5))
.collect()
}
pub fn generate_fast(
&self,
cond: &[Vec<f32>],
text_tokens: &[u32],
max_new_tokens: usize,
seed: u64,
) -> SpeechGen {
let h = cfg::T3_HIDDEN;
let nl = self.layers.len();
let mut cache_c: Vec<(Vec<f32>, Vec<f32>)> = vec![(Vec::new(), Vec::new()); nl];
let mut cache_u = cache_c.clone();
let bos = [cfg::START_SPEECH_TOKEN];
let hc = self.forward_cached(
&self.build_inputs_embeds_cfg(cond, text_tokens, &bos, false),
0,
&mut cache_c,
);
let hu = self.forward_cached(
&self.build_inputs_embeds_cfg(cond, text_tokens, &bos, true),
0,
&mut cache_u,
);
let mut last_c = hc[hc.len() - h..].to_vec();
let mut last_u = hu[hu.len() - h..].to_vec();
let mut abs_pos = cache_c[0].0.len() / h; let mut speech_seen = vec![cfg::START_SPEECH_TOKEN];
let mut out = Vec::new();
let mut rng = SplitMix64::new(seed);
let mut stopped = false;
for step in 0..max_new_tokens {
let cl = self.speech_logits_last(&last_c);
let ul = self.speech_logits_last(&last_u);
let mut logits: Vec<f32> = cl
.iter()
.zip(&ul)
.map(|(c, u)| c + cfg::CFG_WEIGHT * (c - u))
.collect();
for (i, v) in logits.iter_mut().enumerate() {
if i >= cfg::S3_SPEECH_VOCAB && i != cfg::STOP_SPEECH_TOKEN as usize {
*v = f32::NEG_INFINITY;
}
}
if std::env::var("CB_STOPTRACE").is_ok() {
let s = cfg::STOP_SPEECH_TOKEN as usize;
let (mut best_i, mut best_v) = (0usize, f32::NEG_INFINITY);
for (i, &v) in logits.iter().enumerate().take(cfg::S3_SPEECH_VOCAB) {
if v > best_v {
best_v = v;
best_i = i;
}
}
eprintln!(
" step {step:>3}: STOP cond {:>8.3} uncond {:>8.3} blend {:>8.3} | best tok {best_i:>5} blend {best_v:>8.3} | gap {:>8.3}",
cl[s],
ul[s],
logits[s],
logits[s] - best_v
);
}
apply_repetition_penalty(&mut logits, &speech_seen, cfg::SAMPLE_REPETITION_PENALTY);
if cfg::SAMPLE_TEMPERATURE != 1.0 {
for v in logits.iter_mut() {
*v /= cfg::SAMPLE_TEMPERATURE;
}
}
apply_min_p(&mut logits, cfg::SAMPLE_MIN_P);
apply_top_p(&mut logits, cfg::SAMPLE_TOP_P);
let next = sample_multinomial(&logits, &mut rng);
if next == cfg::STOP_SPEECH_TOKEN {
stopped = true;
break;
}
out.push(next);
speech_seen.push(next);
let sidx = step + 1;
let e = Self::embed_row(&self.speech_emb, next as usize);
let p = &self.speech_pos[sidx * h..sidx * h + h];
let emb: Vec<f32> = e.iter().zip(p).map(|(a, b)| a + b).collect();
last_c = self.forward_cached(&emb, abs_pos, &mut cache_c);
last_u = self.forward_cached(&emb, abs_pos, &mut cache_u);
abs_pos += 1;
}
SpeechGen {
tokens: out,
stopped,
}
}
fn rope_tables(&self, seq: usize) -> (Vec<Vec<f32>>, Vec<Vec<f32>>) {
let hd = cfg::T3_HEAD_DIM;
let mut cos = vec![vec![0f32; hd]; seq];
let mut sin = vec![vec![0f32; hd]; seq];
for p in 0..seq {
for (j, &f) in self.inv_freq.iter().enumerate() {
let a = p as f32 * f;
let (s, c) = a.sin_cos();
cos[p][j] = c;
cos[p][j + hd / 2] = c;
sin[p][j] = s;
sin[p][j + hd / 2] = s;
}
}
(cos, sin)
}
}
pub struct SpeechGen {
pub tokens: Vec<u32>,
pub stopped: bool,
}
pub struct Perceiver {
query: Vec<f32>, norm_w: Vec<f32>,
norm_b: Vec<f32>,
to_q: Vec<f32>,
to_q_b: Vec<f32>,
to_k: Vec<f32>,
to_k_b: Vec<f32>,
to_v: Vec<f32>,
to_v_b: Vec<f32>,
proj: Vec<f32>,
proj_b: Vec<f32>,
}
impl Perceiver {
pub fn load(st: &StReader, p: &str) -> Result<Self> {
let h = cfg::T3_HIDDEN;
Ok(Self {
query: st.f32(&format!("{p}.pre_attention_query"))?, norm_w: st.f32(&format!("{p}.attn.norm.weight"))?,
norm_b: st.f32(&format!("{p}.attn.norm.bias"))?,
to_q: st.mat(&format!("{p}.attn.to_q.weight"), h, h)?,
to_q_b: st.f32(&format!("{p}.attn.to_q.bias"))?,
to_k: st.mat(&format!("{p}.attn.to_k.weight"), h, h)?,
to_k_b: st.f32(&format!("{p}.attn.to_k.bias"))?,
to_v: st.mat(&format!("{p}.attn.to_v.weight"), h, h)?,
to_v_b: st.f32(&format!("{p}.attn.to_v.bias"))?,
proj: st.mat(&format!("{p}.attn.proj_out.weight"), h, h)?,
proj_b: st.f32(&format!("{p}.attn.proj_out.bias"))?,
})
}
fn attn_block(&self, x1: &[f32], lq: usize, x2: &[f32], lk: usize) -> Vec<f32> {
let h = cfg::T3_HIDDEN;
let nh = cfg::PERCEIVER_HEADS;
let hd = h / nh;
let scale = 1.0 / (hd as f32).sqrt();
let normed = |src: &[f32], len: usize| -> Vec<f32> {
(0..len)
.flat_map(|i| layer_norm(&src[i * h..i * h + h], &self.norm_w, &self.norm_b, 1e-5))
.collect()
};
let proj_seq = |w: &[f32], b: &[f32], src: &[f32], len: usize| -> Vec<f32> {
(0..len)
.flat_map(|i| add_bias(matvec(w, &src[i * h..i * h + h], h, h), b))
.collect()
};
let x1n = normed(x1, lq);
let x2n = normed(x2, lk);
let q = proj_seq(&self.to_q, &self.to_q_b, &x1n, lq);
let k = proj_seq(&self.to_k, &self.to_k_b, &x2n, lk);
let v = proj_seq(&self.to_v, &self.to_v_b, &x2n, lk);
let mut ctx = vec![0f32; lq * h];
for head in 0..nh {
let off = head * hd;
for i in 0..lq {
let qh = &q[i * h + off..i * h + off + hd];
let mut scores = vec![0f32; lk];
for (j, sc) in scores.iter_mut().enumerate() {
let kh = &k[j * h + off..j * h + off + hd];
*sc = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_inplace(&mut scores);
let out = &mut ctx[i * h + off..i * h + off + hd];
for (j, &w) in scores.iter().enumerate() {
let vh = &v[j * h + off..j * h + off + hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
}
let mut out = vec![0f32; lq * h];
for i in 0..lq {
let p = add_bias(
matvec(&self.proj, &ctx[i * h..i * h + h], h, h),
&self.proj_b,
);
for d in 0..h {
out[i * h + d] = x1[i * h + d] + p[d];
}
}
out
}
pub fn forward(&self, h: &[f32], n: usize) -> Vec<f32> {
let nq = cfg::PERCEIVER_QUERIES;
let pre = self.attn_block(&self.query, nq, h, n);
self.attn_block(&pre, nq, &pre, nq)
}
}
struct SplitMix64(u64);
impl SplitMix64 {
fn new(seed: u64) -> Self {
Self(seed)
}
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
}
}
fn apply_repetition_penalty(logits: &mut [f32], ids: &[u32], penalty: f32) {
let mut seen = std::collections::HashSet::new();
for &id in ids {
let i = id as usize;
if i < logits.len() && seen.insert(id) {
logits[i] = if logits[i] < 0.0 {
logits[i] * penalty
} else {
logits[i] / penalty
};
}
}
}
fn apply_min_p(logits: &mut [f32], min_p: f32) {
let mut probs = logits.to_vec();
softmax_inplace(&mut probs);
let pmax = probs.iter().cloned().fold(0.0f32, f32::max);
let thresh = min_p * pmax;
for (i, &p) in probs.iter().enumerate() {
if p < thresh {
logits[i] = f32::NEG_INFINITY;
}
}
}
fn apply_top_p(logits: &mut [f32], top_p: f32) {
if top_p >= 1.0 {
return;
}
let mut probs = logits.to_vec();
softmax_inplace(&mut probs);
let mut order: Vec<usize> = (0..probs.len()).collect();
order.sort_by(|&a, &b| {
probs[b]
.partial_cmp(&probs[a])
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut cum = 0.0f32;
let mut keep = vec![false; probs.len()];
for (rank, &idx) in order.iter().enumerate() {
cum += probs[idx];
keep[idx] = true;
if cum >= top_p && rank >= 1 {
break;
}
}
for (i, k) in keep.iter().enumerate() {
if !k {
logits[i] = f32::NEG_INFINITY;
}
}
}
fn sample_multinomial(logits: &[f32], rng: &mut SplitMix64) -> u32 {
let mut probs = logits.to_vec();
softmax_inplace(&mut probs);
let u = rng.next_f32();
let mut acc = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
acc += p;
if u < acc {
return i as u32;
}
}
(probs.len() - 1) as u32
}
pub struct CfmSolver {
pub n_timesteps: usize,
pub cfg_rate: f32,
}
impl CfmSolver {
pub fn s3gen() -> Self {
Self {
n_timesteps: cfg::CFM_TIMESTEPS,
cfg_rate: cfg::CFM_INFERENCE_CFG_RATE,
}
}
fn t_span(&self) -> Vec<f32> {
(0..=self.n_timesteps)
.map(|i| {
let ti = i as f32 / self.n_timesteps as f32;
1.0 - (ti * 0.5 * std::f32::consts::PI).cos()
})
.collect()
}
pub fn solve<F>(&self, z: &[f32], mu: &[f32], spks: &[f32], cond: &[f32], est: F) -> Vec<f32>
where
F: Fn(&[f32], &[f32], &[f32], &[f32], f32) -> Vec<f32>,
{
let zero_mu = vec![0f32; mu.len()];
let zero_spks = vec![0f32; spks.len()];
let zero_cond = vec![0f32; cond.len()];
let span = self.t_span();
let mut x = z.to_vec();
for w in span.windows(2) {
let (t, r) = (w[0], w[1]);
let dt = r - t;
let d_cond = est(&x, mu, spks, cond, t);
let d_unc = est(&x, &zero_mu, &zero_spks, &zero_cond, t);
for (i, xi) in x.iter_mut().enumerate() {
*xi += dt * ((1.0 + self.cfg_rate) * d_cond[i] - self.cfg_rate * d_unc[i]);
}
}
x
}
}
pub fn cfm_seed_noise(n: usize, seed: u64) -> Vec<f32> {
let mut rng = SplitMix64::new(seed);
let mut out = Vec::with_capacity(n);
while out.len() < n {
let u1 = rng.next_f32().max(1e-7);
let u2 = rng.next_f32();
let radius = (-2.0 * u1.ln()).sqrt();
out.push(radius * (2.0 * std::f32::consts::PI * u2).cos());
if out.len() < n {
out.push(radius * (2.0 * std::f32::consts::PI * u2).sin());
}
}
out
}
pub fn mish(x: f32) -> f32 {
let softplus = x.max(0.0) + (-(x.abs())).exp().ln_1p();
x * softplus.tanh()
}
pub fn sinusoidal_pos_emb(t: f32, dim: usize, scale: f32) -> Vec<f32> {
let half = dim / 2;
let step = (10000f32).ln() / (half as f32 - 1.0);
let mut out = vec![0f32; dim];
for i in 0..half {
let e = scale * t * (-(i as f32) * step).exp();
out[i] = e.sin();
out[half + i] = e.cos();
}
out
}
pub fn layer_norm_channels(
x: &[f32],
w: &[f32],
b: &[f32],
c: usize,
t: usize,
eps: f32,
) -> Vec<f32> {
let mut out = vec![0f32; c * t];
for ti in 0..t {
let mean = (0..c).map(|ci| x[ci * t + ti]).sum::<f32>() / c as f32;
let var = (0..c)
.map(|ci| (x[ci * t + ti] - mean).powi(2))
.sum::<f32>()
/ c as f32;
let inv = 1.0 / (var + eps).sqrt();
for ci in 0..c {
out[ci * t + ti] = (x[ci * t + ti] - mean) * inv * w[ci] + b[ci];
}
}
out
}
pub fn causal_conv1d(
x: &[f32],
in_ch: usize,
t: usize,
w: &[f32],
bias: Option<&[f32]>,
out_ch: usize,
k: usize,
) -> Vec<f32> {
let pad = k - 1;
let tp = t + pad;
let mut xp = vec![0f32; in_ch * tp];
for ic in 0..in_ch {
xp[ic * tp + pad..ic * tp + tp].copy_from_slice(&x[ic * t..ic * t + t]);
}
let (out, t_out) = conv1d(&xp, in_ch, tp, w, bias, out_ch, k, 1, 0, 1);
debug_assert_eq!(t_out, t);
out
}
pub fn conv_transpose1d(
x: &[f32],
in_ch: usize,
t: usize,
w: &[f32],
bias: Option<&[f32]>,
out_ch: usize,
k: usize,
stride: usize,
pad: usize,
) -> (Vec<f32>, usize) {
let full = (t - 1) * stride + k;
let mut acc = vec![0f32; out_ch * full];
acc.par_chunks_mut(full).enumerate().for_each(|(oc, row)| {
for ic in 0..in_ch {
for ti in 0..t {
let xv = x[ic * t + ti];
for kk in 0..k {
row[ti * stride + kk] += w[(ic * out_ch + oc) * k + kk] * xv;
}
}
}
});
let t_out = full - 2 * pad;
let mut out = vec![0f32; out_ch * t_out];
for oc in 0..out_ch {
let b0 = bias.map(|b| b[oc]).unwrap_or(0.0);
for ot in 0..t_out {
out[oc * t_out + ot] = acc[oc * full + ot + pad] + b0;
}
}
(out, t_out)
}
struct DecResBlock {
b1_cw: Vec<f32>,
b1_cb: Vec<f32>,
b1_lw: Vec<f32>,
b1_lb: Vec<f32>,
b2_cw: Vec<f32>,
b2_cb: Vec<f32>,
b2_lw: Vec<f32>,
b2_lb: Vec<f32>,
mlp_w: Vec<f32>,
mlp_b: Vec<f32>,
res_w: Vec<f32>,
res_b: Vec<f32>,
din: usize,
dout: usize,
}
impl DecResBlock {
fn load(st: &StReader, p: &str, din: usize, dout: usize) -> Result<Self> {
Ok(Self {
b1_cw: st.f32(&format!("{p}.block1.block.0.weight"))?,
b1_cb: st.f32(&format!("{p}.block1.block.0.bias"))?,
b1_lw: st.f32(&format!("{p}.block1.block.2.weight"))?,
b1_lb: st.f32(&format!("{p}.block1.block.2.bias"))?,
b2_cw: st.f32(&format!("{p}.block2.block.0.weight"))?,
b2_cb: st.f32(&format!("{p}.block2.block.0.bias"))?,
b2_lw: st.f32(&format!("{p}.block2.block.2.weight"))?,
b2_lb: st.f32(&format!("{p}.block2.block.2.bias"))?,
mlp_w: st.mat(&format!("{p}.mlp.1.weight"), dout, cfg::DEC_TIME_DIM)?,
mlp_b: st.f32(&format!("{p}.mlp.1.bias"))?,
res_w: st.f32(&format!("{p}.res_conv.weight"))?, res_b: st.f32(&format!("{p}.res_conv.bias"))?,
din,
dout,
})
}
}
struct DecTransformer {
n1_w: Vec<f32>,
n1_b: Vec<f32>,
n3_w: Vec<f32>,
n3_b: Vec<f32>,
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
out_w: Vec<f32>,
out_b: Vec<f32>,
ff1_w: Vec<f32>,
ff1_b: Vec<f32>,
ff2_w: Vec<f32>,
ff2_b: Vec<f32>,
}
impl DecTransformer {
fn load(st: &StReader, p: &str) -> Result<Self> {
let (c, inner) = (cfg::DEC_CHANNELS, cfg::DEC_HEADS * cfg::DEC_HEAD_DIM);
Ok(Self {
n1_w: st.f32(&format!("{p}.norm1.weight"))?,
n1_b: st.f32(&format!("{p}.norm1.bias"))?,
n3_w: st.f32(&format!("{p}.norm3.weight"))?,
n3_b: st.f32(&format!("{p}.norm3.bias"))?,
q: st.mat(&format!("{p}.attn1.to_q.weight"), inner, c)?,
k: st.mat(&format!("{p}.attn1.to_k.weight"), inner, c)?,
v: st.mat(&format!("{p}.attn1.to_v.weight"), inner, c)?,
out_w: st.mat(&format!("{p}.attn1.to_out.0.weight"), c, inner)?,
out_b: st.f32(&format!("{p}.attn1.to_out.0.bias"))?,
ff1_w: st.mat(&format!("{p}.ff.net.0.proj.weight"), c * 4, c)?,
ff1_b: st.f32(&format!("{p}.ff.net.0.proj.bias"))?,
ff2_w: st.mat(&format!("{p}.ff.net.2.weight"), c, c * 4)?,
ff2_b: st.f32(&format!("{p}.ff.net.2.bias"))?,
})
}
}
pub struct Decoder {
t1_w: Vec<f32>,
t1_b: Vec<f32>,
t2_w: Vec<f32>,
t2_b: Vec<f32>,
down_res: DecResBlock,
down_tf: Vec<DecTransformer>,
down_cw: Vec<f32>,
down_cb: Vec<f32>,
mid: Vec<(DecResBlock, Vec<DecTransformer>)>,
up_res: DecResBlock,
up_tf: Vec<DecTransformer>,
up_cw: Vec<f32>,
up_cb: Vec<f32>,
fb_cw: Vec<f32>,
fb_cb: Vec<f32>,
fb_lw: Vec<f32>,
fb_lb: Vec<f32>,
fp_w: Vec<f32>,
fp_b: Vec<f32>,
}
impl Decoder {
pub fn load(st: &StReader, prefix: &str) -> Result<Self> {
let e = |s: &str| format!("{prefix}.{s}");
let (ch, tdim) = (cfg::DEC_CHANNELS, cfg::DEC_TIME_DIM);
let tf_vec = |st: &StReader, base: &str| -> Result<Vec<DecTransformer>> {
(0..cfg::DEC_N_BLOCKS)
.map(|i| DecTransformer::load(st, &format!("{base}.1.{i}")))
.collect()
};
let mut mid = Vec::with_capacity(cfg::DEC_MID_BLOCKS);
for i in 0..cfg::DEC_MID_BLOCKS {
let base = e(&format!("mid_blocks.{i}"));
mid.push((
DecResBlock::load(st, &format!("{base}.0"), ch, ch)?,
tf_vec(st, &base)?,
));
}
Ok(Self {
t1_w: st.mat(&e("time_mlp.linear_1.weight"), tdim, cfg::DEC_IN_CHANNELS)?,
t1_b: st.f32(&e("time_mlp.linear_1.bias"))?,
t2_w: st.mat(&e("time_mlp.linear_2.weight"), tdim, tdim)?,
t2_b: st.f32(&e("time_mlp.linear_2.bias"))?,
down_res: DecResBlock::load(st, &e("down_blocks.0.0"), cfg::DEC_IN_CHANNELS, ch)?,
down_tf: tf_vec(st, &e("down_blocks.0"))?,
down_cw: st.f32(&e("down_blocks.0.2.weight"))?,
down_cb: st.f32(&e("down_blocks.0.2.bias"))?,
mid,
up_res: DecResBlock::load(st, &e("up_blocks.0.0"), 2 * ch, ch)?,
up_tf: tf_vec(st, &e("up_blocks.0"))?,
up_cw: st.f32(&e("up_blocks.0.2.weight"))?,
up_cb: st.f32(&e("up_blocks.0.2.bias"))?,
fb_cw: st.f32(&e("final_block.block.0.weight"))?,
fb_cb: st.f32(&e("final_block.block.0.bias"))?,
fb_lw: st.f32(&e("final_block.block.2.weight"))?,
fb_lb: st.f32(&e("final_block.block.2.bias"))?,
fp_w: st.f32(&e("final_proj.weight"))?, fp_b: st.f32(&e("final_proj.bias"))?,
})
}
fn time_embed(&self, t: f32) -> Vec<f32> {
let se = sinusoidal_pos_emb(t, cfg::DEC_IN_CHANNELS, 1000.0);
let h1 = add_bias(
matvec(&self.t1_w, &se, cfg::DEC_TIME_DIM, cfg::DEC_IN_CHANNELS),
&self.t1_b,
);
let a: Vec<f32> = h1.iter().map(|&v| silu(v)).collect();
add_bias(
matvec(&self.t2_w, &a, cfg::DEC_TIME_DIM, cfg::DEC_TIME_DIM),
&self.t2_b,
)
}
fn causal_block(
cw: &[f32],
cb: &[f32],
lw: &[f32],
lb: &[f32],
x: &[f32],
cin: usize,
cout: usize,
t: usize,
) -> Vec<f32> {
let c = causal_conv1d(x, cin, t, cw, Some(cb), cout, 3);
let n = layer_norm_channels(&c, lw, lb, cout, t, 1e-5);
n.iter().map(|&v| mish(v)).collect()
}
fn res_forward(&self, rb: &DecResBlock, x: &[f32], time_emb: &[f32], t: usize) -> Vec<f32> {
let mut h = Self::causal_block(
&rb.b1_cw, &rb.b1_cb, &rb.b1_lw, &rb.b1_lb, x, rb.din, rb.dout, t,
);
let tm: Vec<f32> = time_emb.iter().map(|&v| mish(v)).collect();
let mlp = add_bias(
matvec(&rb.mlp_w, &tm, rb.dout, cfg::DEC_TIME_DIM),
&rb.mlp_b,
);
for c in 0..rb.dout {
for ti in 0..t {
h[c * t + ti] += mlp[c];
}
}
let h2 = Self::causal_block(
&rb.b2_cw, &rb.b2_cb, &rb.b2_lw, &rb.b2_lb, &h, rb.dout, rb.dout, t,
);
let (res, _) = conv1d(
x,
rb.din,
t,
&rb.res_w,
Some(&rb.res_b),
rb.dout,
1,
1,
0,
1,
);
h2.iter().zip(&res).map(|(a, b)| a + b).collect()
}
fn tf_forward(&self, tf: &DecTransformer, x: &[f32], t: usize) -> Vec<f32> {
let c = cfg::DEC_CHANNELS;
let (nh, hd) = (cfg::DEC_HEADS, cfg::DEC_HEAD_DIM);
let inner = nh * hd;
let scale = 1.0 / (hd as f32).sqrt();
let mut q = vec![0f32; t * inner];
let mut k = vec![0f32; t * inner];
let mut v = vec![0f32; t * inner];
q.par_chunks_mut(inner)
.zip(k.par_chunks_mut(inner))
.zip(v.par_chunks_mut(inner))
.enumerate()
.for_each(|(ti, ((qr, kr), vr))| {
let col: Vec<f32> = (0..c).map(|ci| x[ci * t + ti]).collect();
let n = layer_norm(&col, &tf.n1_w, &tf.n1_b, 1e-5);
qr.copy_from_slice(&matvec(&tf.q, &n, inner, c));
kr.copy_from_slice(&matvec(&tf.k, &n, inner, c));
vr.copy_from_slice(&matvec(&tf.v, &n, inner, c));
});
let mut ctx = vec![0f32; t * inner];
ctx.par_chunks_mut(inner)
.enumerate()
.for_each(|(i, out_row)| {
let mut sc = vec![0f32; t];
for head in 0..nh {
let off = head * hd;
let qh = &q[i * inner + off..i * inner + off + hd];
for (j, s) in sc.iter_mut().enumerate() {
let kh = &k[j * inner + off..j * inner + off + hd];
*s = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_inplace(&mut sc);
let out = &mut out_row[off..off + hd];
for (j, &w) in sc.iter().enumerate() {
let vh = &v[j * inner + off..j * inner + off + hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
});
let mut xr = x.to_vec();
let attn_out: Vec<Vec<f32>> = (0..t)
.into_par_iter()
.map(|ti| {
add_bias(
matvec(&tf.out_w, &ctx[ti * inner..ti * inner + inner], c, inner),
&tf.out_b,
)
})
.collect();
for (ti, ao) in attn_out.iter().enumerate() {
for ci in 0..c {
xr[ci * t + ti] += ao[ci];
}
}
let ff_out: Vec<Vec<f32>> = (0..t)
.into_par_iter()
.map(|ti| {
let col: Vec<f32> = (0..c).map(|ci| xr[ci * t + ti]).collect();
let n = layer_norm(&col, &tf.n3_w, &tf.n3_b, 1e-5);
let f1 = add_bias(matvec(&tf.ff1_w, &n, c * 4, c), &tf.ff1_b);
let g: Vec<f32> = f1.iter().map(|&z| gelu(z)).collect();
add_bias(matvec(&tf.ff2_w, &g, c, c * 4), &tf.ff2_b)
})
.collect();
for (ti, f2) in ff_out.iter().enumerate() {
for ci in 0..c {
xr[ci * t + ti] += f2[ci];
}
}
xr
}
pub fn forward(
&self,
x_mel: &[f32],
mu: &[f32],
spks: &[f32],
cond: &[f32],
t_scalar: f32,
t: usize,
) -> Vec<f32> {
let m = cfg::S3GEN_MELS;
let ch = cfg::DEC_CHANNELS;
let time_emb = self.time_embed(t_scalar);
let mut x = vec![0f32; cfg::DEC_IN_CHANNELS * t];
for c in 0..m {
for ti in 0..t {
x[c * t + ti] = x_mel[c * t + ti];
x[(m + c) * t + ti] = mu[c * t + ti];
x[(2 * m + c) * t + ti] = spks[c];
x[(3 * m + c) * t + ti] = cond[c * t + ti];
}
}
let mut xd = self.res_forward(&self.down_res, &x, &time_emb, t);
for tf in &self.down_tf {
xd = self.tf_forward(tf, &xd, t);
}
let skip = xd.clone();
xd = causal_conv1d(&xd, ch, t, &self.down_cw, Some(&self.down_cb), ch, 3);
for (rb, tfs) in &self.mid {
xd = self.res_forward(rb, &xd, &time_emb, t);
for tf in tfs {
xd = self.tf_forward(tf, &xd, t);
}
}
let mut xu = vec![0f32; 2 * ch * t];
for c in 0..ch {
for ti in 0..t {
xu[c * t + ti] = xd[c * t + ti];
xu[(ch + c) * t + ti] = skip[c * t + ti];
}
}
let mut xr = self.res_forward(&self.up_res, &xu, &time_emb, t);
for tf in &self.up_tf {
xr = self.tf_forward(tf, &xr, t);
}
xr = causal_conv1d(&xr, ch, t, &self.up_cw, Some(&self.up_cb), ch, 3);
let fb = Self::causal_block(
&self.fb_cw,
&self.fb_cb,
&self.fb_lw,
&self.fb_lb,
&xr,
ch,
ch,
t,
);
let (out, _) = conv1d(&fb, ch, t, &self.fp_w, Some(&self.fp_b), m, 1, 1, 0, 1);
out
}
}
fn rel_pos_encoding(t: usize, d: usize) -> Vec<f32> {
let cols = 2 * t - 1;
let mut pe = vec![0f32; cols * d];
let ln10k = (10000f32).ln();
for k in 0..cols {
let r = (t as isize - 1 - k as isize) as f32;
for i in 0..d / 2 {
let f = (-((2 * i) as f32) * ln10k / d as f32).exp();
pe[k * d + 2 * i] = (r * f).sin();
pe[k * d + 2 * i + 1] = (r * f).cos();
}
}
pe
}
fn rel_shift(x: &[f32], t: usize) -> Vec<f32> {
let cols = 2 * t - 1;
let padded_cols = cols + 1; let mut flat = vec![0f32; t * padded_cols];
for i in 0..t {
for c in 0..cols {
flat[i * padded_cols + 1 + c] = x[i * cols + c];
}
}
let mut out = vec![0f32; t * t];
for i in 0..t {
for j in 0..t {
out[i * t + j] = flat[t + i * cols + j];
}
}
out
}
struct ConformerLayer {
nm_w: Vec<f32>,
nm_b: Vec<f32>,
nf_w: Vec<f32>,
nf_b: Vec<f32>,
q_w: Vec<f32>,
q_b: Vec<f32>,
k_w: Vec<f32>,
k_b: Vec<f32>,
v_w: Vec<f32>,
v_b: Vec<f32>,
out_w: Vec<f32>,
out_b: Vec<f32>,
pos_w: Vec<f32>, bias_u: Vec<f32>,
bias_v: Vec<f32>, ff1_w: Vec<f32>,
ff1_b: Vec<f32>,
ff2_w: Vec<f32>,
ff2_b: Vec<f32>,
}
impl ConformerLayer {
fn load(st: &StReader, p: &str) -> Result<Self> {
let (d, u) = (cfg::CONFORMER_DIM, cfg::CONFORMER_UNITS);
Ok(Self {
nm_w: st.f32(&format!("{p}.norm_mha.weight"))?,
nm_b: st.f32(&format!("{p}.norm_mha.bias"))?,
nf_w: st.f32(&format!("{p}.norm_ff.weight"))?,
nf_b: st.f32(&format!("{p}.norm_ff.bias"))?,
q_w: st.mat(&format!("{p}.self_attn.linear_q.weight"), d, d)?,
q_b: st.f32(&format!("{p}.self_attn.linear_q.bias"))?,
k_w: st.mat(&format!("{p}.self_attn.linear_k.weight"), d, d)?,
k_b: st.f32(&format!("{p}.self_attn.linear_k.bias"))?,
v_w: st.mat(&format!("{p}.self_attn.linear_v.weight"), d, d)?,
v_b: st.f32(&format!("{p}.self_attn.linear_v.bias"))?,
out_w: st.mat(&format!("{p}.self_attn.linear_out.weight"), d, d)?,
out_b: st.f32(&format!("{p}.self_attn.linear_out.bias"))?,
pos_w: st.mat(&format!("{p}.self_attn.linear_pos.weight"), d, d)?,
bias_u: st.f32(&format!("{p}.self_attn.pos_bias_u"))?,
bias_v: st.f32(&format!("{p}.self_attn.pos_bias_v"))?,
ff1_w: st.mat(&format!("{p}.feed_forward.w_1.weight"), u, d)?,
ff1_b: st.f32(&format!("{p}.feed_forward.w_1.bias"))?,
ff2_w: st.mat(&format!("{p}.feed_forward.w_2.weight"), d, u)?,
ff2_b: st.f32(&format!("{p}.feed_forward.w_2.bias"))?,
})
}
fn forward(&self, x: &[f32], pos_emb: &[f32], t: usize) -> Vec<f32> {
let d = cfg::CONFORMER_DIM;
let (nh, dk) = (
cfg::CONFORMER_HEADS,
cfg::CONFORMER_DIM / cfg::CONFORMER_HEADS,
);
let scale = 1.0 / (dk as f32).sqrt();
let mut q = vec![0f32; t * d];
let mut k = vec![0f32; t * d];
let mut v = vec![0f32; t * d];
for i in 0..t {
let n = layer_norm(&x[i * d..i * d + d], &self.nm_w, &self.nm_b, 1e-12);
q[i * d..i * d + d].copy_from_slice(&add_bias(matvec(&self.q_w, &n, d, d), &self.q_b));
k[i * d..i * d + d].copy_from_slice(&add_bias(matvec(&self.k_w, &n, d, d), &self.k_b));
v[i * d..i * d + d].copy_from_slice(&add_bias(matvec(&self.v_w, &n, d, d), &self.v_b));
}
let np = 2 * t - 1;
let mut p = vec![0f32; np * d];
for m in 0..np {
p[m * d..m * d + d].copy_from_slice(&matvec(
&self.pos_w,
&pos_emb[m * d..m * d + d],
d,
d,
));
}
let mut ctx = vec![0f32; t * d];
for h in 0..nh {
let off = h * dk;
let mut bd = vec![0f32; t * np];
for i in 0..t {
for m in 0..np {
let mut s = 0.0;
for e in 0..dk {
s += (q[i * d + off + e] + self.bias_v[off + e]) * p[m * d + off + e];
}
bd[i * np + m] = s;
}
}
let bd_shifted = rel_shift(&bd, t);
for i in 0..t {
let mut scores = vec![0f32; t];
for j in 0..t {
let mut ac = 0.0;
for e in 0..dk {
ac += (q[i * d + off + e] + self.bias_u[off + e]) * k[j * d + off + e];
}
scores[j] = (ac + bd_shifted[i * t + j]) * scale;
}
softmax_inplace(&mut scores);
for (j, &w) in scores.iter().enumerate() {
for e in 0..dk {
ctx[i * d + off + e] += w * v[j * d + off + e];
}
}
}
}
let mut x1 = x.to_vec();
for i in 0..t {
let o = add_bias(
matvec(&self.out_w, &ctx[i * d..i * d + d], d, d),
&self.out_b,
);
for e in 0..d {
x1[i * d + e] += o[e];
}
}
let u = cfg::CONFORMER_UNITS;
for i in 0..t {
let n = layer_norm(&x1[i * d..i * d + d], &self.nf_w, &self.nf_b, 1e-12);
let h1 = add_bias(matvec(&self.ff1_w, &n, u, d), &self.ff1_b);
let a: Vec<f32> = h1.iter().map(|&z| silu(z)).collect();
let h2 = add_bias(matvec(&self.ff2_w, &a, d, u), &self.ff2_b);
for e in 0..d {
x1[i * d + e] += h2[e];
}
}
x1
}
}
pub struct UpConformer {
input_emb: Vec<f32>, emb_lin_w: Vec<f32>,
emb_lin_b: Vec<f32>,
emb_ln_w: Vec<f32>,
emb_ln_b: Vec<f32>,
pre_c1_w: Vec<f32>,
pre_c1_b: Vec<f32>,
pre_c2_w: Vec<f32>,
pre_c2_b: Vec<f32>,
encoders: Vec<ConformerLayer>,
up_conv_w: Vec<f32>,
up_conv_b: Vec<f32>,
uemb_lin_w: Vec<f32>,
uemb_lin_b: Vec<f32>,
uemb_ln_w: Vec<f32>,
uemb_ln_b: Vec<f32>,
up_encoders: Vec<ConformerLayer>,
after_w: Vec<f32>,
after_b: Vec<f32>,
proj_w: Vec<f32>,
proj_b: Vec<f32>,
}
impl UpConformer {
pub fn load(st: &StReader, flow: &str) -> Result<Self> {
let d = cfg::CONFORMER_DIM;
let e = |s: &str| format!("{flow}.{s}");
Ok(Self {
input_emb: st.mat(&e("input_embedding.weight"), cfg::S3_SPEECH_VOCAB, d)?,
emb_lin_w: st.mat(&e("encoder.embed.out.0.weight"), d, d)?,
emb_lin_b: st.f32(&e("encoder.embed.out.0.bias"))?,
emb_ln_w: st.f32(&e("encoder.embed.out.1.weight"))?,
emb_ln_b: st.f32(&e("encoder.embed.out.1.bias"))?,
pre_c1_w: st.f32(&e("encoder.pre_lookahead_layer.conv1.weight"))?, pre_c1_b: st.f32(&e("encoder.pre_lookahead_layer.conv1.bias"))?,
pre_c2_w: st.f32(&e("encoder.pre_lookahead_layer.conv2.weight"))?, pre_c2_b: st.f32(&e("encoder.pre_lookahead_layer.conv2.bias"))?,
encoders: (0..cfg::CONFORMER_BLOCKS)
.map(|i| ConformerLayer::load(st, &e(&format!("encoder.encoders.{i}"))))
.collect::<Result<_>>()?,
up_conv_w: st.f32(&e("encoder.up_layer.conv.weight"))?, up_conv_b: st.f32(&e("encoder.up_layer.conv.bias"))?,
uemb_lin_w: st.mat(&e("encoder.up_embed.out.0.weight"), d, d)?,
uemb_lin_b: st.f32(&e("encoder.up_embed.out.0.bias"))?,
uemb_ln_w: st.f32(&e("encoder.up_embed.out.1.weight"))?,
uemb_ln_b: st.f32(&e("encoder.up_embed.out.1.bias"))?,
up_encoders: (0..cfg::UP_CONFORMER_BLOCKS)
.map(|i| ConformerLayer::load(st, &e(&format!("encoder.up_encoders.{i}"))))
.collect::<Result<_>>()?,
after_w: st.f32(&e("encoder.after_norm.weight"))?,
after_b: st.f32(&e("encoder.after_norm.bias"))?,
proj_w: st.mat(&e("encoder_proj.weight"), cfg::S3GEN_MELS, d)?,
proj_b: st.f32(&e("encoder_proj.bias"))?,
})
}
fn embed(
&self,
x: &[f32],
lw: &[f32],
lb: &[f32],
nw: &[f32],
nb: &[f32],
t: usize,
) -> Vec<f32> {
let d = cfg::CONFORMER_DIM;
let xscale = (d as f32).sqrt();
let mut out = vec![0f32; t * d];
for i in 0..t {
let l = add_bias(matvec(lw, &x[i * d..i * d + d], d, d), lb);
let n = layer_norm(&l, nw, nb, 1e-5);
for e in 0..d {
out[i * d + e] = n[e] * xscale;
}
}
out
}
fn pre_lookahead(&self, x: &[f32], t: usize) -> Vec<f32> {
let d = cfg::CONFORMER_DIM;
let mut xc = vec![0f32; d * t];
for i in 0..t {
for e in 0..d {
xc[e * t + i] = x[i * d + e];
}
}
let mut p1 = vec![0f32; d * (t + 3)];
for e in 0..d {
p1[e * (t + 3)..e * (t + 3) + t].copy_from_slice(&xc[e * t..e * t + t]);
}
let (c1, _) = conv1d(
&p1,
d,
t + 3,
&self.pre_c1_w,
Some(&self.pre_c1_b),
d,
4,
1,
0,
1,
);
let c1: Vec<f32> = c1
.iter()
.map(|&z| if z >= 0.0 { z } else { 0.1 * z })
.collect();
let mut p2 = vec![0f32; d * (t + 2)];
for e in 0..d {
p2[e * (t + 2) + 2..e * (t + 2) + t + 2].copy_from_slice(&c1[e * t..e * t + t]);
}
let (c2, _) = conv1d(
&p2,
d,
t + 2,
&self.pre_c2_w,
Some(&self.pre_c2_b),
d,
3,
1,
0,
1,
);
let mut out = x.to_vec();
for i in 0..t {
for e in 0..d {
out[i * d + e] += c2[e * t + i];
}
}
out
}
fn up_layer(&self, x: &[f32], t: usize) -> Vec<f32> {
let d = cfg::CONFORMER_DIM;
let t2 = 2 * t;
let mut interp = vec![0f32; d * t2];
for i in 0..t {
for e in 0..d {
interp[e * t2 + 2 * i] = x[i * d + e];
interp[e * t2 + 2 * i + 1] = x[i * d + e];
}
}
let mut padded = vec![0f32; d * (t2 + 4)];
for e in 0..d {
padded[e * (t2 + 4) + 4..e * (t2 + 4) + t2 + 4]
.copy_from_slice(&interp[e * t2..e * t2 + t2]);
}
let (c, _) = conv1d(
&padded,
d,
t2 + 4,
&self.up_conv_w,
Some(&self.up_conv_b),
d,
5,
1,
0,
1,
);
let mut out = vec![0f32; t2 * d];
for i in 0..t2 {
for e in 0..d {
out[i * d + e] = c[e * t2 + i];
}
}
out
}
pub fn forward(&self, tokens: &[u32]) -> Vec<f32> {
let d = cfg::CONFORMER_DIM;
let t = tokens.len();
let mut x = vec![0f32; t * d];
for (i, &tok) in tokens.iter().enumerate() {
x[i * d..i * d + d]
.copy_from_slice(&self.input_emb[tok as usize * d..tok as usize * d + d]);
}
let mut x = self.embed(
&x,
&self.emb_lin_w,
&self.emb_lin_b,
&self.emb_ln_w,
&self.emb_ln_b,
t,
);
x = self.pre_lookahead(&x, t);
let pe = rel_pos_encoding(t, d);
for enc in &self.encoders {
x = enc.forward(&x, &pe, t);
}
let mut x = self.up_layer(&x, t);
let t2 = 2 * t;
x = self.embed(
&x,
&self.uemb_lin_w,
&self.uemb_lin_b,
&self.uemb_ln_w,
&self.uemb_ln_b,
t2,
);
let pe2 = rel_pos_encoding(t2, d);
for enc in &self.up_encoders {
x = enc.forward(&x, &pe2, t2);
}
let mut mu = vec![0f32; t2 * cfg::S3GEN_MELS];
for i in 0..t2 {
let n = layer_norm(&x[i * d..i * d + d], &self.after_w, &self.after_b, 1e-5);
let m = add_bias(matvec(&self.proj_w, &n, cfg::S3GEN_MELS, d), &self.proj_b);
mu[i * cfg::S3GEN_MELS..i * cfg::S3GEN_MELS + cfg::S3GEN_MELS].copy_from_slice(&m);
}
mu
}
}
fn batchnorm(
x: &[f32],
m: &[f32],
v: &[f32],
w: &[f32],
b: &[f32],
c: usize,
l: usize,
eps: f32,
) -> Vec<f32> {
let mut out = vec![0f32; c * l];
for ci in 0..c {
let inv = 1.0 / (v[ci] + eps).sqrt();
for li in 0..l {
out[ci * l + li] = (x[ci * l + li] - m[ci]) * inv * w[ci] + b[ci];
}
}
out
}
#[allow(clippy::too_many_arguments)]
fn conv2d(
inp: &[f32],
cin: usize,
h: usize,
w: usize,
weight: &[f32],
bias: Option<&[f32]>,
cout: usize,
kh: usize,
kw: usize,
sh: usize,
sw: usize,
ph: usize,
pw: usize,
) -> (Vec<f32>, usize, usize) {
let ho = (h + 2 * ph - kh) / sh + 1;
let wo = (w + 2 * pw - kw) / sw + 1;
let mut out = vec![0f32; cout * ho * wo];
for oc in 0..cout {
let b0 = bias.map(|b| b[oc]).unwrap_or(0.0);
for oy in 0..ho {
for ox in 0..wo {
let mut acc = b0;
for ic in 0..cin {
for ky in 0..kh {
let iy = (oy * sh + ky) as isize - ph as isize;
if iy < 0 || iy as usize >= h {
continue;
}
for kx in 0..kw {
let ix = (ox * sw + kx) as isize - pw as isize;
if ix >= 0 && (ix as usize) < w {
acc += weight[((oc * cin + ic) * kh + ky) * kw + kx]
* inp[(ic * h + iy as usize) * w + ix as usize];
}
}
}
}
out[(oc * ho + oy) * wo + ox] = acc;
}
}
}
(out, ho, wo)
}
fn conv1d_dilated(
x: &[f32],
cin: usize,
t: usize,
w: &[f32],
cout: usize,
k: usize,
pad: usize,
dil: usize,
) -> Vec<f32> {
let mut out = vec![0f32; cout * t];
out.par_chunks_mut(t).enumerate().for_each(|(oc, row)| {
for (ot, o) in row.iter_mut().enumerate() {
let mut acc = 0.0;
for ic in 0..cin {
for kk in 0..k {
let it = (ot + kk * dil) as isize - pad as isize;
if it >= 0 && (it as usize) < t {
acc += w[(oc * cin + ic) * k + kk] * x[ic * t + it as usize];
}
}
}
*o = acc;
}
});
out
}
struct Bn {
m: Vec<f32>,
v: Vec<f32>,
w: Vec<f32>,
b: Vec<f32>,
}
impl Bn {
fn load(st: &StReader, p: &str, ch: usize) -> Result<Self> {
Ok(Self {
m: st.f32(&format!("{p}.running_mean"))?,
v: st.f32(&format!("{p}.running_var"))?,
w: st
.f32(&format!("{p}.weight"))
.unwrap_or_else(|_| vec![1.0; ch]), b: st
.f32(&format!("{p}.bias"))
.unwrap_or_else(|_| vec![0.0; ch]),
})
}
fn apply(&self, x: &[f32], c: usize, l: usize) -> Vec<f32> {
batchnorm(x, &self.m, &self.v, &self.w, &self.b, c, l, 1e-5)
}
}
struct DenseTdnnLayer {
nl1: Bn,
lin1: Vec<f32>, in_ch: usize,
nl2: Bn,
local: Vec<f32>, c1w: Vec<f32>,
c1b: Vec<f32>, c2w: Vec<f32>,
c2b: Vec<f32>, dil: usize,
}
impl DenseTdnnLayer {
fn load(st: &StReader, p: &str, in_ch: usize, dil: usize) -> Result<Self> {
Ok(Self {
nl1: Bn::load(st, &format!("{p}.nonlinear1.batchnorm"), in_ch)?,
lin1: st.f32(&format!("{p}.linear1.weight"))?,
in_ch,
nl2: Bn::load(st, &format!("{p}.nonlinear2.batchnorm"), cfg::CAMP_BN)?,
local: st.f32(&format!("{p}.cam_layer.linear_local.weight"))?,
c1w: st.f32(&format!("{p}.cam_layer.linear1.weight"))?,
c1b: st.f32(&format!("{p}.cam_layer.linear1.bias"))?,
c2w: st.f32(&format!("{p}.cam_layer.linear2.weight"))?,
c2b: st.f32(&format!("{p}.cam_layer.linear2.bias"))?,
dil,
})
}
fn forward(&self, x: &[f32], t: usize) -> Vec<f32> {
let bn = cfg::CAMP_BN; let g = cfg::CAMP_GROWTH; let red = bn / 2; let h = relu_inplace(self.nl1.apply(x, self.in_ch, t));
let (h, _) = conv1d(&h, self.in_ch, t, &self.lin1, None, bn, 1, 1, 0, 1);
let h = relu_inplace(self.nl2.apply(&h, bn, t));
let pad = self.dil; let y = conv1d_dilated(&h, bn, t, &self.local, g, 3, pad, self.dil);
let seg = seg_pooling(&h, bn, t, 100);
let mut ctx = vec![0f32; bn * t];
for c in 0..bn {
let mean: f32 = h[c * t..c * t + t].iter().sum::<f32>() / t as f32;
for ti in 0..t {
ctx[c * t + ti] = mean + seg[c * t + ti];
}
}
let (c1, _) = conv1d(&ctx, bn, t, &self.c1w, Some(&self.c1b), red, 1, 1, 0, 1);
let c1 = relu_inplace(c1);
let (c2, _) = conv1d(&c1, red, t, &self.c2w, Some(&self.c2b), g, 1, 1, 0, 1);
let mut out = vec![0f32; g * t];
for i in 0..g * t {
out[i] = y[i] / (1.0 + (-c2[i]).exp()); }
out
}
}
fn relu_inplace(mut v: Vec<f32>) -> Vec<f32> {
for x in v.iter_mut() {
if *x < 0.0 {
*x = 0.0;
}
}
v
}
fn seg_pooling(x: &[f32], c: usize, t: usize, seg_len: usize) -> Vec<f32> {
let nseg = t.div_ceil(seg_len);
let mut out = vec![0f32; c * t];
for ci in 0..c {
for s in 0..nseg {
let a = s * seg_len;
let b = ((s + 1) * seg_len).min(t);
let mean: f32 = x[ci * t + a..ci * t + b].iter().sum::<f32>() / (b - a) as f32;
for ti in a..b {
out[ci * t + ti] = mean;
}
}
}
out
}
struct ResBlock2D {
c1: Vec<f32>,
bn1: Bn,
c2: Vec<f32>,
bn2: Bn,
sc_conv: Option<Vec<f32>>,
sc_bn: Option<Bn>,
cin: usize,
cout: usize,
stride: usize,
}
impl ResBlock2D {
fn load(st: &StReader, p: &str, cin: usize, cout: usize, stride: usize) -> Result<Self> {
let sc = st.f32(&format!("{p}.shortcut.0.weight")).ok();
let sc_bn = if sc.is_some() {
Some(Bn::load(st, &format!("{p}.shortcut.1"), cout)?)
} else {
None
};
Ok(Self {
c1: st.f32(&format!("{p}.conv1.weight"))?,
bn1: Bn::load(st, &format!("{p}.bn1"), cout)?,
c2: st.f32(&format!("{p}.conv2.weight"))?,
bn2: Bn::load(st, &format!("{p}.bn2"), cout)?,
sc_conv: sc,
sc_bn,
cin,
cout,
stride,
})
}
fn forward(&self, x: &[f32], h: usize, w: usize) -> (Vec<f32>, usize) {
let (o1, h1, w1) = conv2d(
x,
self.cin,
h,
w,
&self.c1,
None,
self.cout,
3,
3,
self.stride,
1,
1,
1,
);
let o1 = relu_inplace(self.bn1.apply(&o1, self.cout, h1 * w1));
let (o2, h2, w2) = conv2d(
&o1, self.cout, h1, w1, &self.c2, None, self.cout, 3, 3, 1, 1, 1, 1,
);
let o2 = self.bn2.apply(&o2, self.cout, h2 * w2);
let sc = if let (Some(cw), Some(bn)) = (&self.sc_conv, &self.sc_bn) {
let (s, sh, sw) = conv2d(
x,
self.cin,
h,
w,
cw,
None,
self.cout,
1,
1,
self.stride,
1,
0,
0,
);
bn.apply(&s, self.cout, sh * sw)
} else {
x.to_vec()
};
let mut out = vec![0f32; self.cout * h2 * w2];
for i in 0..self.cout * h2 * w2 {
out[i] = (o2[i] + sc[i]).max(0.0);
}
(out, h2)
}
}
pub struct CampPlus {
conv1: Vec<f32>,
bn1: Bn,
layer1: Vec<ResBlock2D>,
layer2: Vec<ResBlock2D>,
conv2: Vec<f32>,
bn2: Bn,
tdnn_w: Vec<f32>,
tdnn_bn: Bn,
blocks: Vec<Vec<DenseTdnnLayer>>,
transits: Vec<(Bn, Vec<f32>, usize, usize)>, out_bn: Bn,
dense_w: Vec<f32>,
dense_bn: Bn,
}
impl CampPlus {
pub fn load(st: &StReader, prefix: &str) -> Result<Self> {
let h = |s: &str| format!("{prefix}.head.{s}");
let xv = |s: &str| format!("{prefix}.xvector.{s}");
let mut blocks = Vec::new();
let mut transits = Vec::new();
let mut channels = cfg::CAMP_INIT; for (bi, &(nl, dil)) in cfg::CAMP_BLOCKS.iter().enumerate() {
let mut layers = Vec::with_capacity(nl);
for li in 0..nl {
let in_ch = channels + li * cfg::CAMP_GROWTH;
layers.push(DenseTdnnLayer::load(
st,
&xv(&format!("block{}.tdnnd{}", bi + 1, li + 1)),
in_ch,
dil,
)?);
}
channels += nl * cfg::CAMP_GROWTH;
let out = channels / 2;
transits.push((
Bn::load(
st,
&xv(&format!("transit{}.nonlinear.batchnorm", bi + 1)),
channels,
)?,
st.f32(&xv(&format!("transit{}.linear.weight", bi + 1)))?,
channels,
out,
));
channels = out;
blocks.push(layers);
}
Ok(Self {
conv1: st.f32(&h("conv1.weight"))?,
bn1: Bn::load(st, &h("bn1"), 32)?,
layer1: vec![
ResBlock2D::load(st, &h("layer1.0"), 32, 32, 2)?,
ResBlock2D::load(st, &h("layer1.1"), 32, 32, 1)?,
],
layer2: vec![
ResBlock2D::load(st, &h("layer2.0"), 32, 32, 2)?,
ResBlock2D::load(st, &h("layer2.1"), 32, 32, 1)?,
],
conv2: st.f32(&h("conv2.weight"))?,
bn2: Bn::load(st, &h("bn2"), 32)?,
tdnn_w: st.f32(&xv("tdnn.linear.weight"))?,
tdnn_bn: Bn::load(st, &xv("tdnn.nonlinear.batchnorm"), cfg::CAMP_INIT)?,
blocks,
transits,
out_bn: Bn::load(st, &xv("out_nonlinear.batchnorm"), channels)?,
dense_w: st.f32(&xv("dense.linear.weight"))?,
dense_bn: Bn::load(st, &xv("dense.nonlinear.batchnorm"), cfg::CAMP_EMBED)?,
})
}
pub fn embed(&self, mel: &[f32], frames: usize) -> Vec<f32> {
let feat = cfg::S3GEN_MELS; let mut x = vec![0f32; feat * frames];
for f in 0..frames {
for m in 0..feat {
x[m * frames + f] = mel[f * feat + m];
}
}
let (o, h1, w1) = conv2d(&x, 1, feat, frames, &self.conv1, None, 32, 3, 3, 1, 1, 1, 1);
let mut cur = relu_inplace(self.bn1.apply(&o, 32, h1 * w1));
let (mut ch, cw) = (h1, w1);
for rb in self.layer1.iter().chain(self.layer2.iter()) {
let (o, nh) = rb.forward(&cur, ch, cw);
cur = o;
ch = nh;
}
let (o, ch2, cw2) = conv2d(&cur, 32, ch, cw, &self.conv2, None, 32, 3, 3, 2, 1, 1, 1);
let cur = relu_inplace(self.bn2.apply(&o, 32, ch2 * cw2));
let feat_ch = 32 * ch2;
let t0 = cw2;
let (o, t1) = conv1d(
&cur,
feat_ch,
t0,
&self.tdnn_w,
None,
cfg::CAMP_INIT,
5,
2,
2,
1,
);
let mut xv = relu_inplace(self.tdnn_bn.apply(&o, cfg::CAMP_INIT, t1));
let mut chn = cfg::CAMP_INIT;
for (bi, layers) in self.blocks.iter().enumerate() {
for layer in layers {
let add = layer.forward(&xv, t1); let mut nx = vec![0f32; (chn + cfg::CAMP_GROWTH) * t1];
nx[..chn * t1].copy_from_slice(&xv);
nx[chn * t1..].copy_from_slice(&add);
xv = nx;
chn += cfg::CAMP_GROWTH;
}
let (bn, lin, cin, cout) = &self.transits[bi];
let t = relu_inplace(bn.apply(&xv, *cin, t1));
let (o, _) = conv1d(&t, *cin, t1, lin, None, *cout, 1, 1, 0, 1);
xv = o;
chn = *cout;
}
xv = relu_inplace(self.out_bn.apply(&xv, chn, t1));
let mut stats = vec![0f32; 2 * chn];
for c in 0..chn {
let mean: f32 = xv[c * t1..c * t1 + t1].iter().sum::<f32>() / t1 as f32;
let var: f32 = xv[c * t1..c * t1 + t1]
.iter()
.map(|v| (v - mean).powi(2))
.sum::<f32>()
/ (t1 as f32 - 1.0).max(1.0); stats[c] = mean;
stats[chn + c] = var.max(0.0).sqrt(); }
let (d, _) = conv1d(
&stats,
2 * chn,
1,
&self.dense_w,
None,
cfg::CAMP_EMBED,
1,
1,
0,
1,
);
self.dense_bn.apply(&d, cfg::CAMP_EMBED, 1)
}
}
fn recon_wn(g: &[f32], v: &[f32], d0: usize, rest: usize) -> Vec<f32> {
let mut w = vec![0f32; d0 * rest];
for i in 0..d0 {
let base = i * rest;
let norm = v[base..base + rest]
.iter()
.map(|x| x * x)
.sum::<f32>()
.sqrt();
let s = g[i] / norm;
for j in 0..rest {
w[base + j] = v[base + j] * s;
}
}
w
}
fn load_wn(st: &StReader, p: &str, d0: usize, rest: usize) -> Result<(Vec<f32>, Vec<f32>)> {
let g = st.f32(&format!("{p}.parametrizations.weight.original0"))?;
let v = st.f32(&format!("{p}.parametrizations.weight.original1"))?;
let b = st.f32(&format!("{p}.bias"))?;
Ok((recon_wn(&g, &v, d0, rest), b))
}
fn hift_snake(x: &mut [f32], alpha: &[f32], t: usize) {
for (ci, a) in alpha.iter().enumerate() {
for v in x[ci * t..ci * t + t].iter_mut() {
let s = (a * *v).sin();
*v += s * s / a;
}
}
}
fn leaky(x: &mut [f32], slope: f32) {
for v in x.iter_mut() {
if *v < 0.0 {
*v *= slope;
}
}
}
fn add_bias_ch(x: &mut [f32], b: &[f32], c: usize, t: usize) {
for ci in 0..c {
for ti in 0..t {
x[ci * t + ti] += b[ci];
}
}
}
struct VocResBlock {
convs1: Vec<(Vec<f32>, Vec<f32>)>,
convs2: Vec<(Vec<f32>, Vec<f32>)>,
a1: Vec<Vec<f32>>,
a2: Vec<Vec<f32>>,
ch: usize,
k: usize,
dils: Vec<usize>,
}
impl VocResBlock {
fn load(st: &StReader, p: &str, ch: usize, k: usize, dils: &[usize]) -> Result<Self> {
let (mut c1, mut c2, mut a1, mut a2) = (vec![], vec![], vec![], vec![]);
for i in 0..dils.len() {
c1.push(load_wn(st, &format!("{p}.convs1.{i}"), ch, ch * k)?);
c2.push(load_wn(st, &format!("{p}.convs2.{i}"), ch, ch * k)?);
a1.push(st.f32(&format!("{p}.activations1.{i}.alpha"))?);
a2.push(st.f32(&format!("{p}.activations2.{i}.alpha"))?);
}
Ok(Self {
convs1: c1,
convs2: c2,
a1,
a2,
ch,
k,
dils: dils.to_vec(),
})
}
fn forward(&self, x: &[f32], t: usize) -> Vec<f32> {
let mut x = x.to_vec();
for i in 0..self.dils.len() {
let mut xt = x.clone();
hift_snake(&mut xt, &self.a1[i], t);
let mut y = conv1d_dilated(
&xt,
self.ch,
t,
&self.convs1[i].0,
self.ch,
self.k,
(self.k - 1) / 2 * self.dils[i],
self.dils[i],
);
add_bias_ch(&mut y, &self.convs1[i].1, self.ch, t);
hift_snake(&mut y, &self.a2[i], t);
let mut z = conv1d_dilated(
&y,
self.ch,
t,
&self.convs2[i].0,
self.ch,
self.k,
(self.k - 1) / 2,
1,
);
add_bias_ch(&mut z, &self.convs2[i].1, self.ch, t);
for (a, b) in x.iter_mut().zip(z) {
*a += b;
}
}
x
}
}
fn hift_sine_source(f0: &[f32], up: usize, sr: f32, mw: &[f32], mb: f32) -> Vec<f32> {
let t_in = f0.len();
let l = t_in * up;
let nh = cfg::HIFT_HARMONICS;
let two_pi = 2.0 * std::f64::consts::PI;
let mut sines = vec![0f32; l * nh];
for hi in 0..nh {
let mut phase_dn = vec![0f64; t_in];
let mut acc = 0f64;
for i in 0..t_in {
let v = f0[i] as f64 * (hi as f64 + 1.0) / sr as f64;
acc += v - v.floor();
phase_dn[i] = acc * two_pi * up as f64;
}
for t in 0..l {
let src = ((t as f64 + 0.5) / up as f64 - 0.5)
.max(0.0)
.min(t_in as f64 - 1.0);
let i0 = (src as usize).min(t_in - 1);
let i1 = (i0 + 1).min(t_in - 1);
let w = src - i0 as f64;
let ph = phase_dn[i0] * (1.0 - w) + phase_dn[i1] * w;
sines[t * nh + hi] = ph.rem_euclid(two_pi).sin() as f32;
}
}
let mut har = vec![0f32; l];
for t in 0..l {
let uv = if f0[(t / up).min(t_in - 1)] > 10.0 {
1.0f32
} else {
0.0
};
let mut acc = mb;
for hi in 0..nh {
acc += sines[t * nh + hi] * 0.1 * uv * mw[hi];
}
har[t] = acc.tanh();
}
har
}
fn hift_stft(har: &[f32]) -> (Vec<f32>, usize) {
let (n, hop) = (cfg::HIFT_NFFT, cfg::HIFT_HOP);
let nb = n / 2 + 1;
let l = har.len() as i64;
let frames = har.len() / hop + 1;
let win: Vec<f32> = (0..n)
.map(|i| 0.5 * (1.0 - (2.0 * std::f32::consts::PI * i as f32 / n as f32).cos()))
.collect();
let padded = |i: i64| -> f32 {
let j = if i < 0 {
-i
} else if i >= l {
2 * l - 2 - i
} else {
i
};
har[j.clamp(0, l - 1) as usize]
};
let mut out = vec![0f32; (n + 2) * frames];
for b in 0..nb {
for f in 0..frames {
let start = f as i64 * hop as i64 - (n / 2) as i64;
let (mut re, mut im) = (0f64, 0f64);
for (j, &wj) in win.iter().enumerate() {
let v = (padded(start + j as i64) * wj) as f64;
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / n as f64;
re += v * ang.cos();
im -= v * ang.sin();
}
out[b * frames + f] = re as f32; out[(nb + b) * frames + f] = im as f32; }
}
(out, frames)
}
fn hift_istft(mag: &[f32], phase: &[f32], frames: usize) -> Vec<f32> {
let (n, hop) = (cfg::HIFT_NFFT, cfg::HIFT_HOP);
let nb = n / 2 + 1;
let win: Vec<f64> = (0..n)
.map(|i| 0.5 * (1.0 - (2.0 * std::f64::consts::PI * i as f64 / n as f64).cos()))
.collect();
let full = (frames - 1) * hop + n;
let mut y = vec![0f64; full];
let mut norm = vec![0f64; full];
let inv_n = 1.0 / n as f64;
for f in 0..frames {
for j in 0..n {
let mut acc = 0f64;
for b in 0..nb {
let m = mag[b * frames + f] as f64;
let ph = phase[b * frames + f] as f64;
let (re, im) = (m * ph.cos(), m * ph.sin());
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / n as f64;
let term = re * ang.cos() - im * ang.sin();
acc += if b == 0 || b == nb - 1 {
term
} else {
2.0 * term
};
}
let v = acc * inv_n * win[j];
y[f * hop + j] += v;
norm[f * hop + j] += win[j] * win[j];
}
}
let start = n / 2;
let len = (frames - 1) * hop;
(start..start + len)
.map(|i| (y[i] / norm[i].max(1e-11)) as f32)
.collect()
}
pub struct HiftGan {
conv_pre: (Vec<f32>, Vec<f32>),
ups: Vec<(Vec<f32>, Vec<f32>, usize, usize)>, source_downs: Vec<(Vec<f32>, Vec<f32>, usize, usize, usize)>, source_resblocks: Vec<VocResBlock>,
resblocks: Vec<VocResBlock>,
conv_post: (Vec<f32>, Vec<f32>),
f0_condnet: Vec<(Vec<f32>, Vec<f32>, usize, usize)>, f0_cls_w: Vec<f32>,
f0_cls_b: Vec<f32>,
ms_w: Vec<f32>,
ms_b: f32,
}
impl HiftGan {
pub fn load(st: &StReader, p: &str) -> Result<Self> {
let e = |s: &str| format!("{p}.{s}");
let nb2 = cfg::HIFT_NFFT + 2; let mut ups = Vec::new();
let mut resblocks = Vec::new();
let mut ch = cfg::HIFT_BASE;
for i in 0..3 {
let out = ch / 2;
ups.push((
load_wn(
st,
&e(&format!("ups.{i}")),
ch,
out * cfg::HIFT_UP_KERNELS[i],
)?
.0,
st.f32(&e(&format!("ups.{i}.bias")))?,
ch,
out,
));
for (j, &k) in [3usize, 7, 11].iter().enumerate() {
resblocks.push(VocResBlock::load(
st,
&e(&format!("resblocks.{}", i * 3 + j)),
out,
k,
&[1, 3, 5],
)?);
}
ch = out;
}
let mut source_downs = Vec::new();
let mut source_resblocks = Vec::new();
let su = [15usize, 3, 1];
let sk = [7usize, 7, 11]; let mut sch = cfg::HIFT_BASE;
for i in 0..3 {
let out = sch / 2;
let (k, stride) = if su[i] == 1 {
(1, 1)
} else {
(su[i] * 2, su[i])
};
source_downs.push((
st.f32(&e(&format!("source_downs.{i}.weight")))?,
st.f32(&e(&format!("source_downs.{i}.bias")))?,
out,
k,
stride,
));
source_resblocks.push(VocResBlock::load(
st,
&e(&format!("source_resblocks.{i}")),
out,
sk[i],
&[1, 3, 5],
)?);
sch = out;
}
let mut f0_condnet = Vec::new();
for (idx, &(cin, cout)) in [(80, 512), (512, 512), (512, 512), (512, 512), (512, 512)]
.iter()
.enumerate()
{
let (w, b) = load_wn(
st,
&e(&format!("f0_predictor.condnet.{}", idx * 2)),
cout,
cin * 3,
)?;
f0_condnet.push((w, b, cin, cout));
}
let ms_w = st.f32(&e("m_source.l_linear.weight"))?;
let ms_b = st.f32(&e("m_source.l_linear.bias"))?[0];
Ok(Self {
conv_pre: load_wn(st, &e("conv_pre"), cfg::HIFT_BASE, 80 * 7)?,
ups,
source_downs,
source_resblocks,
resblocks,
conv_post: load_wn(st, &e("conv_post"), nb2, ch * 7)?,
f0_condnet,
f0_cls_w: st.f32(&e("f0_predictor.classifier.weight"))?,
f0_cls_b: st.f32(&e("f0_predictor.classifier.bias"))?,
ms_w,
ms_b,
})
}
fn f0_predict(&self, mel: &[f32], t: usize) -> Vec<f32> {
let mut x = mel.to_vec();
let n = self.f0_condnet.len();
for (idx, (w, b, cin, cout)) in self.f0_condnet.iter().enumerate() {
let (o, _) = conv1d(&x, *cin, t, w, Some(b), *cout, 3, 1, 1, 1);
x = o;
if idx < n - 1 {
for v in x.iter_mut() {
if *v < 0.0 {
*v = v.exp() - 1.0; }
}
}
}
let ch = 512;
(0..t)
.map(|ti| {
let mut acc = self.f0_cls_b[0];
for c in 0..ch {
acc += self.f0_cls_w[c] * x[c * t + ti];
}
acc.abs()
})
.collect()
}
fn decode(&self, mel: &[f32], t_mel: usize, s_stft: &[f32], sframes: usize) -> Vec<f32> {
let nb2 = cfg::HIFT_NFFT + 2;
let (mut x, _) = conv1d(
mel,
80,
t_mel,
&self.conv_pre.0,
Some(&self.conv_pre.1),
cfg::HIFT_BASE,
7,
1,
3,
1,
);
let mut t = t_mel;
for i in 0..3 {
leaky(&mut x, 0.1);
let (w, b, cin, cout) = &self.ups[i];
let (rate, k) = (cfg::HIFT_UP_RATES[i], cfg::HIFT_UP_KERNELS[i]);
let (mut xu, mut tu) =
conv_transpose1d(&x, *cin, t, w, Some(b), *cout, k, rate, (k - rate) / 2);
if i == 2 {
let mut padded = vec![0f32; cout * (tu + 1)];
for c in 0..*cout {
padded[c * (tu + 1)] = xu[c * tu + 1];
padded[c * (tu + 1) + 1..c * (tu + 1) + tu + 1]
.copy_from_slice(&xu[c * tu..c * tu + tu]);
}
xu = padded;
tu += 1;
}
let (sw, sb, scout, sk, sstride) = &self.source_downs[i];
let spad = if *sstride == 1 { 0 } else { sstride / 2 };
let (si, _) = conv1d(
s_stft,
nb2,
sframes,
sw,
Some(sb),
*scout,
*sk,
*sstride,
spad,
1,
);
let si = self.source_resblocks[i].forward(&si, tu);
for (a, b) in xu.iter_mut().zip(&si) {
*a += b;
}
let mut xs = vec![0f32; cout * tu];
for j in 0..3 {
let r = self.resblocks[i * 3 + j].forward(&xu, tu);
for (a, b) in xs.iter_mut().zip(r) {
*a += b;
}
}
for v in xs.iter_mut() {
*v /= 3.0;
}
x = xs;
t = tu;
}
leaky(&mut x, 0.01);
let (xp, _) = conv1d(
&x,
cfg::HIFT_BASE / 8,
t,
&self.conv_post.0,
Some(&self.conv_post.1),
nb2,
7,
1,
3,
1,
);
let nb = cfg::HIFT_NFFT / 2 + 1;
let mut mag = vec![0f32; nb * t];
let mut phase = vec![0f32; nb * t];
for b in 0..nb {
for ti in 0..t {
mag[b * t + ti] = xp[b * t + ti].exp().min(1e2);
phase[b * t + ti] = xp[(nb + b) * t + ti].sin();
}
}
hift_istft(&mag, &phase, t)
}
pub fn generate(&self, mel: &[f32], frames: usize) -> Vec<f32> {
let mut mc = vec![0f32; 80 * frames];
for f in 0..frames {
for m in 0..80 {
mc[m * frames + f] = mel[f * 80 + m];
}
}
let f0 = self.f0_predict(&mc, frames);
let up = cfg::HIFT_UP_RATES.iter().product::<usize>() * cfg::HIFT_HOP; let har = hift_sine_source(&f0, up, cfg::S3GEN_SR as f32, &self.ms_w, self.ms_b);
let (s_stft, sframes) = hift_stft(&har);
let mut wav = self.decode(&mc, frames, &s_stft, sframes);
for v in wav.iter_mut() {
*v = v.clamp(-0.99, 0.99);
}
wav
}
}
pub struct S3Gen {
encoder: UpConformer,
pub campplus: CampPlus,
decoder: Decoder,
vocoder: HiftGan,
spk_w: Vec<f32>, spk_b: Vec<f32>,
}
impl S3Gen {
pub fn load(path: &Path) -> Result<Self> {
let st = StReader::open(path).with_context(|| format!("open s3gen {}", path.display()))?;
Ok(Self {
encoder: UpConformer::load(&st, "flow")?,
campplus: CampPlus::load(&st, "speaker_encoder")?,
decoder: Decoder::load(&st, "flow.decoder.estimator")?,
vocoder: HiftGan::load(&st, "mel2wav")?,
spk_w: st.mat(
"flow.spk_embed_affine_layer.weight",
cfg::S3GEN_MELS,
cfg::S3GEN_SPK_EMBED_DIM,
)?,
spk_b: st.f32("flow.spk_embed_affine_layer.bias")?,
})
}
pub fn project_spk(&self, xvector: &[f32]) -> Vec<f32> {
let norm = xvector.iter().map(|v| v * v).sum::<f32>().sqrt().max(1e-12);
let xn: Vec<f32> = xvector.iter().map(|v| v / norm).collect();
add_bias(
matvec(&self.spk_w, &xn, cfg::S3GEN_MELS, cfg::S3GEN_SPK_EMBED_DIM),
&self.spk_b,
)
}
pub fn token_to_wav(
&self,
prompt_tokens: &[u32],
gen_tokens: &[u32],
xvector: &[f32],
prompt_mel: &[f32],
seed: u64,
) -> Vec<f32> {
let mels = cfg::S3GEN_MELS;
let spks = self.project_spk(xvector);
let prompt_frames = prompt_mel.len() / mels;
let tok1 = prompt_tokens.len().min(prompt_frames / 2);
let mut all = prompt_tokens[..tok1].to_vec();
all.extend_from_slice(gen_tokens);
let t_enc = std::time::Instant::now();
let mu_tm = self.encoder.forward(&all);
let d_enc = t_enc.elapsed();
let mel_len = 2 * all.len();
let mel_len1 = 2 * tok1;
let mut mu = vec![0f32; mels * mel_len];
let mut cond = vec![0f32; mels * mel_len];
for f in 0..mel_len {
for m in 0..mels {
mu[m * mel_len + f] = mu_tm[f * mels + m];
if f < mel_len1 {
cond[m * mel_len + f] = prompt_mel[f * mels + m];
}
}
}
let z = cfm_seed_noise(mels * mel_len, seed);
let est = |x: &[f32], muc: &[f32], sc: &[f32], cc: &[f32], t: f32| {
self.decoder.forward(x, muc, sc, cc, t, mel_len)
};
let t_cfm = std::time::Instant::now();
let feat = CfmSolver::s3gen().solve(&z, &mu, &spks, &cond, est);
let d_cfm = t_cfm.elapsed();
let gen_len = mel_len - mel_len1;
let mut gen_mel = vec![0f32; gen_len * mels];
for f in 0..gen_len {
for m in 0..mels {
gen_mel[f * mels + m] = feat[m * mel_len + (mel_len1 + f)];
}
}
let t_voc = std::time::Instant::now();
let mut wav = self.vocoder.generate(&gen_mel, gen_len);
if std::env::var("CB_PROFILE").is_ok() {
eprintln!(
" [S3Gen] UpConformer {:.0} ms | CFM {:.0} ms | HiFT {:.0} ms",
d_enc.as_secs_f64() * 1000.0,
d_cfm.as_secs_f64() * 1000.0,
t_voc.elapsed().as_secs_f64() * 1000.0
);
}
let n_trim = cfg::S3GEN_SR as usize / 50;
for (i, v) in wav.iter_mut().take(2 * n_trim).enumerate() {
*v *= if i < n_trim {
0.0
} else {
let t = (i - n_trim) as f32 / (n_trim - 1) as f32; ((std::f32::consts::PI * (1.0 - t)).cos() + 1.0) / 2.0
};
}
wav
}
}
pub fn resample_sinc(x: &[f32], sr_in: u32, sr_out: u32) -> Vec<f32> {
if sr_in == sr_out || x.is_empty() {
return x.to_vec();
}
let g = gcd(sr_in, sr_out);
let (orig, new) = ((sr_in / g) as usize, (sr_out / g) as usize);
const LPF_W: f64 = 6.0;
let base = (orig.min(new) as f64) * 0.99;
let width = (LPF_W * orig as f64 / base).ceil() as usize;
let klen = 2 * width + orig;
let mut kernels = vec![0f64; new * klen];
for i in 0..new {
for j in 0..klen {
let idx = (j as f64 - width as f64) / orig as f64;
let mut t = (idx - i as f64 / new as f64) * base;
t = t.clamp(-LPF_W, LPF_W);
let window = (t * std::f64::consts::PI / LPF_W / 2.0).cos().powi(2);
let tp = t * std::f64::consts::PI;
let sinc = if tp == 0.0 { 1.0 } else { tp.sin() / tp };
kernels[i * klen + j] = sinc * window * base / orig as f64;
}
}
let n_out = (x.len() * new).div_ceil(orig);
let mut out = Vec::with_capacity(n_out);
'outer: for k in 0.. {
for i in 0..new {
if out.len() == n_out {
break 'outer;
}
let mut acc = 0f64;
for j in 0..klen {
let src = k * orig + j;
if src >= width && src - width < x.len() {
acc += x[src - width] as f64 * kernels[i * klen + j];
}
}
out.push(acc as f32);
}
}
out
}
fn gcd(a: u32, b: u32) -> u32 {
if b == 0 { a } else { gcd(b, a % b) }
}
pub fn resample_linear(x: &[f32], sr_in: u32, sr_out: u32) -> Vec<f32> {
if sr_in == sr_out || x.is_empty() {
return x.to_vec();
}
let ratio = sr_in as f64 / sr_out as f64;
let n = (x.len() as f64 / ratio) as usize;
(0..n)
.map(|i| {
let p = i as f64 * ratio;
let a = p as usize;
let t = (p - a as f64) as f32;
x[a] * (1.0 - t) + x[(a + 1).min(x.len() - 1)] * t
})
.collect()
}
pub fn fbank_kaldi(audio_16k: &[f32]) -> (Vec<f32>, usize) {
const WIN: usize = 400;
const HOP: usize = 160;
const NFFT: usize = 512;
const NBINS: usize = NFFT / 2; const NMEL: usize = 80;
let n = audio_16k.len();
let frames = if n < WIN { 0 } else { 1 + (n - WIN) / HOP };
assert!(
frames > 0,
"fbank_kaldi: input shorter than one 25 ms frame"
);
let win: Vec<f64> = (0..WIN)
.map(|i| {
(0.5 - 0.5 * (2.0 * std::f64::consts::PI * i as f64 / (WIN - 1) as f64).cos())
.powf(0.85)
})
.collect();
let mel = |f: f64| 1127.0 * (1.0 + f / 700.0).ln();
let (mlo, mhi) = (mel(20.0), mel(8000.0));
let delta = (mhi - mlo) / (NMEL + 1) as f64;
let bin_mel: Vec<f64> = (0..NBINS)
.map(|b| mel(b as f64 * 16000.0 / NFFT as f64))
.collect();
let mut fb = vec![0f32; NMEL * NBINS];
for m in 0..NMEL {
let (l, c, r) = (
mlo + m as f64 * delta,
mlo + (m + 1) as f64 * delta,
mlo + (m + 2) as f64 * delta,
);
for b in 0..NBINS {
let up = (bin_mel[b] - l) / (c - l);
let down = (r - bin_mel[b]) / (r - c);
fb[m * NBINS + b] = up.min(down).max(0.0) as f32;
}
}
let mut out = vec![0f32; frames * NMEL];
let mut fr = vec![0f64; NFFT];
let mut power = vec![0f32; NBINS];
for f in 0..frames {
let s = f * HOP;
let frame = &audio_16k[s..s + WIN];
let mean: f64 = frame.iter().map(|&v| v as f64).sum::<f64>() / WIN as f64; for j in 0..WIN {
let x = frame[j] as f64 - mean;
let xp = frame[j.saturating_sub(1)] as f64 - mean; fr[j] = (x - 0.97 * xp) * win[j];
}
fr[WIN..NFFT].fill(0.0);
for (b, p) in power.iter_mut().enumerate() {
let (mut re, mut im) = (0f64, 0f64);
for (j, &v) in fr.iter().enumerate().take(WIN) {
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / NFFT as f64;
re += v * ang.cos();
im -= v * ang.sin();
}
*p = (re * re + im * im) as f32;
}
for m in 0..NMEL {
let row = &fb[m * NBINS..(m + 1) * NBINS];
let e: f32 = row.iter().zip(&power).map(|(w, p)| w * p).sum();
out[f * NMEL + m] = e.max(f32::EPSILON).ln();
}
}
for m in 0..NMEL {
let mean: f32 = (0..frames).map(|f| out[f * NMEL + m]).sum::<f32>() / frames as f32;
for f in 0..frames {
out[f * NMEL + m] -= mean;
}
}
(out, frames)
}
pub fn mel_s3gen(audio_24k: &[f32]) -> (Vec<f32>, usize) {
let (n_fft, hop, n_mels) = (1920usize, 480usize, 80usize);
let nb = n_fft / 2 + 1;
let fb = librosa_mel(24000.0, n_fft, n_mels, 0.0, 8000.0); let pad = (n_fft - hop) / 2;
let n = audio_24k.len();
let idx = |i: isize| -> f32 {
let j = if i < 0 {
(-i) as usize
} else if i as usize >= n {
2 * (n - 1) - i as usize
} else {
i as usize
};
audio_24k[j]
};
let padded_len = n + 2 * pad;
let frames = if padded_len < n_fft {
1
} else {
(padded_len - n_fft) / hop + 1
};
let win: 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();
let mut out = vec![0f32; frames * n_mels];
out.par_chunks_mut(n_mels)
.enumerate()
.for_each(|(f, orow)| {
let mut mag = vec![0f32; nb];
let start = f as isize * hop as isize - pad as isize;
let fr: Vec<f64> = (0..n_fft)
.map(|j| (idx(start + j as isize) * win[j]) as f64)
.collect();
for (b, mg) in mag.iter_mut().enumerate() {
let (mut re, mut im) = (0f64, 0f64);
for (j, &v) in fr.iter().enumerate() {
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / n_fft as f64;
re += v * ang.cos();
im -= v * ang.sin();
}
*mg = ((re * re + im * im + 1e-9) as f32).sqrt();
}
for (m, o) in orow.iter_mut().enumerate() {
let row = &fb[m * nb..(m + 1) * nb];
let e: f32 = row.iter().zip(&mag).map(|(w, g)| w * g).sum();
*o = e.max(1e-5).ln();
}
});
(out, frames)
}
pub fn punc_norm(text: &str) -> String {
if text.is_empty() {
return "You need to add some text for me to talk.".to_string();
}
let mut s: String = {
let mut cs = text.chars();
let first = cs.next().unwrap();
if first.is_lowercase() {
first.to_uppercase().collect::<String>() + cs.as_str()
} else {
text.to_string()
}
};
if s.contains(' ') {
s = s.split_whitespace().collect::<Vec<_>>().join(" ");
}
for (old, new) in [
("...", ", "),
("…", ", "),
(":", ","),
(" - ", ", "),
(";", ", "),
("—", "-"),
("–", "-"),
(" ,", ","),
("“", "\""),
("”", "\""),
("‘", "'"),
("’", "'"),
] {
s = s.replace(old, new);
}
if !s.ends_with(['.', '!', '?', '-', ',']) {
s.push('.');
}
s
}
pub const WATERMARK_ALPHA: f32 = 0.03;
pub const WATERMARK_KEY: u64 = 0x43_48_41_54_54_45_52_42; const WATERMARK_BLOCK: usize = 1024;
fn watermark_carrier(wav: &[f32], key: u64) -> Vec<f32> {
let n = wav.len();
const SMOOTH: usize = 32;
const CARRIER_HZ: f32 = 10_000.0; let mut rng = SplitMix64::new(key);
let prn: Vec<f32> = (0..n + SMOOTH)
.map(|_| if rng.next_f32() < 0.5 { -1.0f32 } else { 1.0 })
.collect();
let taps: Vec<f32> = (0..SMOOTH)
.map(|k| 0.5 - 0.5 * (2.0 * std::f32::consts::PI * k as f32 / (SMOOTH - 1) as f32).cos())
.collect();
let tap_sum: f32 = taps.iter().sum::<f32>().max(1e-12);
let w = 2.0 * std::f32::consts::PI * CARRIER_HZ / cfg::S3GEN_SR as f32;
let mut carrier: Vec<f32> = (0..n)
.map(|i| {
let base: f32 = prn[i..i + SMOOTH]
.iter()
.zip(&taps)
.map(|(p, t)| p * t)
.sum::<f32>()
/ tap_sum;
base * (w * i as f32).cos()
})
.collect();
let crms = (carrier.iter().map(|v| v * v).sum::<f32>() / n.max(1) as f32).sqrt();
if crms > 1e-12 {
for c in carrier.iter_mut() {
*c /= crms;
}
}
let nb = n.div_ceil(WATERMARK_BLOCK).max(1);
let rms: Vec<f32> = (0..nb)
.map(|b| {
let s = b * WATERMARK_BLOCK;
let e = (s + WATERMARK_BLOCK).min(n);
if e <= s {
return 0.0;
}
(wav[s..e].iter().map(|v| v * v).sum::<f32>() / (e - s) as f32).sqrt()
})
.collect();
for (i, c) in carrier.iter_mut().enumerate() {
let b = i / WATERMARK_BLOCK;
let t = (i % WATERMARK_BLOCK) as f32 / WATERMARK_BLOCK as f32;
let e0 = rms[b];
let e1 = rms[(b + 1).min(nb - 1)];
*c *= e0 + (e1 - e0) * t;
}
carrier
}
pub fn watermark_embed(wav: &mut [f32], key: u64) {
let carrier = watermark_carrier(wav, key);
for (s, c) in wav.iter_mut().zip(&carrier) {
*s = (*s + WATERMARK_ALPHA * c).clamp(-1.0, 1.0);
}
}
pub fn watermark_detect(wav: &[f32], key: u64) -> f32 {
if wav.len() < WATERMARK_BLOCK {
return 0.0;
}
let carrier = watermark_carrier(wav, key);
let dot: f64 = wav
.iter()
.zip(&carrier)
.map(|(a, b)| *a as f64 * *b as f64)
.sum();
let ex: f64 = wav
.iter()
.map(|v| (*v as f64) * (*v as f64))
.sum::<f64>()
.sqrt();
let ec: f64 = carrier
.iter()
.map(|v| (*v as f64) * (*v as f64))
.sum::<f64>()
.sqrt();
if ex <= 0.0 || ec <= 0.0 {
return 0.0;
}
(dot / (ex * ec)) as f32
}
#[derive(Clone)]
pub struct RefConditionals {
pub ve_embed: Vec<f32>,
pub t3_prompt: Vec<u32>,
pub s3_prompt: Vec<u32>,
pub xvector: Vec<f32>,
pub prompt_mel: Vec<f32>,
}
pub struct CloneOutput {
pub wav: Vec<f32>,
pub stopped_naturally: bool,
pub tokens: usize,
}
pub struct Chatterbox {
pub t3: T3,
pub ve: VoiceEncoder,
pub s3tok: S3Tokenizer,
pub s3gen: S3Gen,
}
impl Chatterbox {
pub fn load(dir: &Path) -> Result<Self> {
Ok(Self {
t3: T3::load(&dir.join(cfg::FILE_T3_MTL))?,
ve: VoiceEncoder::load(&dir.join(cfg::FILE_VE))?,
s3tok: S3Tokenizer::load(&dir.join(cfg::FILE_S3GEN_MTL))?,
s3gen: S3Gen::load(&dir.join(cfg::FILE_S3GEN_MTL))?,
})
}
pub fn clone_voice(
&self,
ref_audio_24k: &[f32],
text_tokens: &[u32],
max_new: usize,
seed: u64,
) -> Vec<f32> {
self.clone(ref_audio_24k, text_tokens, max_new, seed).wav
}
pub fn embed_ref(&self, ref_audio_24k: &[f32]) -> RefConditionals {
let ref_16k = resample_sinc(ref_audio_24k, 24000, 16000);
let enc_len = (cfg::ENC_COND_SECS * cfg::S3_SR as usize).min(ref_16k.len()); let dec_len24 = (cfg::DEC_COND_SECS * cfg::S3GEN_SR as usize).min(ref_audio_24k.len()); let ve_embed = self.ve.embed(&ref_16k);
let mut t3_prompt = self.s3tok.encode(&ref_16k[..enc_len]);
t3_prompt.truncate(cfg::SPEECH_COND_PROMPT_LEN);
let s3gen_ref_16k = resample_sinc(&ref_audio_24k[..dec_len24], 24000, 16000);
let s3_prompt = self.s3tok.encode(&s3gen_ref_16k);
let (fbank, ff) = fbank_kaldi(&s3gen_ref_16k);
let xvector = self.s3gen.campplus.embed(&fbank, ff);
let (prompt_mel, _) = mel_s3gen(&ref_audio_24k[..dec_len24]);
RefConditionals {
ve_embed,
t3_prompt,
s3_prompt,
xvector,
prompt_mel,
}
}
pub fn clone_from_ref(
&self,
r: &RefConditionals,
text_tokens: &[u32],
max_new: usize,
seed: u64,
) -> CloneOutput {
let cond = self
.t3
.build_cond(&r.ve_embed, &r.t3_prompt, cfg::DEFAULT_EMOTION_ADV);
let generated = self.t3.generate_fast(&cond, text_tokens, max_new, seed);
let mut wav = self.s3gen.token_to_wav(
&r.s3_prompt,
&generated.tokens,
&r.xvector,
&r.prompt_mel,
seed,
);
watermark_embed(&mut wav, WATERMARK_KEY);
CloneOutput {
wav,
stopped_naturally: generated.stopped,
tokens: generated.tokens.len(),
}
}
pub fn clone(
&self,
ref_audio_24k: &[f32],
text_tokens: &[u32],
max_new: usize,
seed: u64,
) -> CloneOutput {
self.clone_from_ref(&self.embed_ref(ref_audio_24k), text_tokens, max_new, seed)
}
}
fn apply_rope(x: &mut [f32], cos: &[f32], sin: &[f32], nh: usize, hd: usize) {
let half = hd / 2;
for head in 0..nh {
let base = head * hd;
let mut rotated = vec![0f32; hd];
for d in 0..half {
rotated[d] = -x[base + half + d];
rotated[half + d] = x[base + d];
}
for d in 0..hd {
x[base + d] = x[base + d] * cos[d] + rotated[d] * sin[d];
}
}
}
fn llama3_inv_freq(head_dim: usize) -> Vec<f32> {
let base = 500_000.0_f64;
let (factor, low_ff, high_ff, orig_max) = (8.0_f64, 1.0_f64, 4.0_f64, 8192.0_f64);
let low_wavelen = orig_max / low_ff;
let high_wavelen = orig_max / high_ff;
let mut out = Vec::with_capacity(head_dim / 2);
for i in (0..head_dim).step_by(2) {
let inv = base.powf(-(i as f64) / head_dim as f64);
let wavelen = 2.0 * std::f64::consts::PI / inv;
let scaled = if wavelen > low_wavelen {
inv / factor
} else if wavelen < high_wavelen {
inv
} else {
let smooth = (orig_max / wavelen - low_ff) / (high_ff - low_ff);
(1.0 - smooth) * inv / factor + smooth * inv
};
out.push(scaled as f32);
}
out
}
const VE_FREQ: usize = cfg::VE_N_FFT / 2 + 1; const VE_STEP: usize = 77;
pub struct VoiceEncoder {
lstm_ih: [Vec<f32>; 3], lstm_hh: [Vec<f32>; 3], lstm_bih: [Vec<f32>; 3],
lstm_bhh: [Vec<f32>; 3],
proj_w: Vec<f32>, proj_b: Vec<f32>, mel_filters: Vec<f32>, window: Vec<f32>, dft_cos: Vec<f64>, dft_sin: Vec<f64>,
}
impl VoiceEncoder {
pub fn load(path: &Path) -> Result<Self> {
let st = StReader::open(path).with_context(|| format!("open VE {}", path.display()))?;
let g = cfg::VE_HIDDEN * 4; let ins = [cfg::VE_NUM_MELS, cfg::VE_HIDDEN, cfg::VE_HIDDEN];
let mut ih: [Vec<f32>; 3] = Default::default();
let mut hh: [Vec<f32>; 3] = Default::default();
let mut bih: [Vec<f32>; 3] = Default::default();
let mut bhh: [Vec<f32>; 3] = Default::default();
for l in 0..3 {
ih[l] = st.mat(&format!("lstm.weight_ih_l{l}"), g, ins[l])?;
hh[l] = st.mat(&format!("lstm.weight_hh_l{l}"), g, cfg::VE_HIDDEN)?;
bih[l] = st.f32(&format!("lstm.bias_ih_l{l}"))?;
bhh[l] = st.f32(&format!("lstm.bias_hh_l{l}"))?;
}
let n_fft = cfg::VE_N_FFT;
let window: Vec<f32> = (0..n_fft)
.map(|i| {
(0.5 - 0.5 * (2.0 * std::f64::consts::PI * i as f64 / n_fft as f64).cos()) as f32
})
.collect();
let mut dft_cos = vec![0f64; VE_FREQ * n_fft];
let mut dft_sin = vec![0f64; VE_FREQ * n_fft];
for b in 0..VE_FREQ {
for j in 0..n_fft {
let ang = -2.0 * std::f64::consts::PI * b as f64 * j as f64 / n_fft as f64;
dft_cos[b * n_fft + j] = ang.cos();
dft_sin[b * n_fft + j] = ang.sin();
}
}
Ok(Self {
lstm_ih: ih,
lstm_hh: hh,
lstm_bih: bih,
lstm_bhh: bhh,
proj_w: st.mat("proj.weight", cfg::VE_HIDDEN, cfg::VE_HIDDEN)?,
proj_b: st.f32("proj.bias")?,
mel_filters: librosa_mel(
cfg::VE_SR as f64,
cfg::VE_N_FFT,
cfg::VE_NUM_MELS,
0.0,
cfg::VE_FMAX as f64,
),
window,
dft_cos,
dft_sin,
})
}
fn mel(&self, audio: &[f32]) -> (Vec<f32>, usize) {
let (n_fft, hop, pad) = (cfg::VE_N_FFT, cfg::VE_HOP, cfg::VE_N_FFT / 2);
let mut sig = Vec::with_capacity(audio.len() + 2 * pad);
for i in 0..pad {
sig.push(audio[(pad - i).min(audio.len().saturating_sub(1))]);
}
sig.extend_from_slice(audio);
for i in 0..pad {
let idx = audio.len().saturating_sub(2 + i);
sig.push(audio[idx.min(audio.len().saturating_sub(1))]);
}
let t = 1 + audio.len() / hop;
let mut mel = vec![0f32; t * cfg::VE_NUM_MELS];
let mut frame = vec![0f32; n_fft];
let mut power = vec![0f64; VE_FREQ];
for f in 0..t {
let start = f * hop;
for j in 0..n_fft {
frame[j] = sig.get(start + j).copied().unwrap_or(0.0) * self.window[j];
}
for (b, p) in power.iter_mut().enumerate() {
let cb = &self.dft_cos[b * n_fft..][..n_fft];
let sb = &self.dft_sin[b * n_fft..][..n_fft];
let (mut re, mut im) = (0f64, 0f64);
for j in 0..n_fft {
let x = frame[j] as f64;
re += x * cb[j];
im += x * sb[j];
}
*p = re * re + im * im; }
for m in 0..cfg::VE_NUM_MELS {
let fb = &self.mel_filters[m * VE_FREQ..][..VE_FREQ];
let acc: f64 = (0..VE_FREQ).map(|b| fb[b] as f64 * power[b]).sum();
mel[f * cfg::VE_NUM_MELS + m] = acc as f32;
}
}
(mel, t)
}
fn lstm_layer(&self, input: &[f32], t: usize, in_dim: usize, l: usize) -> Vec<f32> {
let h = cfg::VE_HIDDEN;
let (wih, whh, bih, bhh) = (
&self.lstm_ih[l],
&self.lstm_hh[l],
&self.lstm_bih[l],
&self.lstm_bhh[l],
);
let mut hs = vec![0f32; t * h];
let mut hprev = vec![0f32; h];
let mut c = vec![0f32; h];
for step in 0..t {
let x = &input[step * in_dim..step * in_dim + in_dim];
for gi in 0..h {
let (i_row, f_row, g_row, o_row) = (gi, h + gi, 2 * h + gi, 3 * h + gi);
let mut gate = [0f32; 4];
for (slot, &row) in [i_row, f_row, g_row, o_row].iter().enumerate() {
let wih_r = &wih[row * in_dim..row * in_dim + in_dim];
let whh_r = &whh[row * h..row * h + h];
let mut acc = bih[row] + bhh[row];
for j in 0..in_dim {
acc += wih_r[j] * x[j];
}
for j in 0..h {
acc += whh_r[j] * hprev[j];
}
gate[slot] = acc;
}
let i = sigmoid(gate[0]);
let f = sigmoid(gate[1]);
let g = gate[2].tanh();
let o = sigmoid(gate[3]);
let cell = f * c[gi] + i * g;
c[gi] = cell;
hs[step * h + gi] = o * cell.tanh();
}
hprev.copy_from_slice(&hs[step * h..step * h + h]);
}
hs
}
fn embed_partial(&self, mel_partial: &[f32], frames: usize) -> Vec<f32> {
let h = cfg::VE_HIDDEN;
let l0 = self.lstm_layer(mel_partial, frames, cfg::VE_NUM_MELS, 0);
let l1 = self.lstm_layer(&l0, frames, h, 1);
let l2 = self.lstm_layer(&l1, frames, h, 2);
let last = &l2[(frames - 1) * h..frames * h];
let mut e = add_bias(matvec(&self.proj_w, last, h, h), &self.proj_b);
for v in &mut e {
*v = v.max(0.0); }
l2_normalize(&mut e);
e
}
pub fn trim_silence(audio: &[f32], top_db: f32) -> &[f32] {
let (frame, hop, pad) = (2048usize, 512usize, 1024usize);
let n = audio.len();
let n_frames = (n + 2 * pad).saturating_sub(frame) / hop + 1;
let mse: Vec<f64> = (0..n_frames)
.map(|i| {
let mut acc = 0f64;
for j in 0..frame {
let idx = (i * hop + j) as isize - pad as isize;
if idx >= 0 && (idx as usize) < n {
let v = audio[idx as usize] as f64;
acc += v * v;
}
}
acc / frame as f64
})
.collect();
let refv = mse.iter().cloned().fold(1e-10f64, f64::max);
let thresh = |p: f64| 10.0 * (p.max(1e-10) / refv).log10() > -(top_db as f64);
let first = mse.iter().position(|&p| thresh(p));
let (Some(f0), Some(f1)) = (first, mse.iter().rposition(|&p| thresh(p))) else {
return audio;
};
&audio[(f0 * hop).min(n)..((f1 + 1) * hop).min(n)]
}
pub fn embed(&self, audio: &[f32]) -> Vec<f32> {
let audio = Self::trim_silence(audio, 20.0);
let (mel, t) = self.mel(audio);
let win = cfg::VE_PARTIAL_FRAMES;
let (mut n_wins, remainder) = if t >= win {
let span = t - win + VE_STEP;
(span / VE_STEP, span % VE_STEP)
} else {
(0, 0)
};
if n_wins == 0 || (remainder + (win - VE_STEP)) as f32 / win as f32 >= 0.8 {
n_wins += 1;
}
let target_n = win + VE_STEP * (n_wins - 1);
let mut mel = mel;
if target_n > t {
mel.resize(target_n * cfg::VE_NUM_MELS, 0.0);
}
let mut acc = vec![0f32; cfg::VE_HIDDEN];
for i in 0..n_wins {
let s = i * VE_STEP;
let part = &mel[s * cfg::VE_NUM_MELS..(s + win) * cfg::VE_NUM_MELS];
let e = self.embed_partial(part, win);
for (a, v) in acc.iter_mut().zip(&e) {
*a += v;
}
}
for v in &mut acc {
*v /= n_wins as f32;
}
l2_normalize(&mut acc);
acc
}
}
fn sigmoid(v: f32) -> f32 {
1.0 / (1.0 + (-v).exp())
}
fn l2_normalize(v: &mut [f32]) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-12);
for x in v.iter_mut() {
*x /= norm;
}
}
#[cfg(feature = "cli")]
pub struct MtlTokenizer {
inner: tokenizers::Tokenizer,
}
#[cfg(feature = "cli")]
impl MtlTokenizer {
pub fn load(path: &Path) -> Result<Self> {
let inner = tokenizers::Tokenizer::from_file(path)
.map_err(|e| anyhow::anyhow!("load tokenizer {}: {e}", path.display()))?;
Ok(Self { inner })
}
pub fn encode(&self, text: &str) -> Result<Vec<u32>> {
let prepared = text.replace(' ', "[SPACE]");
let enc = self
.inner
.encode(prepared, false)
.map_err(|e| anyhow::anyhow!("encode: {e}"))?;
Ok(enc.get_ids().to_vec())
}
pub fn decode(&self, ids: &[u32]) -> Result<String> {
self.inner
.decode(ids, false)
.map_err(|e| anyhow::anyhow!("decode: {e}"))
.map(|s| s.replace("[SPACE]", " "))
}
pub fn vocab_size(&self) -> usize {
self.inner.get_vocab_size(true)
}
}
fn gelu(v: f32) -> f32 {
0.5 * v * (1.0 + libm::erff(v * std::f32::consts::FRAC_1_SQRT_2))
}
fn layer_norm(x: &[f32], w: &[f32], b: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let mean = x.iter().sum::<f32>() / n;
let var = x.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / n;
let inv = 1.0 / (var + eps).sqrt();
(0..x.len())
.map(|i| (x[i] - mean) * inv * w[i] + b[i])
.collect()
}
fn conv1d(
inp: &[f32],
in_ch: usize,
t_in: usize,
w: &[f32],
bias: Option<&[f32]>,
out_ch: usize,
k: usize,
stride: usize,
pad: usize,
groups: usize,
) -> (Vec<f32>, usize) {
let t_out = (t_in + 2 * pad - k) / stride + 1;
let in_pg = in_ch / groups;
let out_pg = out_ch / groups;
let mut out = vec![0f32; out_ch * t_out];
out.par_chunks_mut(t_out).enumerate().for_each(|(oc, row)| {
let g = oc / out_pg;
let b0 = bias.map(|b| b[oc]).unwrap_or(0.0);
for (ot, o) in row.iter_mut().enumerate() {
let mut acc = b0;
for icg in 0..in_pg {
let ic = g * in_pg + icg;
let wbase = oc * in_pg * k + icg * k;
let ibase = ic * t_in;
for kk in 0..k {
let it = (ot * stride + kk) as isize - pad as isize;
if it >= 0 && (it as usize) < t_in {
acc += w[wbase + kk] * inp[ibase + it as usize];
}
}
}
*o = acc;
}
});
(out, t_out)
}
fn rope_tables(seq: usize, head_dim: usize, inv_freq: &[f32]) -> (Vec<Vec<f32>>, Vec<Vec<f32>>) {
let half = head_dim / 2;
let mut cos = vec![vec![0f32; head_dim]; seq];
let mut sin = vec![vec![0f32; head_dim]; seq];
for p in 0..seq {
for (j, &f) in inv_freq.iter().enumerate() {
let (s, c) = (p as f32 * f).sin_cos();
cos[p][j] = c;
cos[p][j + half] = c;
sin[p][j] = s;
sin[p][j + half] = s;
}
}
(cos, sin)
}
fn hz_to_mel_slaney(f: f64) -> f64 {
let f_sp = 200.0 / 3.0;
let min_log_hz = 1000.0;
let min_log_mel = min_log_hz / f_sp; let logstep = 6.4_f64.ln() / 27.0;
if f >= min_log_hz {
min_log_mel + (f / min_log_hz).ln() / logstep
} else {
f / f_sp
}
}
fn mel_to_hz_slaney(m: f64) -> f64 {
let f_sp = 200.0 / 3.0;
let min_log_hz = 1000.0;
let min_log_mel = min_log_hz / f_sp;
let logstep = 6.4_f64.ln() / 27.0;
if m >= min_log_mel {
min_log_hz * (logstep * (m - min_log_mel)).exp()
} else {
f_sp * m
}
}
pub(crate) fn librosa_mel(sr: f64, n_fft: usize, n_mels: usize, fmin: f64, fmax: f64) -> Vec<f32> {
let n_freq = n_fft / 2 + 1;
let fftfreqs: Vec<f64> = (0..n_freq).map(|i| i as f64 * sr / n_fft as f64).collect();
let (mmin, mmax) = (hz_to_mel_slaney(fmin), hz_to_mel_slaney(fmax));
let mel_pts: Vec<f64> = (0..n_mels + 2)
.map(|i| mel_to_hz_slaney(mmin + (mmax - mmin) * i as f64 / (n_mels + 1) as f64))
.collect();
let fdiff: Vec<f64> = (0..n_mels + 1)
.map(|i| mel_pts[i + 1] - mel_pts[i])
.collect();
let mut w = vec![0f32; n_mels * n_freq];
for m in 0..n_mels {
let enorm = 2.0 / (mel_pts[m + 2] - mel_pts[m]);
for (f, &ff) in fftfreqs.iter().enumerate() {
let lower = (ff - mel_pts[m]) / fdiff[m];
let upper = (mel_pts[m + 2] - ff) / fdiff[m + 1];
let val = lower.min(upper).max(0.0) * enorm;
w[m * n_freq + f] = val as f32;
}
}
w
}
struct S3Block {
attn_ln_w: Vec<f32>,
attn_ln_b: Vec<f32>,
q_w: Vec<f32>,
q_b: Vec<f32>,
k_w: Vec<f32>,
k_b: Vec<f32>,
v_w: Vec<f32>,
v_b: Vec<f32>,
out_w: Vec<f32>,
out_b: Vec<f32>,
fsmn: Vec<f32>, mlp_ln_w: Vec<f32>,
mlp_ln_b: Vec<f32>,
mlp0_w: Vec<f32>,
mlp0_b: Vec<f32>,
mlp2_w: Vec<f32>,
mlp2_b: Vec<f32>,
}
pub struct S3Tokenizer {
conv1_w: Vec<f32>,
conv1_b: Vec<f32>,
conv2_w: Vec<f32>,
conv2_b: Vec<f32>,
blocks: Vec<S3Block>,
proj_down_w: Vec<f32>, proj_down_b: Vec<f32>, inv_freq: Vec<f32>, mel_filters: Vec<f32>, window: Vec<f32>, dft_cos: Vec<f64>, dft_sin: Vec<f64>,
}
const S3_DIM: usize = 1280;
const S3_HEADS: usize = 20;
const S3_HEAD_DIM: usize = S3_DIM / S3_HEADS; const S3_MEL: usize = 128;
const S3_NFFT: usize = 400;
const S3_HOP: usize = 160;
const S3_FREQ: usize = S3_NFFT / 2 + 1; const S3_FSMN_K: usize = 31;
impl S3Tokenizer {
pub fn load(path: &Path) -> Result<Self> {
let st = StReader::open(path).with_context(|| format!("open s3gen {}", path.display()))?;
let mut blocks = Vec::with_capacity(6);
for l in 0..6 {
let p = format!("tokenizer.encoder.blocks.{l}");
blocks.push(S3Block {
attn_ln_w: st.f32(&format!("{p}.attn_ln.weight"))?,
attn_ln_b: st.f32(&format!("{p}.attn_ln.bias"))?,
q_w: st.mat(&format!("{p}.attn.query.weight"), S3_DIM, S3_DIM)?,
q_b: st.f32(&format!("{p}.attn.query.bias"))?,
k_w: st.mat(&format!("{p}.attn.key.weight"), S3_DIM, S3_DIM)?,
k_b: vec![0.0; S3_DIM], v_w: st.mat(&format!("{p}.attn.value.weight"), S3_DIM, S3_DIM)?,
v_b: st.f32(&format!("{p}.attn.value.bias"))?,
out_w: st.mat(&format!("{p}.attn.out.weight"), S3_DIM, S3_DIM)?,
out_b: st.f32(&format!("{p}.attn.out.bias"))?,
fsmn: st.f32(&format!("{p}.attn.fsmn_block.weight"))?,
mlp_ln_w: st.f32(&format!("{p}.mlp_ln.weight"))?,
mlp_ln_b: st.f32(&format!("{p}.mlp_ln.bias"))?,
mlp0_w: st.mat(&format!("{p}.mlp.0.weight"), 5120, S3_DIM)?,
mlp0_b: st.f32(&format!("{p}.mlp.0.bias"))?,
mlp2_w: st.mat(&format!("{p}.mlp.2.weight"), S3_DIM, 5120)?,
mlp2_b: st.f32(&format!("{p}.mlp.2.bias"))?,
});
}
let window: Vec<f32> = (0..S3_NFFT)
.map(|i| 0.5 - 0.5 * (2.0 * std::f64::consts::PI * i as f64 / S3_NFFT as f64).cos())
.map(|v| v as f32)
.collect();
let mut dft_cos = vec![0f64; S3_FREQ * S3_NFFT];
let mut dft_sin = vec![0f64; S3_FREQ * S3_NFFT];
for b in 0..S3_FREQ {
for j in 0..S3_NFFT {
let ang = -2.0 * std::f64::consts::PI * b as f64 * j as f64 / S3_NFFT as f64;
dft_cos[b * S3_NFFT + j] = ang.cos();
dft_sin[b * S3_NFFT + j] = ang.sin();
}
}
let inv_freq: Vec<f32> = (0..S3_HEAD_DIM / 2)
.map(|i| (10000f64).powf(-((2 * i) as f64) / S3_HEAD_DIM as f64) as f32)
.collect();
Ok(Self {
conv1_w: st.f32("tokenizer.encoder.conv1.weight")?,
conv1_b: st.f32("tokenizer.encoder.conv1.bias")?,
conv2_w: st.f32("tokenizer.encoder.conv2.weight")?,
conv2_b: st.f32("tokenizer.encoder.conv2.bias")?,
blocks,
proj_down_w: st.mat(
"tokenizer.quantizer._codebook.project_down.weight",
8,
S3_DIM,
)?,
proj_down_b: st.f32("tokenizer.quantizer._codebook.project_down.bias")?,
inv_freq,
mel_filters: librosa_mel(S3_SR_F, S3_NFFT, S3_MEL, 0.0, 8000.0),
window,
dft_cos,
dft_sin,
})
}
fn log_mel(&self, audio: &[f32]) -> (Vec<f32>, usize) {
let pad = S3_NFFT / 2; let mut sig = Vec::with_capacity(audio.len() + 2 * pad);
for i in 0..pad {
sig.push(audio[(pad - i).min(audio.len().saturating_sub(1))]);
}
sig.extend_from_slice(audio);
for i in 0..pad {
let idx = audio.len().saturating_sub(2 + i);
sig.push(audio[idx.min(audio.len().saturating_sub(1))]);
}
let n_frames_full = 1 + (sig.len() - S3_NFFT) / S3_HOP;
let t_mel = n_frames_full.saturating_sub(1); let mut power = vec![0f64; S3_FREQ * t_mel]; let mut frame = vec![0f32; S3_NFFT];
for f in 0..t_mel {
let start = f * S3_HOP;
for j in 0..S3_NFFT {
frame[j] = sig[start + j] * self.window[j];
}
for b in 0..S3_FREQ {
let cb = &self.dft_cos[b * S3_NFFT..][..S3_NFFT];
let sb = &self.dft_sin[b * S3_NFFT..][..S3_NFFT];
let mut re = 0f64;
let mut im = 0f64;
for j in 0..S3_NFFT {
let x = frame[j] as f64;
re += x * cb[j];
im += x * sb[j];
}
power[b * t_mel + f] = re * re + im * im;
}
}
let mut mel = vec![0f32; S3_MEL * t_mel];
for m in 0..S3_MEL {
let fb = &self.mel_filters[m * S3_FREQ..][..S3_FREQ];
for t in 0..t_mel {
let mut acc = 0f64;
for b in 0..S3_FREQ {
acc += fb[b] as f64 * power[b * t_mel + t];
}
mel[m * t_mel + t] = acc.max(1e-10).log10() as f32;
}
}
let maxv = mel.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let floor = maxv - 8.0;
for v in &mut mel {
*v = (v.max(floor) + 4.0) / 4.0;
}
(mel, t_mel)
}
pub fn encode(&self, audio: &[f32]) -> Vec<u32> {
let (mel, t_mel) = self.log_mel(audio);
let (c1, t1) = conv1d(
&mel,
S3_MEL,
t_mel,
&self.conv1_w,
Some(&self.conv1_b),
S3_DIM,
3,
2,
1,
1,
);
let c1: Vec<f32> = c1.iter().map(|&v| gelu(v)).collect();
let (c2, t2) = conv1d(
&c1,
S3_DIM,
t1,
&self.conv2_w,
Some(&self.conv2_b),
S3_DIM,
3,
2,
1,
1,
);
let c2: Vec<f32> = c2.iter().map(|&v| gelu(v)).collect();
let seq = t2;
let mut x = vec![0f32; seq * S3_DIM];
for t in 0..seq {
for d in 0..S3_DIM {
x[t * S3_DIM + d] = c2[d * seq + t];
}
}
let (cos, sin) = rope_tables(seq, S3_HEAD_DIM, &self.inv_freq);
for blk in &self.blocks {
self.block_forward(blk, &mut x, seq, &cos, &sin);
}
const FSQ_SCALE: f32 = 0.999_000_012_874_603_3;
let mut codes = Vec::with_capacity(seq);
for t in 0..seq {
let h = matvec(
&self.proj_down_w,
&x[t * S3_DIM..t * S3_DIM + S3_DIM],
8,
S3_DIM,
);
let mut code = 0u32;
let mut pow = 1u32;
for (d, &hd) in h.iter().enumerate() {
let level = round_half_even((hd + self.proj_down_b[d]).tanh() * FSQ_SCALE) + 1.0;
code += (level as u32) * pow;
if d + 1 < 8 {
pow *= 3;
}
}
codes.push(code);
}
codes
}
fn block_forward(
&self,
blk: &S3Block,
x: &mut [f32],
seq: usize,
cos: &[Vec<f32>],
sin: &[Vec<f32>],
) {
let d = S3_DIM;
let normed: Vec<f32> = (0..seq)
.flat_map(|s| layer_norm(&x[s * d..s * d + d], &blk.attn_ln_w, &blk.attn_ln_b, 1e-5))
.collect();
let mut q = vec![0f32; seq * d];
let mut k = vec![0f32; seq * d];
let mut v = vec![0f32; seq * d];
for s in 0..seq {
let xn = &normed[s * d..s * d + d];
let (mut qs, mut ks) = (
add_bias(matvec(&blk.q_w, xn, d, d), &blk.q_b),
add_bias(matvec(&blk.k_w, xn, d, d), &blk.k_b),
);
let vs = add_bias(matvec(&blk.v_w, xn, d, d), &blk.v_b);
apply_rope(&mut qs, &cos[s], &sin[s], S3_HEADS, S3_HEAD_DIM);
apply_rope(&mut ks, &cos[s], &sin[s], S3_HEADS, S3_HEAD_DIM);
q[s * d..s * d + d].copy_from_slice(&qs);
k[s * d..s * d + d].copy_from_slice(&ks);
v[s * d..s * d + d].copy_from_slice(&vs);
}
let mut v_cm = vec![0f32; d * seq]; for s in 0..seq {
for c in 0..d {
v_cm[c * seq + s] = v[s * d + c];
}
}
let (fsm_cm, _) = conv1d(
&v_cm,
d,
seq,
&blk.fsmn,
None,
d,
S3_FSMN_K,
1,
(S3_FSMN_K - 1) / 2,
d,
);
let scale = 1.0 / (S3_HEAD_DIM as f32).sqrt();
for s in 0..seq {
let mut wv = vec![0f32; d];
for head in 0..S3_HEADS {
let off = head * S3_HEAD_DIM;
let qh = &q[s * d + off..s * d + off + S3_HEAD_DIM];
let mut scores = vec![0f32; seq];
for (t, sc) in scores.iter_mut().enumerate() {
let kh = &k[t * d + off..t * d + off + S3_HEAD_DIM];
*sc = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_inplace(&mut scores);
for (t, &wt) in scores.iter().enumerate() {
let vh = &v[t * d + off..t * d + off + S3_HEAD_DIM];
for dd in 0..S3_HEAD_DIM {
wv[off + dd] += wt * vh[dd];
}
}
}
let out = add_bias(matvec(&blk.out_w, &wv, d, d), &blk.out_b);
for c in 0..d {
let fsm = fsm_cm[c * seq + s] + v[s * d + c]; x[s * d + c] += out[c] + fsm;
}
}
for s in 0..seq {
let xn = layer_norm(&x[s * d..s * d + d], &blk.mlp_ln_w, &blk.mlp_ln_b, 1e-5);
let h0 = add_bias(matvec(&blk.mlp0_w, &xn, 5120, d), &blk.mlp0_b);
let act: Vec<f32> = h0.iter().map(|&v| gelu(v)).collect();
let h2 = add_bias(matvec(&blk.mlp2_w, &act, d, 5120), &blk.mlp2_b);
for c in 0..d {
x[s * d + c] += h2[c];
}
}
}
}
const S3_SR_F: f64 = 16000.0;
fn add_bias(mut v: Vec<f32>, b: &[f32]) -> Vec<f32> {
for (x, bb) in v.iter_mut().zip(b) {
*x += bb;
}
v
}
fn round_half_even(x: f32) -> f32 {
let r = x.round();
if (x - x.floor() - 0.5).abs() < f32::EPSILON {
let f = x.floor();
if (f as i64) % 2 == 0 { f } else { f + 1.0 }
} else {
r
}
}