use std::collections::HashMap;
use std::path::Path;
use anyhow::{Context, Result};
use rayon::prelude::*;
use crate::cpu_gemm::{PackedWeight, gemm_packed, gemv_packed};
use crate::weights::LazySt;
pub const FRAME_SIZE: usize = 1920;
const UNK_PENALTY: f32 = 10.0;
pub struct SpUnigram {
pieces: Vec<String>,
scores: Vec<f32>,
kinds: Vec<u8>,
index: HashMap<String, u32>,
byte_ids: [Option<u32>; 256],
unk_id: u32,
min_score: f32,
max_piece_bytes: usize,
add_dummy_prefix: bool,
escape_whitespace: bool,
remove_extra_whitespaces: bool,
}
enum Pb<'a> {
Varint(u64),
Fixed32([u8; 4]),
Bytes(&'a [u8]),
}
fn pb_fields(mut buf: &[u8]) -> Result<Vec<(u32, Pb<'_>)>> {
let mut out = Vec::new();
while !buf.is_empty() {
let (key, rest) = pb_varint(buf)?;
buf = rest;
let (field, wire) = ((key >> 3) as u32, (key & 7) as u8);
match wire {
0 => {
let (v, rest) = pb_varint(buf)?;
buf = rest;
out.push((field, Pb::Varint(v)));
}
1 => {
anyhow::ensure!(buf.len() >= 8, "truncated fixed64");
buf = &buf[8..];
}
2 => {
let (len, rest) = pb_varint(buf)?;
let len = len as usize;
anyhow::ensure!(rest.len() >= len, "truncated length-delimited field");
out.push((field, Pb::Bytes(&rest[..len])));
buf = &rest[len..];
}
5 => {
anyhow::ensure!(buf.len() >= 4, "truncated fixed32");
out.push((field, Pb::Fixed32([buf[0], buf[1], buf[2], buf[3]])));
buf = &buf[4..];
}
w => anyhow::bail!("unsupported protobuf wire type {w}"),
}
}
Ok(out)
}
fn pb_varint(buf: &[u8]) -> Result<(u64, &[u8])> {
let (mut v, mut shift) = (0u64, 0u32);
for (i, b) in buf.iter().enumerate() {
v |= u64::from(b & 0x7f) << shift;
if b & 0x80 == 0 {
return Ok((v, &buf[i + 1..]));
}
shift += 7;
anyhow::ensure!(shift < 64, "varint too long");
}
anyhow::bail!("truncated varint")
}
impl SpUnigram {
pub fn load(path: &Path) -> Result<Self> {
let raw = std::fs::read(path).with_context(|| format!("read {}", path.display()))?;
Self::from_bytes(&raw)
}
pub fn from_bytes(raw: &[u8]) -> Result<Self> {
let top = pb_fields(raw)?;
let (mut pieces, mut scores, mut kinds) = (Vec::new(), Vec::new(), Vec::new());
let mut normalizer = None;
let mut trainer = None;
for (field, val) in &top {
match (field, val) {
(1, Pb::Bytes(b)) => {
let (mut piece, mut score, mut kind) = (None, 0f32, 1u8);
for (f, v) in pb_fields(b)? {
match (f, v) {
(1, Pb::Bytes(s)) => piece = Some(String::from_utf8(s.to_vec())?),
(2, Pb::Fixed32(x)) => score = f32::from_le_bytes(x),
(3, Pb::Varint(k)) => kind = k as u8,
_ => {}
}
}
pieces.push(piece.context("sentencepiece entry without a piece")?);
scores.push(score);
kinds.push(kind);
}
(2, Pb::Bytes(b)) => trainer = Some(*b),
(3, Pb::Bytes(b)) => normalizer = Some(*b),
_ => {}
}
}
anyhow::ensure!(!pieces.is_empty(), "sentencepiece model has no pieces");
if let Some(t) = trainer {
for (f, v) in pb_fields(t)? {
if let (3, Pb::Varint(m)) = (f, v) {
anyhow::ensure!(m == 1, "sentencepiece model_type {m} is not UNIGRAM");
}
}
}
let (mut add_dummy_prefix, mut remove_extra_whitespaces, mut escape_whitespace) =
(true, true, true);
if let Some(n) = normalizer {
for (f, v) in pb_fields(n)? {
match (f, v) {
(1, Pb::Bytes(name)) => {
let name = String::from_utf8_lossy(name).to_string();
anyhow::ensure!(
name == "identity",
"sentencepiece normalizer {name:?} is unsupported (only `identity` \
has no precompiled charsmap to replay)"
);
}
(2, Pb::Bytes(map)) => anyhow::ensure!(
map.is_empty(),
"sentencepiece ships a precompiled charsmap ({} bytes) — unsupported",
map.len()
),
(3, Pb::Varint(v)) => add_dummy_prefix = v != 0,
(4, Pb::Varint(v)) => remove_extra_whitespaces = v != 0,
(5, Pb::Varint(v)) => escape_whitespace = v != 0,
_ => {}
}
}
}
let mut index = HashMap::new();
let mut byte_ids = [None; 256];
let mut unk_id = 0u32;
let mut min_score = f32::INFINITY;
for (id, ((piece, score), kind)) in pieces
.iter()
.zip(scores.iter())
.zip(kinds.iter())
.enumerate()
{
let id = id as u32;
match kind {
1 | 4 => {
index.insert(piece.clone(), id);
min_score = min_score.min(*score);
}
2 => unk_id = id,
6 => {
let hex = piece
.strip_prefix("<0x")
.and_then(|s| s.strip_suffix('>'))
.with_context(|| format!("byte piece {piece:?} is not <0xNN>"))?;
let b = u8::from_str_radix(hex, 16)
.with_context(|| format!("byte piece {piece:?}"))?;
byte_ids[b as usize] = Some(id);
}
_ => {}
}
}
let max_piece_bytes = index.keys().map(String::len).max().unwrap_or(1);
Ok(Self {
pieces,
scores,
kinds,
index,
byte_ids,
unk_id,
min_score,
max_piece_bytes,
add_dummy_prefix,
escape_whitespace,
remove_extra_whitespaces,
})
}
pub fn vocab_size(&self) -> usize {
self.pieces.len()
}
fn normalize(&self, text: &str) -> String {
if text.is_empty() {
return String::new();
}
let mut s = if self.remove_extra_whitespaces {
let mut out = String::with_capacity(text.len());
let mut prev_space = false;
for c in text.chars() {
let is_space = c == ' ';
if !(is_space && prev_space) {
out.push(c);
}
prev_space = is_space;
}
out.trim_end_matches(' ').to_string()
} else {
text.to_string()
};
if self.add_dummy_prefix {
s.insert(0, ' ');
}
if self.escape_whitespace {
s = s.replace(' ', "\u{2581}");
}
s
}
pub fn encode(&self, text: &str) -> Vec<u32> {
let norm = self.normalize(text);
let bytes = norm.as_bytes();
let n = bytes.len();
if n == 0 {
return Vec::new();
}
let mut char_len = vec![0usize; n + 1];
for (i, c) in norm.char_indices() {
char_len[i] = c.len_utf8();
}
let unk_score = self.min_score - UNK_PENALTY;
let mut best = vec![(f32::NEG_INFINITY, usize::MAX, u32::MAX); n + 1];
best[0].0 = 0.0;
for i in 0..n {
if char_len[i] == 0 || best[i].0 == f32::NEG_INFINITY {
continue; }
let mut single_char_piece = false;
let limit = (i + self.max_piece_bytes).min(n);
for j in i + 1..=limit {
if j < n && char_len[j] == 0 {
continue; }
let Some(&id) = self.index.get(&norm[i..j]) else {
continue;
};
if j - i == char_len[i] {
single_char_piece = true;
}
let cand = best[i].0 + self.scores[id as usize];
if cand > best[j].0 {
best[j] = (cand, i, id);
}
}
if !single_char_piece {
let j = i + char_len[i];
let cand = best[i].0 + unk_score;
if cand > best[j].0 {
best[j] = (cand, i, self.unk_id);
}
}
}
let mut ids = Vec::new();
let mut pos = n;
while pos > 0 {
let (_, start, id) = best[pos];
debug_assert_ne!(start, usize::MAX, "unigram lattice has a gap");
if id == self.unk_id {
for b in bytes[start..pos].iter().rev() {
if let Some(bid) = self.byte_ids[*b as usize] {
ids.push(bid);
} else {
ids.push(self.unk_id);
}
}
} else {
ids.push(id);
}
pos = start;
}
ids.reverse();
ids
}
pub fn decode(&self, ids: &[u32]) -> String {
let mut out: Vec<u8> = Vec::new();
for &id in ids {
let id = id as usize;
if id >= self.pieces.len() {
continue;
}
match self.kinds[id] {
6 => {
let hex = &self.pieces[id][3..5];
if let Ok(b) = u8::from_str_radix(hex, 16) {
out.push(b);
}
}
3 => {}
_ => out.extend_from_slice(self.pieces[id].as_bytes()),
}
}
let s = String::from_utf8_lossy(&out).replace('\u{2581}', " ");
if self.add_dummy_prefix {
s.strip_prefix(' ').unwrap_or(&s).to_string()
} else {
s
}
}
}
#[derive(Debug, Clone)]
pub struct Config {
pub d_model: usize,
pub num_heads: usize,
pub num_layers: usize,
pub ffn_dim: usize,
pub latent_dim: usize,
pub n_bins: usize,
pub flow_dim: usize,
pub flow_depth: usize,
pub time_freq_dim: usize,
pub max_period: f32,
pub mimi_dim: usize,
pub mimi_heads: usize,
pub mimi_layers: usize,
pub mimi_ffn: usize,
pub mimi_context: usize,
pub n_filters: usize,
pub ratios: Vec<usize>,
pub resample_stride: usize,
pub sample_rate: usize,
pub frame_rate: f32,
}
impl Config {
fn from_manifest(st: &LazySt) -> Result<Self> {
let dim = |name: &str, axis: usize| -> Result<usize> {
let s = st.shape(name)?;
s.get(axis)
.copied()
.with_context(|| format!("{name} has no axis {axis} (shape {s:?})"))
};
let mut num_layers = 0usize;
while st.has(&format!(
"flow_lm.transformer.layers.{num_layers}.linear1.weight"
)) {
num_layers += 1;
}
anyhow::ensure!(num_layers > 0, "no flow_lm transformer layers found");
let mut mimi_layers = 0usize;
while st.has(&format!(
"mimi.decoder_transformer.transformer.layers.{mimi_layers}.linear1.weight"
)) {
mimi_layers += 1;
}
anyhow::ensure!(mimi_layers > 0, "no mimi decoder transformer layers found");
let d_model = dim("flow_lm.input_linear.weight", 0)?;
let latent_dim = dim("flow_lm.input_linear.weight", 1)?;
let flow_dim = dim("flow_lm.flow_net.input_proj.weight", 0)?;
let mut flow_depth = 0usize;
while st.has(&format!(
"flow_lm.flow_net.res_blocks.{flow_depth}.in_ln.weight"
)) {
flow_depth += 1;
}
let mut ratios = Vec::new();
let mut idx = 2usize; while st.has(&format!("mimi.decoder.model.{idx}.convtr.weight")) {
ratios.push(dim(&format!("mimi.decoder.model.{idx}.convtr.weight"), 2)? / 2);
idx += 3;
}
anyhow::ensure!(!ratios.is_empty(), "no SEANet decoder upsampling stages");
let resample_stride = dim("mimi.downsample.conv.conv.weight", 2)? / 2;
let hop: usize = ratios.iter().product();
let sample_rate = 24_000usize;
Ok(Self {
d_model,
num_heads: d_model / 64,
num_layers,
ffn_dim: dim("flow_lm.transformer.layers.0.linear1.weight", 0)?,
latent_dim,
n_bins: dim("flow_lm.conditioner.embed.weight", 0)? - 1,
flow_dim,
flow_depth,
time_freq_dim: dim("flow_lm.flow_net.time_embed.0.mlp.0.weight", 1)?,
max_period: 10_000.0,
mimi_dim: dim("mimi.decoder.model.0.conv.weight", 0)?,
mimi_heads: 8,
mimi_layers,
mimi_ffn: dim(
"mimi.decoder_transformer.transformer.layers.0.linear1.weight",
0,
)?,
mimi_context: 250,
n_filters: dim("mimi.encoder.model.0.conv.weight", 0)?,
ratios,
resample_stride,
sample_rate,
frame_rate: sample_rate as f32 / (hop * resample_stride) as f32,
})
}
pub fn steps_per_latent(&self) -> usize {
self.resample_stride
}
pub fn frame_size(&self) -> usize {
self.ratios.iter().product::<usize>() * self.resample_stride
}
}
pub const HEAD_DIM: usize = 64;
fn gelu(x: f32) -> f32 {
0.5 * x * (1.0 + libm::erff(x * std::f32::consts::FRAC_1_SQRT_2))
}
fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
fn elu(x: f32) -> f32 {
if x > 0.0 { x } else { x.exp() - 1.0 }
}
pub(crate) struct Linear {
pub(crate) packed: PackedWeight,
pub(crate) b: Option<Vec<f32>>,
pub(crate) out: usize,
inp: usize,
}
impl Linear {
fn load(st: &LazySt, prefix: &str, out: usize, inp: usize, bias: bool) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == out * inp,
"{prefix}.weight has {} elements, expected {out}×{inp}",
w.len()
);
let b = if bias {
let b = st.tensor_f32(&format!("{prefix}.bias"))?;
anyhow::ensure!(b.len() == out, "{prefix}.bias has {} elements", b.len());
Some(b)
} else {
None
};
Ok(Self {
packed: PackedWeight::new(&w, out, inp),
b,
out,
inp,
})
}
fn forward(&self, x: &[f32], m: usize) -> Vec<f32> {
debug_assert_eq!(x.len(), m * self.inp);
let mut y = vec![0f32; m * self.out];
if m == 1 {
gemv_packed(&mut y, x, &self.packed, self.b.as_deref());
} else {
gemm_packed(&mut y, x, &self.packed, m, self.b.as_deref());
}
y
}
}
pub(crate) struct LayerNorm {
pub(crate) w: Option<Vec<f32>>,
pub(crate) b: Option<Vec<f32>>,
eps: f32,
dim: usize,
}
impl LayerNorm {
fn load(st: &LazySt, prefix: &str, dim: usize, eps: f32) -> Result<Self> {
Ok(Self {
w: Some(st.tensor_f32(&format!("{prefix}.weight"))?),
b: Some(st.tensor_f32(&format!("{prefix}.bias"))?),
eps,
dim,
})
}
fn plain(dim: usize, eps: f32) -> Self {
Self {
w: None,
b: None,
eps,
dim,
}
}
fn forward(&self, x: &[f32], m: usize) -> Vec<f32> {
let mut out = vec![0f32; x.len()];
for t in 0..m {
let (src, dst) = (
&x[t * self.dim..(t + 1) * self.dim],
&mut out[t * self.dim..(t + 1) * self.dim],
);
let mean = src.iter().sum::<f32>() / self.dim as f32;
let var = src.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / self.dim as f32;
let inv = 1.0 / (var + self.eps).sqrt();
for (i, d) in dst.iter_mut().enumerate() {
let n = (src[i] - mean) * inv;
*d = match (&self.w, &self.b) {
(Some(w), Some(b)) => n * w[i] + b[i],
_ => n,
};
}
}
out
}
}
struct RmsNormVar {
alpha: Vec<f32>,
eps: f32,
}
impl RmsNormVar {
fn forward(&self, x: &mut [f32]) {
let n = x.len();
let mean = x.iter().sum::<f32>() / n as f32;
let var = x.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / (n - 1) as f32;
let inv = 1.0 / (self.eps + var).sqrt();
for (i, v) in x.iter_mut().enumerate() {
*v *= self.alpha[i] * inv;
}
}
}
#[derive(Clone)]
pub struct Kv {
k: Vec<f32>,
v: Vec<f32>,
offset: usize,
heads: usize,
hd: usize,
}
impl Kv {
pub(crate) fn rows(&self) -> (&[f32], &[f32], usize) {
(&self.k, &self.v, self.offset)
}
fn new(heads: usize, hd: usize) -> Self {
Self {
k: Vec::new(),
v: Vec::new(),
offset: 0,
heads,
hd,
}
}
fn reserve(&mut self, positions: usize) {
let need = positions * self.heads * self.hd;
if self.k.len() < need {
self.k.resize(need, 0.0);
self.v.resize(need, 0.0);
}
}
}
pub(crate) struct Attention {
pub(crate) in_proj: Linear,
pub(crate) out_proj: Linear,
heads: usize,
hd: usize,
context: Option<usize>,
}
impl Attention {
fn load(
st: &LazySt,
prefix: &str,
d_model: usize,
heads: usize,
context: Option<usize>,
) -> Result<Self> {
Ok(Self {
in_proj: Linear::load(
st,
&format!("{prefix}.in_proj"),
3 * d_model,
d_model,
false,
)?,
out_proj: Linear::load(st, &format!("{prefix}.out_proj"), d_model, d_model, false)?,
heads,
hd: d_model / heads,
context,
})
}
fn rope(&self, v: &mut [f32], pos: usize, max_period: f32) {
let half = self.hd / 2;
for h in 0..self.heads {
let row = &mut v[h * self.hd..(h + 1) * self.hd];
for i in 0..half {
let freq = (-(max_period.ln()) * 2.0 * i as f32 / self.hd as f32).exp();
let (s, c) = ((freq * pos as f32).sin(), (freq * pos as f32).cos());
let (r, im) = (row[2 * i], row[2 * i + 1]);
row[2 * i] = r * c - im * s;
row[2 * i + 1] = r * s + im * c;
}
}
}
fn forward(&self, x: &[f32], m: usize, kv: &mut Kv, max_period: f32) -> Vec<f32> {
let d = self.heads * self.hd;
let proj = self.in_proj.forward(x, m);
kv.reserve(kv.offset + m);
let mut q = vec![0f32; m * d];
for t in 0..m {
let pos = kv.offset + t;
let base = t * 3 * d;
q[t * d..(t + 1) * d].copy_from_slice(&proj[base..base + d]);
self.rope(&mut q[t * d..(t + 1) * d], pos, max_period);
let kdst = pos * d;
kv.k[kdst..kdst + d].copy_from_slice(&proj[base + d..base + 2 * d]);
self.rope(&mut kv.k[kdst..kdst + d], pos, max_period);
kv.v[kdst..kdst + d].copy_from_slice(&proj[base + 2 * d..base + 3 * d]);
}
let scale = 1.0 / (self.hd as f32).sqrt();
let mut ctx = vec![0f32; m * d];
for t in 0..m {
let pos = kv.offset + t;
let first = match self.context {
Some(c) => (pos + 1).saturating_sub(c),
None => 0,
};
for h in 0..self.heads {
let qh = &q[t * d + h * self.hd..t * d + (h + 1) * self.hd];
let mut scores = Vec::with_capacity(pos + 1 - first);
let mut top = f32::NEG_INFINITY;
for p in first..=pos {
let kh = &kv.k[p * d + h * self.hd..p * d + (h + 1) * self.hd];
let dot: f32 = qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale;
top = top.max(dot);
scores.push(dot);
}
let mut sum = 0f32;
for s in scores.iter_mut() {
*s = (*s - top).exp();
sum += *s;
}
let out = &mut ctx[t * d + h * self.hd..t * d + (h + 1) * self.hd];
for (i, p) in (first..=pos).enumerate() {
let w = scores[i] / sum;
let vh = &kv.v[p * d + h * self.hd..p * d + (h + 1) * self.hd];
for (o, v) in out.iter_mut().zip(vh) {
*o += w * v;
}
}
}
}
self.out_proj.forward(&ctx, m)
}
}
pub(crate) struct TrLayer {
pub(crate) attn: Attention,
pub(crate) norm1: LayerNorm,
pub(crate) norm2: LayerNorm,
pub(crate) linear1: Linear,
pub(crate) linear2: Linear,
pub(crate) ls1: Option<Vec<f32>>,
pub(crate) ls2: Option<Vec<f32>>,
}
pub(crate) struct Transformer {
pub(crate) layers: Vec<TrLayer>,
d_model: usize,
heads: usize,
max_period: f32,
}
#[derive(Clone)]
pub struct TrState {
kv: Vec<Kv>,
}
impl TrState {
pub fn offset(&self) -> usize {
self.kv.first().map_or(0, |kv| kv.offset)
}
pub(crate) fn layers(&self) -> &[Kv] {
&self.kv
}
}
impl Transformer {
fn load(
st: &LazySt,
prefix: &str,
d_model: usize,
heads: usize,
layers: usize,
ffn: usize,
context: Option<usize>,
layer_scale: bool,
max_period: f32,
) -> Result<Self> {
let mut out = Vec::with_capacity(layers);
for i in 0..layers {
let p = format!("{prefix}.layers.{i}");
out.push(TrLayer {
attn: Attention::load(st, &format!("{p}.self_attn"), d_model, heads, context)?,
norm1: LayerNorm::load(st, &format!("{p}.norm1"), d_model, 1e-5)?,
norm2: LayerNorm::load(st, &format!("{p}.norm2"), d_model, 1e-5)?,
linear1: Linear::load(st, &format!("{p}.linear1"), ffn, d_model, false)?,
linear2: Linear::load(st, &format!("{p}.linear2"), d_model, ffn, false)?,
ls1: layer_scale
.then(|| st.tensor_f32(&format!("{p}.layer_scale_1.scale")))
.transpose()?,
ls2: layer_scale
.then(|| st.tensor_f32(&format!("{p}.layer_scale_2.scale")))
.transpose()?,
});
}
Ok(Self {
layers: out,
d_model,
heads,
max_period,
})
}
fn state(&self) -> TrState {
TrState {
kv: (0..self.layers.len())
.map(|_| Kv::new(self.heads, self.d_model / self.heads))
.collect(),
}
}
fn forward(&self, mut x: Vec<f32>, m: usize, state: &mut TrState) -> Vec<f32> {
for (layer, kv) in self.layers.iter().zip(state.kv.iter_mut()) {
let normed = layer.norm1.forward(&x, m);
let upd = layer.attn.forward(&normed, m, kv, self.max_period);
for (i, u) in upd.iter().enumerate() {
x[i] += layer.ls1.as_ref().map_or(*u, |s| s[i % self.d_model] * u);
}
let normed = layer.norm2.forward(&x, m);
let mut hidden = layer.linear1.forward(&normed, m);
for h in hidden.iter_mut() {
*h = gelu(*h);
}
let upd = layer.linear2.forward(&hidden, m);
for (i, u) in upd.iter().enumerate() {
x[i] += layer.ls2.as_ref().map_or(*u, |s| s[i % self.d_model] * u);
}
}
for kv in state.kv.iter_mut() {
kv.offset += m;
}
x
}
}
struct TimeEmbed {
freqs: Vec<f32>,
l0: Linear,
l2: Linear,
norm: RmsNormVar,
}
impl TimeEmbed {
fn load(st: &LazySt, prefix: &str, dim: usize, freq_dim: usize) -> Result<Self> {
Ok(Self {
freqs: st.tensor_f32(&format!("{prefix}.freqs"))?,
l0: Linear::load(st, &format!("{prefix}.mlp.0"), dim, freq_dim, true)?,
l2: Linear::load(st, &format!("{prefix}.mlp.2"), dim, dim, true)?,
norm: RmsNormVar {
alpha: st.tensor_f32(&format!("{prefix}.mlp.3.alpha"))?,
eps: 1e-5,
},
})
}
fn forward(&self, t: f32) -> Vec<f32> {
let mut emb = vec![0f32; self.freqs.len() * 2];
for (i, f) in self.freqs.iter().enumerate() {
emb[i] = (t * f).cos();
emb[i + self.freqs.len()] = (t * f).sin();
}
let mut h = self.l0.forward(&emb, 1);
for v in h.iter_mut() {
*v = silu(*v);
}
let mut out = self.l2.forward(&h, 1);
self.norm.forward(&mut out);
out
}
}
pub(crate) struct ResBlock {
pub(crate) in_ln: LayerNorm,
pub(crate) m0: Linear,
pub(crate) m2: Linear,
pub(crate) ada: Linear,
}
pub(crate) struct FinalLayer {
pub(crate) norm: LayerNorm,
pub(crate) lin: Linear,
pub(crate) ada: Linear,
}
pub(crate) struct FlowNet {
pub(crate) input_proj: Linear,
pub(crate) time: Vec<TimeEmbed>,
pub(crate) cond: Linear,
pub(crate) blocks: Vec<ResBlock>,
pub(crate) final_layer: FinalLayer,
pub(crate) dim: usize,
}
impl FlowNet {
fn load(st: &LazySt, cfg: &Config) -> Result<Self> {
let p = "flow_lm.flow_net";
let (fd, ld, dm) = (cfg.flow_dim, cfg.latent_dim, cfg.d_model);
let mut blocks = Vec::with_capacity(cfg.flow_depth);
for i in 0..cfg.flow_depth {
let b = format!("{p}.res_blocks.{i}");
blocks.push(ResBlock {
in_ln: LayerNorm::load(st, &format!("{b}.in_ln"), fd, 1e-6)?,
m0: Linear::load(st, &format!("{b}.mlp.0"), fd, fd, true)?,
m2: Linear::load(st, &format!("{b}.mlp.2"), fd, fd, true)?,
ada: Linear::load(st, &format!("{b}.adaLN_modulation.1"), 3 * fd, fd, true)?,
});
}
Ok(Self {
input_proj: Linear::load(st, &format!("{p}.input_proj"), fd, ld, true)?,
time: (0..2)
.map(|i| TimeEmbed::load(st, &format!("{p}.time_embed.{i}"), fd, cfg.time_freq_dim))
.collect::<Result<_>>()?,
cond: Linear::load(st, &format!("{p}.cond_embed"), fd, dm, true)?,
blocks,
final_layer: FinalLayer {
norm: LayerNorm::plain(fd, 1e-6),
lin: Linear::load(st, &format!("{p}.final_layer.linear"), ld, fd, true)?,
ada: Linear::load(
st,
&format!("{p}.final_layer.adaLN_modulation.1"),
2 * fd,
fd,
true,
)?,
},
dim: fd,
})
}
pub(crate) fn time_constant(&self, s: f32, t: f32) -> Vec<f32> {
let (t0, t1) = (self.time[0].forward(s), self.time[1].forward(t));
(0..self.dim).map(|i| (t0[i] + t1[i]) / 2.0).collect()
}
fn forward(&self, c: &[f32], s: f32, t: f32, x: &[f32]) -> Vec<f32> {
let mut h = self.input_proj.forward(x, 1);
let (t0, t1) = (self.time[0].forward(s), self.time[1].forward(t));
let cond = self.cond.forward(c, 1);
let y: Vec<f32> = (0..self.dim)
.map(|i| (t0[i] + t1[i]) / 2.0 + cond[i])
.collect();
let y_act: Vec<f32> = y.iter().map(|v| silu(*v)).collect();
for b in &self.blocks {
let m = b.ada.forward(&y_act, 1);
let (shift, scale, gate) = (
&m[..self.dim],
&m[self.dim..2 * self.dim],
&m[2 * self.dim..],
);
let normed = b.in_ln.forward(&h, 1);
let modulated: Vec<f32> = (0..self.dim)
.map(|i| normed[i] * (1.0 + scale[i]) + shift[i])
.collect();
let mut inner = b.m0.forward(&modulated, 1);
for v in inner.iter_mut() {
*v = silu(*v);
}
let upd = b.m2.forward(&inner, 1);
for i in 0..self.dim {
h[i] += gate[i] * upd[i];
}
}
let m = self.final_layer.ada.forward(&y_act, 1);
let (shift, scale) = (&m[..self.dim], &m[self.dim..]);
let normed = self.final_layer.norm.forward(&h, 1);
let modulated: Vec<f32> = (0..self.dim)
.map(|i| normed[i] * (1.0 + scale[i]) + shift[i])
.collect();
self.final_layer.lin.forward(&modulated, 1)
}
}
pub(crate) struct FlowLm {
pub(crate) embed: Vec<f32>,
pub(crate) input_linear: Linear,
pub(crate) tr: Transformer,
pub(crate) out_norm: LayerNorm,
pub(crate) out_eos: Linear,
pub(crate) flow: FlowNet,
pub(crate) emb_mean: Vec<f32>,
pub(crate) emb_std: Vec<f32>,
pub(crate) bos_emb: Vec<f32>,
pub(crate) speaker_proj: Option<Linear>,
pub(crate) bos_before_voice: Option<Vec<f32>>,
pub(crate) d_model: usize,
pub(crate) latent_dim: usize,
}
impl FlowLm {
fn load(st: &LazySt, cfg: &Config) -> Result<Self> {
Ok(Self {
embed: st.tensor_f32("flow_lm.conditioner.embed.weight")?,
input_linear: Linear::load(
st,
"flow_lm.input_linear",
cfg.d_model,
cfg.latent_dim,
false,
)?,
tr: Transformer::load(
st,
"flow_lm.transformer",
cfg.d_model,
cfg.num_heads,
cfg.num_layers,
cfg.ffn_dim,
None,
false,
cfg.max_period,
)?,
out_norm: LayerNorm::load(st, "flow_lm.out_norm", cfg.d_model, 1e-5)?,
out_eos: Linear::load(st, "flow_lm.out_eos", 1, cfg.d_model, true)?,
flow: FlowNet::load(st, cfg)?,
emb_mean: st.tensor_f32("flow_lm.emb_mean")?,
emb_std: st.tensor_f32("flow_lm.emb_std")?,
bos_emb: st.tensor_f32("flow_lm.bos_emb")?,
speaker_proj: st
.tensor_f32("flow_lm.speaker_proj_weight")
.ok()
.filter(|w| w.len() == cfg.d_model * cfg.latent_dim)
.map(|w| Linear {
packed: PackedWeight::new(&w, cfg.d_model, cfg.latent_dim),
b: None,
out: cfg.d_model,
inp: cfg.latent_dim,
}),
bos_before_voice: st.tensor_f32("flow_lm.bos_before_voice").ok(),
d_model: cfg.d_model,
latent_dim: cfg.latent_dim,
})
}
fn prompt_text(&self, ids: &[u32], state: &mut TrState) {
if ids.is_empty() {
return;
}
let mut x = vec![0f32; ids.len() * self.d_model];
for (t, &id) in ids.iter().enumerate() {
let src = id as usize * self.d_model;
x[t * self.d_model..(t + 1) * self.d_model]
.copy_from_slice(&self.embed[src..src + self.d_model]);
}
self.prompt_rows(x, ids.len(), state);
}
fn prompt_rows(&self, x: Vec<f32>, n: usize, state: &mut TrState) {
if n == 0 {
return;
}
self.tr.forward(x, n, state);
}
fn step(&self, prev: Option<&[f32]>, state: &mut TrState) -> (Vec<f32>, f32) {
let latent = prev.unwrap_or(&self.bos_emb);
let x = self.input_linear.forward(latent, 1);
let out = self.tr.forward(x, 1, state);
let normed = self.out_norm.forward(&out, 1);
let eos = self.out_eos.forward(&normed, 1)[0];
(normed, eos)
}
fn sample_latent(&self, c: &[f32], temp: f32, steps: usize, rng: &mut u64) -> Vec<f32> {
let std = temp.sqrt();
let mut cur: Vec<f32> = (0..self.latent_dim).map(|_| gauss(rng) * std).collect();
for i in 0..steps {
let (s, t) = (i as f32 / steps as f32, (i + 1) as f32 / steps as f32);
let dir = self.flow.forward(c, s, t, &cur);
for (v, d) in cur.iter_mut().zip(dir) {
*v += d / steps as f32;
}
}
cur
}
}
fn gauss(state: &mut u64) -> f32 {
let u1 = ((crate::sampling::xorshift64(state) >> 11) as f64 / (1u64 << 53) as f64).max(1e-12);
let u2 = (crate::sampling::xorshift64(state) >> 11) as f64 / (1u64 << 53) as f64;
((-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()) as f32
}
struct Ct {
c: usize,
t: usize,
d: Vec<f32>,
}
pub(crate) struct Conv1dS {
pub(crate) w: Vec<f32>,
pub(crate) b: Option<Vec<f32>>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) stride: usize,
pub(crate) groups: usize,
pub(crate) replicate: bool,
packed: std::sync::OnceLock<PackedWeight>,
}
#[derive(Clone)]
struct ConvState {
prev: Vec<f32>,
first: bool,
}
impl Conv1dS {
fn load(
st: &LazySt,
prefix: &str,
stride: usize,
groups: usize,
replicate: bool,
) -> Result<Self> {
let shape = st.shape(&format!("{prefix}.weight"))?.to_vec();
anyhow::ensure!(shape.len() == 3, "{prefix}.weight rank {}", shape.len());
Ok(Self {
w: st.tensor_f32(&format!("{prefix}.weight"))?,
b: st.tensor_f32(&format!("{prefix}.bias")).ok(),
in_c: shape[1] * groups,
out_c: shape[0],
k: shape[2],
stride,
groups,
replicate,
packed: std::sync::OnceLock::new(),
})
}
fn state(&self) -> ConvState {
ConvState {
prev: vec![0.0; self.in_c * (self.k - self.stride.min(self.k))],
first: true,
}
}
fn forward(&self, x: &Ct, st: &mut ConvState) -> Ct {
debug_assert_eq!(x.c, self.in_c);
let tp = self.k - self.stride.min(self.k);
if tp > 0 && self.replicate && st.first {
for c in 0..self.in_c {
let v = x.d[c * x.t];
st.prev[c * tp..(c + 1) * tp].fill(v);
}
}
let ta = x.t + tp;
let mut xa = vec![0f32; self.in_c * ta];
for c in 0..self.in_c {
xa[c * ta..c * ta + tp].copy_from_slice(&st.prev[c * tp..(c + 1) * tp]);
xa[c * ta + tp..(c + 1) * ta].copy_from_slice(&x.d[c * x.t..(c + 1) * x.t]);
}
let t_out = (ta - self.k) / self.stride + 1;
let mut out = Ct {
c: self.out_c,
t: t_out,
d: vec![0f32; self.out_c * t_out],
};
if self.groups == 1 {
let packed = self
.packed
.get_or_init(|| PackedWeight::new(&self.w, self.out_c, self.in_c * self.k));
let ck = self.in_c * self.k;
let mut cols = vec![0f32; t_out * ck];
cols.par_chunks_mut(ck).enumerate().for_each(|(ot, crow)| {
let base = ot * self.stride;
for ic in 0..self.in_c {
let row = &xa[ic * ta..];
crow[ic * self.k..(ic + 1) * self.k].copy_from_slice(&row[base..base + self.k]);
}
});
let mut y = vec![0f32; t_out * self.out_c];
gemm_packed(&mut y, &cols, packed, t_out, self.b.as_deref());
out.d
.par_chunks_mut(t_out)
.enumerate()
.for_each(|(oc, orow)| {
for (ot, o) in orow.iter_mut().enumerate() {
*o = y[ot * self.out_c + oc];
}
});
} else {
let icg = self.in_c / self.groups;
let ocg = self.out_c / self.groups;
for oc in 0..self.out_c {
let g = oc / ocg;
let bias = self.b.as_ref().map_or(0.0, |b| b[oc]);
for ot in 0..t_out {
let base = ot * self.stride;
let mut acc = bias;
for ic in 0..icg {
let row = &xa[(g * icg + ic) * ta..];
let w = &self.w[(oc * icg + ic) * self.k..(oc * icg + ic + 1) * self.k];
for (kk, wv) in w.iter().enumerate() {
acc += wv * row[base + kk];
}
}
out.d[oc * t_out + ot] = acc;
}
}
}
if tp > 0 {
for c in 0..self.in_c {
st.prev[c * tp..(c + 1) * tp].copy_from_slice(&xa[(c + 1) * ta - tp..(c + 1) * ta]);
}
st.first = false;
}
out
}
}
pub(crate) struct ConvTr1dS {
pub(crate) w: Vec<f32>,
pub(crate) b: Option<Vec<f32>>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) stride: usize,
pub(crate) groups: usize,
packed_t: std::sync::OnceLock<PackedWeight>,
}
#[derive(Clone)]
struct ConvTrState {
partial: Vec<f32>,
}
impl ConvTr1dS {
fn load(st: &LazySt, prefix: &str, stride: usize, groups: usize) -> Result<Self> {
let shape = st.shape(&format!("{prefix}.weight"))?.to_vec();
anyhow::ensure!(shape.len() == 3, "{prefix}.weight rank {}", shape.len());
Ok(Self {
w: st.tensor_f32(&format!("{prefix}.weight"))?,
b: st.tensor_f32(&format!("{prefix}.bias")).ok(),
in_c: shape[0],
out_c: shape[1] * groups,
k: shape[2],
stride,
groups,
packed_t: std::sync::OnceLock::new(),
})
}
fn state(&self) -> ConvTrState {
ConvTrState {
partial: vec![0.0; self.out_c * (self.k - self.stride)],
}
}
fn forward(&self, x: &Ct, st: &mut ConvTrState) -> Ct {
debug_assert_eq!(x.c, self.in_c);
let (t, s, k) = (x.t, self.stride, self.k);
let tp = k - s;
let full_t = (t - 1) * s + k;
let (icg, ocg) = (self.in_c / self.groups, self.out_c / self.groups);
let mut full = vec![0f32; self.out_c * full_t];
if let Some(b) = &self.b {
for oc in 0..self.out_c {
full[oc * full_t..(oc + 1) * full_t].fill(b[oc]);
}
}
if self.groups == 1 {
let packed = self.packed_t.get_or_init(|| {
let mut wt = vec![0f32; self.out_c * k * self.in_c];
for ic in 0..self.in_c {
for oc in 0..self.out_c {
for kk in 0..k {
wt[(oc * k + kk) * self.in_c + ic] =
self.w[(ic * self.out_c + oc) * k + kk];
}
}
}
PackedWeight::new(&wt, self.out_c * k, self.in_c)
});
let mut a = vec![0f32; t * self.in_c];
for ic in 0..self.in_c {
for ti in 0..t {
a[ti * self.in_c + ic] = x.d[ic * t + ti];
}
}
let mut y = vec![0f32; t * self.out_c * k];
gemm_packed(&mut y, &a, packed, t, None);
full.par_chunks_mut(full_t)
.enumerate()
.for_each(|(oc, orow)| {
for ti in 0..t {
let taps =
&y[ti * self.out_c * k + oc * k..ti * self.out_c * k + (oc + 1) * k];
for (kk, v) in taps.iter().enumerate() {
orow[ti * s + kk] += v;
}
}
});
} else {
for ic in 0..self.in_c {
let g = ic / icg;
for oc_local in 0..ocg {
let oc = g * ocg + oc_local;
let w = &self.w[(ic * ocg + oc_local) * k..(ic * ocg + oc_local + 1) * k];
let orow = &mut full[oc * full_t..(oc + 1) * full_t];
for ti in 0..t {
let v = x.d[ic * t + ti];
for (kk, wv) in w.iter().enumerate() {
orow[ti * s + kk] += wv * v;
}
}
}
}
}
let out_t = full_t - tp;
let mut out = Ct {
c: self.out_c,
t: out_t,
d: vec![0f32; self.out_c * out_t],
};
for oc in 0..self.out_c {
let row = &mut full[oc * full_t..(oc + 1) * full_t];
for i in 0..tp {
row[i] += st.partial[oc * tp + i];
}
let bias = self.b.as_ref().map_or(0.0, |b| b[oc]);
for i in 0..tp {
st.partial[oc * tp + i] = row[out_t + i] - bias;
}
out.d[oc * out_t..(oc + 1) * out_t].copy_from_slice(&row[..out_t]);
}
out
}
}
pub(crate) struct ResnetBlock {
pub(crate) c1: Conv1dS,
pub(crate) c2: Conv1dS,
}
#[derive(Clone)]
struct ResnetState {
c1: ConvState,
c2: ConvState,
}
impl ResnetBlock {
fn forward(&self, x: &Ct, st: &mut ResnetState) -> Ct {
let mut v = Ct {
c: x.c,
t: x.t,
d: x.d.iter().map(|a| elu(*a)).collect(),
};
v = self.c1.forward(&v, &mut st.c1);
for a in v.d.iter_mut() {
*a = elu(*a);
}
v = self.c2.forward(&v, &mut st.c2);
for (o, a) in v.d.iter_mut().zip(&x.d) {
*o += a;
}
v
}
}
pub(crate) struct SeaNetDec {
pub(crate) first: Conv1dS,
pub(crate) stages: Vec<(ConvTr1dS, Vec<ResnetBlock>)>,
pub(crate) last: Conv1dS,
}
#[derive(Clone)]
struct SeaNetDecState {
first: ConvState,
stages: Vec<(ConvTrState, Vec<ResnetState>)>,
last: ConvState,
}
impl SeaNetDec {
fn load(st: &LazySt, ratios: &[usize]) -> Result<Self> {
let first = Conv1dS::load(st, "mimi.decoder.model.0.conv", 1, 1, false)?;
let mut stages = Vec::new();
let mut idx = 2usize; for &ratio in ratios {
let convtr =
ConvTr1dS::load(st, &format!("mimi.decoder.model.{idx}.convtr"), ratio, 1)?;
anyhow::ensure!(
convtr.k == 2 * ratio,
"decoder stage {idx}: kernel {} for ratio {ratio}",
convtr.k
);
idx += 1;
let mut blocks = Vec::new();
while st.has(&format!("mimi.decoder.model.{idx}.block.1.conv.weight")) {
blocks.push(ResnetBlock {
c1: Conv1dS::load(
st,
&format!("mimi.decoder.model.{idx}.block.1.conv"),
1,
1,
false,
)?,
c2: Conv1dS::load(
st,
&format!("mimi.decoder.model.{idx}.block.3.conv"),
1,
1,
false,
)?,
});
idx += 1;
}
anyhow::ensure!(!blocks.is_empty(), "decoder stage has no residual blocks");
stages.push((convtr, blocks));
idx += 1; }
let last = Conv1dS::load(st, &format!("mimi.decoder.model.{idx}.conv"), 1, 1, false)?;
anyhow::ensure!(last.out_c == 1, "decoder emits {} channels", last.out_c);
Ok(Self {
first,
stages,
last,
})
}
fn state(&self) -> SeaNetDecState {
SeaNetDecState {
first: self.first.state(),
stages: self
.stages
.iter()
.map(|(tr, blocks)| {
(
tr.state(),
blocks
.iter()
.map(|b| ResnetState {
c1: b.c1.state(),
c2: b.c2.state(),
})
.collect(),
)
})
.collect(),
last: self.last.state(),
}
}
fn forward(&self, x: &Ct, st: &mut SeaNetDecState) -> Vec<f32> {
let mut z = self.first.forward(x, &mut st.first);
for ((convtr, blocks), (tr_st, block_st)) in self.stages.iter().zip(st.stages.iter_mut()) {
for a in z.d.iter_mut() {
*a = elu(*a);
}
z = convtr.forward(&z, tr_st);
for (b, bs) in blocks.iter().zip(block_st.iter_mut()) {
z = b.forward(&z, bs);
}
}
for a in z.d.iter_mut() {
*a = elu(*a);
}
self.last.forward(&z, &mut st.last).d
}
}
pub(crate) struct MimiEnc {
first: Conv1dS,
stages: Vec<(Vec<ResnetBlock>, Conv1dS)>,
last: Conv1dS,
tr: Transformer,
downsample: Conv1dS,
dim: usize,
}
impl MimiEnc {
fn load(st: &LazySt, cfg: &Config) -> Result<Self> {
let first = Conv1dS::load(st, "mimi.encoder.model.0.conv", 1, 1, false)?;
let mut stages = Vec::new();
let mut idx = 1usize;
for &ratio in cfg.ratios.iter().rev() {
let mut blocks = Vec::new();
while st.has(&format!("mimi.encoder.model.{idx}.block.1.conv.weight")) {
blocks.push(ResnetBlock {
c1: Conv1dS::load(
st,
&format!("mimi.encoder.model.{idx}.block.1.conv"),
1,
1,
false,
)?,
c2: Conv1dS::load(
st,
&format!("mimi.encoder.model.{idx}.block.3.conv"),
1,
1,
false,
)?,
});
idx += 1;
}
anyhow::ensure!(
!blocks.is_empty(),
"encoder stage {idx} has no residual blocks"
);
idx += 1; let conv = Conv1dS::load(
st,
&format!("mimi.encoder.model.{idx}.conv"),
ratio,
1,
false,
)?;
anyhow::ensure!(
conv.k == 2 * ratio,
"encoder stage {idx}: kernel {} for ratio {ratio}",
conv.k
);
stages.push((blocks, conv));
idx += 1;
}
idx += 1; let last = Conv1dS::load(st, &format!("mimi.encoder.model.{idx}.conv"), 1, 1, false)?;
Ok(Self {
first,
stages,
last,
tr: Transformer::load(
st,
"mimi.encoder_transformer.transformer",
cfg.mimi_dim,
cfg.mimi_heads,
cfg.mimi_layers,
cfg.mimi_ffn,
Some(cfg.mimi_context),
true,
cfg.max_period,
)?,
downsample: Conv1dS::load(
st,
"mimi.downsample.conv.conv",
cfg.resample_stride,
1,
true,
)?,
dim: cfg.mimi_dim,
})
}
fn encode_to_latent(&self, pcm: &[f32], frame_size: usize) -> Ct {
let frames = pcm.len().div_ceil(frame_size).max(1);
let mut d = pcm.to_vec();
d.resize(frames * frame_size, 0.0);
let mut y = Ct {
c: 1,
t: d.len(),
d,
};
let mut st = self.first.state();
y = self.first.forward(&y, &mut st);
for (blocks, conv) in &self.stages {
for b in blocks {
let mut bs = ResnetState {
c1: b.c1.state(),
c2: b.c2.state(),
};
y = b.forward(&y, &mut bs);
}
for a in y.d.iter_mut() {
*a = elu(*a);
}
let mut cs = conv.state();
y = conv.forward(&y, &mut cs);
}
for a in y.d.iter_mut() {
*a = elu(*a);
}
let mut ls = self.last.state();
y = self.last.forward(&y, &mut ls);
let mut tm = vec![0f32; y.t * self.dim];
for c in 0..self.dim {
for t in 0..y.t {
tm[t * self.dim + c] = y.d[c * y.t + t];
}
}
let mut trs = self.tr.state();
let tm = self.tr.forward(tm, y.t, &mut trs);
let mut back = Ct {
c: self.dim,
t: y.t,
d: vec![0f32; y.t * self.dim],
};
for c in 0..self.dim {
for t in 0..y.t {
back.d[c * y.t + t] = tm[t * self.dim + c];
}
}
let mut ds = self.downsample.state();
self.downsample.forward(&back, &mut ds)
}
}
pub(crate) struct MimiDec {
pub(crate) quant_out: Linear,
pub(crate) upsample: ConvTr1dS,
pub(crate) tr: Transformer,
pub(crate) dec: SeaNetDec,
pub(crate) dim: usize,
}
pub struct MimiState {
up: ConvTrState,
tr: TrState,
dec: SeaNetDecState,
}
impl MimiDec {
fn load(st: &LazySt, cfg: &Config) -> Result<Self> {
let quant_w = st.tensor_f32("mimi.quantizer.output_proj.weight")?;
anyhow::ensure!(
quant_w.len() == cfg.mimi_dim * cfg.latent_dim,
"quantizer.output_proj has {} elements",
quant_w.len()
);
Ok(Self {
quant_out: Linear {
packed: PackedWeight::new(&quant_w, cfg.mimi_dim, cfg.latent_dim),
b: None,
out: cfg.mimi_dim,
inp: cfg.latent_dim,
},
upsample: ConvTr1dS::load(
st,
"mimi.upsample.convtr.convtr",
cfg.resample_stride,
cfg.mimi_dim,
)?,
tr: Transformer::load(
st,
"mimi.decoder_transformer.transformer",
cfg.mimi_dim,
cfg.mimi_heads,
cfg.mimi_layers,
cfg.mimi_ffn,
Some(cfg.mimi_context),
true,
cfg.max_period,
)?,
dec: SeaNetDec::load(st, &cfg.ratios)?,
dim: cfg.mimi_dim,
})
}
fn state(&self) -> MimiState {
MimiState {
up: self.upsample.state(),
tr: self.tr.state(),
dec: self.dec.state(),
}
}
fn decode_latents(&self, latents: &[f32], n: usize, st: &mut MimiState) -> Vec<f32> {
let z = self.quant_out.forward(latents, n);
let mut zc = vec![0f32; n * self.dim];
for t in 0..n {
for c in 0..self.dim {
zc[c * n + t] = z[t * self.dim + c];
}
}
let up = self.upsample.forward(
&Ct {
c: self.dim,
t: n,
d: zc,
},
&mut st.up,
);
let mut tm = vec![0f32; up.t * self.dim];
for c in 0..self.dim {
for t in 0..up.t {
tm[t * self.dim + c] = up.d[c * up.t + t];
}
}
let tm = self.tr.forward(tm, up.t, &mut st.tr);
let mut back = Ct {
c: self.dim,
t: up.t,
d: vec![0f32; up.t * self.dim],
};
for c in 0..self.dim {
for t in 0..up.t {
back.d[c * up.t + t] = tm[t * self.dim + c];
}
}
self.dec.forward(&back, &mut st.dec)
}
}
#[derive(Clone)]
pub struct Voice {
pub(crate) state: TrState,
}
impl Voice {
pub fn kv_rows(&self, layer: usize) -> Option<(&[f32], &[f32], usize)> {
self.state.layers().get(layer).map(|kv| kv.rows())
}
pub fn layer_count(&self) -> usize {
self.state.layers().len()
}
pub fn len(&self) -> usize {
self.state.offset()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[derive(Debug, Clone)]
pub struct GenOpts {
pub temperature: f32,
pub lsd_steps: usize,
pub eos_threshold: f32,
pub frames_after_eos: Option<usize>,
pub max_tokens_per_chunk: usize,
pub seed: u64,
}
impl Default for GenOpts {
fn default() -> Self {
Self {
temperature: 0.3,
lsd_steps: 1,
eos_threshold: -4.0,
frames_after_eos: None,
max_tokens_per_chunk: 50,
seed: 0x5eed_c0ff_ee01,
}
}
}
pub struct PocketTts {
cfg: Config,
tok: SpUnigram,
lm: FlowLm,
mimi: MimiDec,
enc: Option<MimiEnc>,
voice_cloning: bool,
}
pub struct Session {
state: TrState,
mimi: MimiState,
rng: u64,
eos_step: Option<usize>,
step: usize,
}
impl Session {
pub fn eos_step(&self) -> Option<usize> {
self.eos_step
}
pub fn offset(&self) -> usize {
self.state.offset()
}
}
pub trait FrameSource {
fn start(&mut self, tts: &PocketTts, voice: &Voice, ids: &[u32]) -> Result<()>;
fn next(
&mut self,
tts: &PocketTts,
prev: Option<&[f32]>,
noise: &[f32],
opts: &GenOpts,
) -> Result<Step>;
}
#[derive(Default)]
pub struct CpuFrames {
session: Option<Session>,
}
impl FrameSource for CpuFrames {
fn start(&mut self, tts: &PocketTts, voice: &Voice, ids: &[u32]) -> Result<()> {
let mut s = tts.session(voice, 0);
tts.prompt_tokens(&mut s, ids);
self.session = Some(s);
Ok(())
}
fn next(
&mut self,
tts: &PocketTts,
prev: Option<&[f32]>,
noise: &[f32],
opts: &GenOpts,
) -> Result<Step> {
let session = self.session.as_mut().context("start() was not called")?;
let (cond, eos_logit) = tts.lm.step(prev, &mut session.state);
let latent = tts.sample_latent_with_noise(&cond, noise, opts.lsd_steps);
let pcm = tts
.mimi
.decode_latents(&tts.denormalize(&latent), 1, &mut session.mimi);
Ok(Step {
latent,
eos_logit,
pcm,
})
}
}
pub struct Step {
pub latent: Vec<f32>,
pub eos_logit: f32,
pub pcm: Vec<f32>,
}
impl PocketTts {
pub fn load(dir: impl AsRef<Path>) -> Result<Self> {
let dir = dir.as_ref();
let st = LazySt::open(dir)?;
let tok = SpUnigram::load(&dir.join("tokenizer.model"))?;
Self::from_st(st, tok)
}
pub fn from_bytes(model: Vec<u8>, tokenizer: &[u8]) -> Result<Self> {
Self::from_st(
LazySt::from_bytes(vec![model])?,
SpUnigram::from_bytes(tokenizer)?,
)
}
fn from_st(st: LazySt, tok: SpUnigram) -> Result<Self> {
let cfg = Config::from_manifest(&st)?;
anyhow::ensure!(
tok.vocab_size() == cfg.n_bins,
"tokenizer has {} pieces but the LUT has {} bins",
tok.vocab_size(),
cfg.n_bins
);
let lm = FlowLm::load(&st, &cfg)?;
let mimi = MimiDec::load(&st, &cfg)?;
let voice_cloning = st
.tensor_f32("mimi.encoder.model.0.conv.weight")
.map(|w| w.iter().any(|v| *v != 0.0))
.unwrap_or(false);
let enc = if voice_cloning {
Some(MimiEnc::load(&st, &cfg)?)
} else {
None
};
Ok(Self {
cfg,
tok,
lm,
mimi,
enc,
voice_cloning,
})
}
pub fn config(&self) -> &Config {
&self.cfg
}
pub fn supports_voice_cloning(&self) -> bool {
self.voice_cloning
}
pub fn tokenizer(&self) -> &SpUnigram {
&self.tok
}
pub fn sample_rate(&self) -> usize {
self.cfg.sample_rate
}
pub fn encode_audio(&self, pcm: &[f32]) -> Result<Vec<f32>> {
let enc = self
.enc
.as_ref()
.filter(|_| self.voice_cloning)
.context("this checkpoint has no Mimi encoder (the ungated weights zero it)")?;
let lat = enc.encode_to_latent(pcm, self.cfg.frame_size());
let mut rows = vec![0f32; lat.t * lat.c];
for c in 0..lat.c {
for t in 0..lat.t {
rows[t * lat.c + c] = lat.d[c * lat.t + t];
}
}
Ok(rows)
}
pub fn clone_voice(&self, pcm: &[f32]) -> Result<Voice> {
let enc = self.enc.as_ref().filter(|_| self.voice_cloning).context(
"this checkpoint cannot clone voices — its Mimi encoder is zeroed (the \
`kyutai/pocket-tts-without-voice-cloning` weights). Accept the terms at \
https://huggingface.co/kyutai/pocket-tts and use those instead.",
)?;
let proj = self
.lm
.speaker_proj
.as_ref()
.context("checkpoint has no speaker projection")?;
anyhow::ensure!(!pcm.is_empty(), "empty audio");
let lat = enc.encode_to_latent(pcm, self.cfg.frame_size());
let mut rows = vec![0f32; lat.t * lat.c];
for c in 0..lat.c {
for t in 0..lat.t {
rows[t * lat.c + c] = lat.d[c * lat.t + t];
}
}
let cond = proj.forward(&rows, lat.t);
let d = self.lm.d_model;
let mut prompt = Vec::with_capacity((lat.t + 1) * d);
if let Some(bos) = &self.lm.bos_before_voice {
prompt.extend_from_slice(bos);
}
prompt.extend_from_slice(&cond);
let n = prompt.len() / d;
let mut state = self.lm.tr.state();
self.lm.prompt_rows(prompt, n, &mut state);
Ok(Voice { state })
}
pub fn load_voice(&self, path: impl AsRef<Path>) -> Result<Voice> {
let path = path.as_ref();
let bytes = std::fs::read(path).with_context(|| format!("read {}", path.display()))?;
self.load_voice_bytes(bytes)
}
pub fn load_voice_bytes(&self, bytes: Vec<u8>) -> Result<Voice> {
let offsets = safetensors_i64_scalars(&bytes)?;
let st = LazySt::from_bytes(vec![bytes])?;
let mut state = self.lm.tr.state();
for (i, kv) in state.kv.iter_mut().enumerate() {
let name = format!("transformer.layers.{i}.self_attn/cache");
let shape = st.shape(&name)?.to_vec();
anyhow::ensure!(
shape.len() == 5 && shape[0] == 2 && shape[1] == 1,
"voice cache {name} has shape {shape:?}, expected [2, 1, T, H, D]"
);
let (positions, heads, hd) = (shape[2], shape[3], shape[4]);
anyhow::ensure!(
heads == kv.heads && hd == kv.hd,
"voice cache is {heads}×{hd} per position, model wants {}×{}",
kv.heads,
kv.hd
);
let data = st.tensor_f32(&name)?;
let half = positions * heads * hd;
kv.k = data[..half].to_vec();
kv.v = data[half..2 * half].to_vec();
kv.offset = offsets
.get(&format!("transformer.layers.{i}.self_attn/offset"))
.map(|v| *v as usize)
.unwrap_or(positions);
anyhow::ensure!(
kv.offset <= positions,
"voice offset {} exceeds the {positions} cached positions",
kv.offset
);
}
Ok(Voice { state })
}
pub fn session(&self, voice: &Voice, seed: u64) -> Session {
Session {
state: voice.state.clone(),
mimi: self.mimi.state(),
rng: seed | 1,
eos_step: None,
step: 0,
}
}
pub fn prompt_tokens(&self, session: &mut Session, ids: &[u32]) {
self.lm.prompt_text(ids, &mut session.state);
}
pub fn sample_latent_with_noise(&self, cond: &[f32], noise: &[f32], steps: usize) -> Vec<f32> {
let mut cur = noise.to_vec();
for i in 0..steps {
let (s, t) = (i as f32 / steps as f32, (i + 1) as f32 / steps as f32);
let dir = self.lm.flow.forward(cond, s, t, &cur);
for (v, d) in cur.iter_mut().zip(dir) {
*v += d / steps as f32;
}
}
cur
}
pub fn draw_noise(&self, opts: &GenOpts, rng: &mut u64) -> Vec<f32> {
let std = opts.temperature.sqrt();
(0..self.lm.latent_dim).map(|_| gauss(rng) * std).collect()
}
pub fn sample_latent_from(&self, cond: &[f32], opts: &GenOpts, rng: &mut u64) -> Vec<f32> {
self.lm
.sample_latent(cond, opts.temperature, opts.lsd_steps, rng)
}
pub fn denormalize(&self, latent: &[f32]) -> Vec<f32> {
latent
.iter()
.enumerate()
.map(|(i, v)| v * self.lm.emb_std[i] + self.lm.emb_mean[i])
.collect()
}
pub fn backbone_step(&self, session: &mut Session, prev: Option<&[f32]>) -> (Vec<f32>, f32) {
self.lm.step(prev, &mut session.state)
}
pub fn text_embedding(&self, id: u32) -> &[f32] {
let d = self.lm.d_model;
&self.lm.embed[id as usize * d..(id as usize + 1) * d]
}
pub fn latent_input(&self, prev: Option<&[f32]>) -> Vec<f32> {
self.lm
.input_linear
.forward(prev.unwrap_or(&self.lm.bos_emb), 1)
}
pub fn next_frame(&self, session: &mut Session, prev: Option<&[f32]>, opts: &GenOpts) -> Step {
let (cond, eos_logit) = self.lm.step(prev, &mut session.state);
if eos_logit > opts.eos_threshold && session.eos_step.is_none() {
session.eos_step = Some(session.step);
}
let latent =
self.lm
.sample_latent(&cond, opts.temperature, opts.lsd_steps, &mut session.rng);
let denorm: Vec<f32> = latent
.iter()
.enumerate()
.map(|(i, v)| v * self.lm.emb_std[i] + self.lm.emb_mean[i])
.collect();
let pcm = self.mimi.decode_latents(&denorm, 1, &mut session.mimi);
session.step += 1;
Step {
latent,
eos_logit,
pcm,
}
}
pub fn decode_latents(&self, latents: &[f32], n: usize, state: &mut MimiState) -> Vec<f32> {
self.mimi.decode_latents(latents, n, state)
}
pub fn mimi_state(&self) -> MimiState {
self.mimi.state()
}
pub(crate) fn flow_lm(&self) -> &FlowLm {
&self.lm
}
pub(crate) fn mimi_dec(&self) -> &MimiDec {
&self.mimi
}
pub fn latent_stats(&self) -> (&[f32], &[f32]) {
(&self.lm.emb_mean, &self.lm.emb_std)
}
pub fn generate(&self, voice: &Voice, text: &str, opts: &GenOpts) -> Result<Vec<f32>> {
let mut pcm = Vec::new();
self.generate_streaming(voice, text, opts, |frame| pcm.extend_from_slice(frame))?;
Ok(pcm)
}
pub fn generate_streaming(
&self,
voice: &Voice,
text: &str,
opts: &GenOpts,
on_frame: impl FnMut(&[f32]),
) -> Result<()> {
self.generate_with(&mut CpuFrames::default(), voice, text, opts, on_frame)
}
pub fn generate_with(
&self,
src: &mut impl FrameSource,
voice: &Voice,
text: &str,
opts: &GenOpts,
mut on_frame: impl FnMut(&[f32]),
) -> Result<()> {
let mut rng = opts.seed | 1;
for chunk in self.split_into_best_sentences(text, opts.max_tokens_per_chunk)? {
let (_, guess) = prepare_text_prompt(&chunk)?;
let frames_after_eos = opts.frames_after_eos.unwrap_or(guess + 2);
let ids = self.tok.encode(&chunk);
let max_gen = ((ids.len() as f32 / 3.0 + 2.0) * self.cfg.frame_rate).ceil() as usize;
src.start(self, voice, &ids)?;
let mut prev: Option<Vec<f32>> = None;
let mut eos_step: Option<usize> = None;
for step in 0..max_gen {
let noise = self.draw_noise(opts, &mut rng);
let frame = src.next(self, prev.as_deref(), &noise, opts)?;
if frame.eos_logit > opts.eos_threshold && eos_step.is_none() {
eos_step = Some(step);
}
if eos_step.is_some_and(|e| step >= e + frames_after_eos) {
break;
}
on_frame(&frame.pcm);
prev = Some(frame.latent);
}
}
Ok(())
}
pub fn split_into_best_sentences(&self, text: &str, max_tokens: usize) -> Result<Vec<String>> {
let (text, _) = prepare_text_prompt(text)?;
let tokens = self.tok.encode(text.trim());
let enders = &self.tok.encode(".!...?")[1..];
let fallback = &self.tok.encode(",;:")[1..];
let segments = self.segments(&tokens, enders);
let mut refined: Vec<(usize, String)> = Vec::new();
for (n, seg) in segments {
if n <= max_tokens {
refined.push((n, seg));
continue;
}
let sub_tokens = self.tok.encode(seg.trim());
let sub = self.segments(&sub_tokens, fallback);
if sub.len() > 1 {
refined.extend(sub);
} else {
refined.push((n, seg));
}
}
let mut chunks = Vec::new();
let mut current = String::new();
let mut count = 0usize;
for (n, seg) in refined {
if current.is_empty() {
current = seg;
count = n;
} else if count + n > max_tokens {
chunks.push(current.trim().to_string());
current = seg;
count = n;
} else {
current.push(' ');
current.push_str(&seg);
count += n;
}
}
if !current.is_empty() {
chunks.push(current.trim().to_string());
}
Ok(chunks)
}
fn segments(&self, tokens: &[u32], boundary: &[u32]) -> Vec<(usize, String)> {
let mut indices = vec![0usize];
let mut previous_was_boundary = false;
for (i, t) in tokens.iter().enumerate() {
if boundary.contains(t) {
previous_was_boundary = true;
} else {
if previous_was_boundary {
indices.push(i);
}
previous_was_boundary = false;
}
}
indices.push(tokens.len());
indices
.windows(2)
.map(|w| (w[1] - w[0], self.tok.decode(&tokens[w[0]..w[1]])))
.collect()
}
}
fn safetensors_i64_scalars(bytes: &[u8]) -> Result<HashMap<String, i64>> {
anyhow::ensure!(bytes.len() > 8, "safetensors blob too small");
let hlen = u64::from_le_bytes(bytes[..8].try_into().expect("8 bytes")) as usize;
let start = 8 + hlen;
anyhow::ensure!(start <= bytes.len(), "safetensors header exceeds the blob");
let hdr: serde_json::Value = serde_json::from_slice(&bytes[8..start])?;
let mut out = HashMap::new();
for (name, meta) in hdr
.as_object()
.context("safetensors header is not an object")?
{
if meta.get("dtype").and_then(|d| d.as_str()) != Some("I64") {
continue;
}
let Some(off) = meta
.get("data_offsets")
.and_then(|o| o.get(0))
.and_then(serde_json::Value::as_u64)
else {
continue;
};
let at = start + off as usize;
if at + 8 <= bytes.len() {
out.insert(
name.clone(),
i64::from_le_bytes(bytes[at..at + 8].try_into().expect("8 bytes")),
);
}
}
Ok(out)
}
pub fn prepare_text_prompt(text: &str) -> Result<(String, usize)> {
let text = text.trim();
anyhow::ensure!(!text.is_empty(), "text prompt cannot be empty");
let mut text = text.replace(['\n', '\r'], " ").replace(" ", " ");
let guess = if text.split_whitespace().count() <= 4 {
3
} else {
1
};
let first = text.chars().next().expect("non-empty");
if !first.is_uppercase() {
let rest: String = text.chars().skip(1).collect();
text = first.to_uppercase().collect::<String>() + &rest;
}
if text.chars().last().is_some_and(|c| c.is_alphanumeric()) {
text.push('.');
}
Ok((text, guess))
}