use std::path::Path;
use anyhow::{Context, Result};
use crate::cpu_gemm::{PackedWeight, PackedWeightI8, gemm_i8, gemm_packed};
use crate::weights::LazySt;
pub const N_STREAMS: usize = 17;
pub const DEP_Q: usize = 8;
pub const SILENCE_TOKENS: [u32; 8] = [948, 243, 1178, 546, 1736, 1030, 1978, 2008];
pub const SINE_TOKENS: [u32; 8] = [430, 1268, 381, 1611, 1095, 1495, 56, 472];
pub const PRIME_SILENCE_FRAMES: usize = 6;
pub const CARD: u32 = 2048; pub const TEXT_CARD: u32 = 32000; const UNGENERATED: i64 = -2;
const DELAYS: [usize; N_STREAMS] = [0, 0, 1, 1, 1, 1, 1, 1, 1, 0, 1, 1, 1, 1, 1, 1, 1];
const MAX_DELAY: usize = 1;
const CT: usize = MAX_DELAY + 2;
const DIM: usize = 4096;
const LAYERS: usize = 32;
const HEADS: usize = 32;
const HEAD_DIM: usize = DIM / HEADS;
const CONTEXT: usize = 3000;
const HIDDEN: usize = 11264; const RMS_EPS: f32 = 1e-8;
const DDIM: usize = 1024;
const DLAYERS: usize = 6;
const DHEADS: usize = 16;
const DHEAD_DIM: usize = DDIM / DHEADS;
const DHIDDEN: usize = 2816;
enum Band {
F32(PackedWeight),
I8(PackedWeightI8),
}
impl Band {
fn mv(&self, x: &[f32], out: &mut [f32]) {
match self {
Band::F32(w) => gemm_packed(out, x, w, 1, None),
Band::I8(w) => gemm_i8(out, x, w, 1, None),
}
}
}
fn rms_norm(x: &[f32], alpha: &[f32], out: &mut [f32]) {
let var = RMS_EPS + x.iter().map(|v| v * v).sum::<f32>() / x.len() as f32;
let inv = 1.0 / var.sqrt();
for ((o, v), a) in out.iter_mut().zip(x).zip(alpha) {
*o = v * a * inv;
}
}
fn q4_roundtrip(v: &mut [f32], rows: usize, cols: usize) {
let (sc, qz) = crate::weights::quantize_q4_0(v, rows, cols);
let nblk = cols / 32;
for r in 0..rows {
for b in 0..nblk {
let d = sc[r * nblk + b];
for wi in 0..4 {
let word = qz[(r * nblk + b) * 4 + wi];
for byte in 0..4 {
let bt = ((word >> (8 * byte)) & 0xff) as u8;
let i = wi * 4 + byte;
v[r * cols + b * 32 + i] = ((bt & 0x0f) as i32 - 8) as f32 * d;
v[r * cols + b * 32 + i + 16] = ((bt >> 4) as i32 - 8) as f32 * d;
}
}
}
}
}
fn q4_affine_roundtrip(v: &mut [f32], rows: usize, cols: usize) {
for r in 0..rows {
for b in 0..cols / 32 {
let blk = &mut v[r * cols + b * 32..r * cols + b * 32 + 32];
let (mut lo, mut hi) = (f32::INFINITY, f32::NEG_INFINITY);
for &x in blk.iter() {
lo = lo.min(x);
hi = hi.max(x);
}
let d = (hi - lo) / 15.0;
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
for x in blk.iter_mut() {
let q = (((*x - lo) * id).round() as i32).clamp(0, 15);
*x = lo + q as f32 * d;
}
}
}
}
fn qn_roundtrip(v: &mut [f32], rows: usize, cols: usize, bits: u32, group: usize) {
let qmin = -(1i32 << (bits - 1));
let qmax = (1i32 << (bits - 1)) - 1;
for r in 0..rows {
for b in 0..cols / group {
let blk = &mut v[r * cols + b * group..r * cols + (b + 1) * group];
let (mut amax, mut mx) = (0f32, 0f32);
for &x in blk.iter() {
if x.abs() > amax {
amax = x.abs();
mx = x;
}
}
let d = mx / qmin as f32;
let id = if d != 0.0 { 1.0 / d } else { 0.0 };
for x in blk.iter_mut() {
*x = ((*x * id).round() as i32).clamp(qmin, qmax) as f32 * d;
}
}
}
}
fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
fn argmax(v: &[f32]) -> u32 {
let mut best = (f32::NEG_INFINITY, 0u32);
for (i, &x) in v.iter().enumerate() {
if x > best.0 {
best = (x, i as u32);
}
}
best.1
}
fn xorshift_u01(st: &mut u64) -> f64 {
let mut x = *st;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
*st = x;
(x >> 11) as f64 / (1u64 << 53) as f64
}
pub fn topk_sample(logits: &[f32], temp: f32, k: usize, st: &mut u64) -> u32 {
let mut idx: Vec<(f32, u32)> = logits.iter().enumerate().map(|(i, &l)| (l, i as u32)).collect();
let k = k.min(idx.len());
idx.select_nth_unstable_by(k - 1, |a, b| b.0.total_cmp(&a.0));
idx.truncate(k);
let mx = idx.iter().map(|v| v.0).fold(f32::NEG_INFINITY, f32::max);
let mut probs: Vec<f64> = idx.iter().map(|v| (((v.0 - mx) / temp) as f64).exp()).collect();
let sum: f64 = probs.iter().sum();
for p in probs.iter_mut() {
*p /= sum;
}
let mut u = xorshift_u01(st);
for (p, v) in probs.iter().zip(&idx) {
u -= p;
if u <= 0.0 {
return v.1;
}
}
idx.last().unwrap().1
}
pub fn load_voice_embeddings(path: &Path) -> Result<Vec<Vec<f32>>> {
let bytes = std::fs::read(path).with_context(|| format!("voice prompt {}", path.display()))?;
let hlen = u64::from_le_bytes(bytes[..8].try_into().unwrap()) as usize;
let hdr: serde_json::Value = serde_json::from_slice(&bytes[8..8 + hlen])?;
let meta = hdr.get("embeddings").context("no `embeddings` tensor")?;
anyhow::ensure!(meta["dtype"] == "F32", "voice embeddings must be F32");
let shape: Vec<usize> = meta["shape"]
.as_array()
.context("shape")?
.iter()
.map(|v| v.as_u64().unwrap() as usize)
.collect();
anyhow::ensure!(shape.len() == 2 && shape[1] == 4096, "want [T,4096], got {shape:?}");
let o = meta["data_offsets"].as_array().context("offsets")?;
let a = o[0].as_u64().unwrap() as usize + 8 + hlen;
Ok(bytes[a..a + shape[0] * 4096 * 4]
.chunks_exact(4096 * 4)
.map(|r| r.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().unwrap())).collect())
.collect())
}
pub fn load_persona_tokens(spec: &str) -> Result<Vec<u32>> {
let (path, key) = spec.rsplit_once(':').context("persona spec wants file.json:key")?;
let v: serde_json::Value = serde_json::from_slice(&std::fs::read(path)?)?;
let arr = v.get(key).with_context(|| format!("persona `{key}` not in {path}"))?;
Ok(arr
.as_array()
.context("token array")?
.iter()
.map(|t| t.as_u64().unwrap() as u32)
.collect())
}
pub fn gumbel_full_pick(logits: &[f32], temp: f32, st: &mut u64) -> u32 {
let mut best = (f32::NEG_INFINITY, 0u32);
for (i, &l) in logits.iter().enumerate() {
let u = xorshift_u01(st).max(1e-300);
let g = -(-u.ln()).ln() as f32;
let s = l / temp + g;
if s > best.0 {
best = (s, i as u32);
}
}
best.1
}
fn rope_row(row: &mut [f32], pos: usize, head_dim: usize) {
let ts = pos as f32;
for head in row.chunks_exact_mut(head_dim) {
for i in 0..head_dim / 2 {
let freq = (-(10000f32).ln() * 2.0 * i as f32 / head_dim as f32).exp();
let (sin, cos) = (freq * ts).sin_cos();
let (r, im) = (head[2 * i], head[2 * i + 1]);
head[2 * i] = r * cos - im * sin;
head[2 * i + 1] = r * sin + im * cos;
}
}
}
fn attend(q: &[f32], ks: &[Vec<f32>], vs: &[Vec<f32>], heads: usize, head_dim: usize, out: &mut [f32]) {
let scale = 1.0 / (head_dim as f32).sqrt();
out.fill(0.0);
for h in 0..heads {
let qh = &q[h * head_dim..(h + 1) * head_dim];
let mut scores: Vec<f32> = ks
.iter()
.map(|k| qh.iter().zip(&k[h * head_dim..(h + 1) * head_dim]).map(|(a, b)| a * b).sum::<f32>() * scale)
.collect();
let mx = scores.iter().fold(f32::NEG_INFINITY, |a, &b| a.max(b));
let mut den = 0f32;
for s in scores.iter_mut() {
*s = (*s - mx).exp();
den += *s;
}
let oh = &mut out[h * head_dim..(h + 1) * head_dim];
for (s, v) in scores.iter().zip(vs) {
let w = s / den;
for (o, vv) in oh.iter_mut().zip(&v[h * head_dim..(h + 1) * head_dim]) {
*o += w * vv;
}
}
}
}
struct TemporalLayer {
norm1: Vec<f32>,
in_proj: Band, out_proj: Band, norm2: Vec<f32>,
gate_in: Band, gate_out: Band, }
struct DepLayer {
norm1: Vec<f32>, in_projs: Vec<Band>, out_projs: Vec<Band>, norm2: Vec<f32>,
gate_in: Vec<Band>, gate_out: Vec<Band>, }
pub fn spm_pieces(model_path: &Path) -> Result<Vec<String>> {
fn varint(b: &[u8], mut o: usize) -> Result<(u64, usize)> {
let (mut v, mut sh, start) = (0u64, 0u32, o);
loop {
let byte = *b.get(o).context("varint eof")?;
v |= ((byte & 0x7f) as u64) << sh;
o += 1;
if byte & 0x80 == 0 {
return Ok((v, o - start));
}
sh += 7;
}
}
let b = std::fs::read(model_path)?;
let mut pieces = Vec::with_capacity(32001);
let mut o = 0usize;
while o < b.len() {
let (tag, n) = varint(&b, o)?;
o += n;
let (field, wire) = ((tag >> 3) as u32, (tag & 7) as u8);
match wire {
2 => {
let (len, n2) = varint(&b, o)?;
o += n2;
let body = &b[o..o + len as usize];
o += len as usize;
if field == 1 {
let (t2, m) = varint(body, 0)?;
if t2 >> 3 == 1 && t2 & 7 == 2 {
let (l2, m2) = varint(body, m)?;
pieces.push(
String::from_utf8_lossy(&body[m + m2..m + m2 + l2 as usize])
.into_owned(),
);
} else {
pieces.push(String::new());
}
}
}
0 => {
let (_, n2) = varint(&b, o)?;
o += n2;
}
5 => o += 4,
1 => o += 8,
w => anyhow::bail!("unexpected wire type {w}"),
}
}
anyhow::ensure!(pieces.len() >= 32000, "spm parse found only {} pieces", pieces.len());
Ok(pieces)
}
pub fn find_spm_model() -> Result<std::path::PathBuf> {
fn walk(dir: &Path, name: &str) -> Option<std::path::PathBuf> {
for e in std::fs::read_dir(dir).ok()? {
let p = e.ok()?.path();
if p.is_dir() {
if let Some(f) = walk(&p, name) {
return Some(f);
}
} else if p.file_name().is_some_and(|f| f == name) {
return Some(p);
}
}
None
}
let hub = std::path::PathBuf::from(std::env::var("HOME")?)
.join(".cache/huggingface/hub/models--kyutai--moshiko-pytorch-bf16");
walk(&hub, "tokenizer_spm_32k_3.model")
.context("tokenizer_spm_32k_3.model not in the HF cache (run export_moshi_lm.py once)")
}
#[derive(Clone, Copy)]
pub struct Sampling {
pub temp: f32,
pub temp_text: f32,
pub top_k: usize,
pub top_k_text: usize,
pub gumbel_full: bool,
}
impl Default for Sampling {
fn default() -> Self {
Self { temp: 0.8, temp_text: 0.7, top_k: 250, top_k_text: 25, gumbel_full: false }
}
}
pub struct MoshiLm {
emb: Vec<Vec<f32>>, text_emb: Vec<f32>, layers: Vec<TemporalLayer>,
out_norm: Vec<f32>,
text_linear: Band, dep_in: Vec<Band>, dep_text_emb: Vec<f32>, dep_emb: Vec<Vec<f32>>, dep_layers: Vec<DepLayer>,
linears: Vec<Band>, sampler: Option<(Sampling, std::cell::RefCell<u64>)>,
}
pub struct MoshiState {
cache: [[i64; CT]; N_STREAMS],
offset: usize,
k_cache: Vec<Vec<Vec<f32>>>, v_cache: Vec<Vec<Vec<f32>>>,
}
impl MoshiState {
pub fn snapshot(&self) -> Self {
Self {
cache: self.cache,
offset: self.offset,
k_cache: self.k_cache.clone(),
v_cache: self.v_cache.clone(),
}
}
pub fn clock(&self) -> usize {
self.offset
}
pub fn advance_clock(&mut self) {
self.offset += 1;
}
}
impl MoshiState {
pub fn force_generated(&mut self, text: u32, audio: [u32; DEP_Q]) {
let pos = self.offset % CT;
self.cache[0][pos] = text as i64;
for (cb, &tok) in audio.iter().enumerate() {
self.cache[1 + cb][pos] = tok as i64;
}
}
}
pub struct StepTrace {
pub input_tokens: [i64; N_STREAMS],
pub emb_sum: Vec<f32>,
pub transformer_out: Vec<f32>,
pub text_logits: Vec<f32>,
pub audio_logits: Vec<Vec<f32>>, pub text_token: u32,
pub audio_tokens: [u32; DEP_Q],
pub out: Option<[i32; DEP_Q + 1]>,
}
impl MoshiLm {
pub fn load(dir: &Path) -> Result<Self> {
Self::load_band(dir, false, false)
}
pub fn load_i8(dir: &Path) -> Result<Self> {
Self::load_band(dir, true, false)
}
pub fn load_q4sim(dir: &Path) -> Result<Self> {
Self::load_band(dir, false, true)
}
fn load_band(dir: &Path, i8: bool, q4sim: bool) -> Result<Self> {
let st = LazySt::open(dir).context("moshi-lm checkpoint")?;
let t = |n: &str| st.tensor_f32(n);
let parse = |k: &str| -> Vec<String> {
std::env::var(k)
.map(|v| v.split(',').map(|s| s.trim().to_string()).filter(|s| !s.is_empty()).collect())
.unwrap_or_default()
};
let (only, spare) = (parse("MOSHI_Q4SIM_ONLY"), parse("MOSHI_Q4SIM_SPARE"));
let pw = |n: &str, rows: usize, k: usize| -> Result<Band> {
let mut v = st.tensor_f32(n)?;
anyhow::ensure!(v.len() == rows * k, "{n} len {} != {rows}x{k}", v.len());
let rt = q4sim
&& (only.is_empty() || only.iter().any(|o| n.contains(o.as_str())))
&& !spare.iter().any(|o| n.contains(o.as_str()));
if rt {
let q8_list = parse("MOSHI_Q4SIM_Q8");
let group: usize = std::env::var("MOSHI_Q4SIM_GROUP")
.ok()
.and_then(|g| g.parse().ok())
.unwrap_or(32);
if q8_list.iter().any(|o| n.contains(o.as_str())) {
qn_roundtrip(&mut v, rows, k, 8, 32);
} else if std::env::var_os("MOSHI_Q4SIM_AFFINE").is_some() {
q4_affine_roundtrip(&mut v, rows, k);
} else if group != 32 {
qn_roundtrip(&mut v, rows, k, 4, group);
} else {
q4_roundtrip(&mut v, rows, k);
}
}
Ok(if i8 {
Band::I8(PackedWeightI8::new(&v, rows, k))
} else {
Band::F32(PackedWeight::new(&v, rows, k))
})
};
let mut emb = Vec::with_capacity(16);
for k in 0..16 {
emb.push(t(&format!("emb.{k}.weight"))?);
}
let mut layers = Vec::with_capacity(LAYERS);
for l in 0..LAYERS {
let p = format!("transformer.layers.{l}");
layers.push(TemporalLayer {
norm1: t(&format!("{p}.norm1.alpha"))?,
in_proj: pw(&format!("{p}.self_attn.in_projs.0.weight"), 3 * DIM, DIM)?,
out_proj: pw(&format!("{p}.self_attn.out_projs.0.weight"), DIM, DIM)?,
norm2: t(&format!("{p}.norm2.alpha"))?,
gate_in: pw(&format!("{p}.gating.linear_in.weight"), 2 * HIDDEN, DIM)?,
gate_out: pw(&format!("{p}.gating.linear_out.weight"), DIM, HIDDEN)?,
});
}
let mut dep_layers = Vec::with_capacity(DLAYERS);
for l in 0..DLAYERS {
let p = format!("depformer.layers.{l}");
let mut dl = DepLayer {
norm1: t(&format!("{p}.norm1.alpha"))?,
in_projs: Vec::with_capacity(DEP_Q),
out_projs: Vec::with_capacity(DEP_Q),
norm2: t(&format!("{p}.norm2.alpha"))?,
gate_in: Vec::with_capacity(DEP_Q),
gate_out: Vec::with_capacity(DEP_Q),
};
for s in 0..DEP_Q {
dl.in_projs.push(pw(&format!("{p}.self_attn.in_projs.{s}.weight"), 3 * DDIM, DDIM)?);
dl.out_projs.push(pw(&format!("{p}.self_attn.out_projs.{s}.weight"), DDIM, DDIM)?);
dl.gate_in.push(pw(&format!("{p}.gating.{s}.linear_in.weight"), 2 * DHIDDEN, DDIM)?);
dl.gate_out.push(pw(&format!("{p}.gating.{s}.linear_out.weight"), DDIM, DHIDDEN)?);
}
dep_layers.push(dl);
}
let mut dep_in = Vec::with_capacity(DEP_Q);
let mut linears = Vec::with_capacity(DEP_Q);
for s in 0..DEP_Q {
dep_in.push(pw(&format!("depformer_in.{s}.weight"), DDIM, DIM)?);
linears.push(pw(&format!("linears.{s}.weight"), CARD as usize, DDIM)?);
}
let mut dep_emb = Vec::with_capacity(DEP_Q - 1);
for k in 0..DEP_Q - 1 {
dep_emb.push(t(&format!("depformer_emb.{k}.weight"))?);
}
Ok(Self {
emb,
text_emb: t("text_emb.weight")?,
layers,
out_norm: t("out_norm.alpha")?,
text_linear: pw("text_linear.weight", TEXT_CARD as usize, DIM)?,
dep_in,
dep_text_emb: t("depformer_text_emb.weight")?,
dep_emb,
dep_layers,
linears,
sampler: None,
})
}
pub fn set_sampling(&mut self, s: Sampling, seed: u64) {
self.sampler = Some((s, std::cell::RefCell::new(seed.max(1))));
}
fn pick(&self, logits: &[f32], text: bool) -> u32 {
let Some((cfg, rng)) = &self.sampler else {
return argmax(logits);
};
let (temp, k) = if text { (cfg.temp_text, cfg.top_k_text) } else { (cfg.temp, cfg.top_k) };
let mut st = rng.borrow_mut();
if cfg.gumbel_full {
gumbel_full_pick(logits, temp, &mut st)
} else {
topk_sample(logits, temp, k, &mut st)
}
}
pub fn state(&self) -> MoshiState {
MoshiState {
cache: [[UNGENERATED; CT]; N_STREAMS],
offset: 0,
k_cache: vec![Vec::new(); LAYERS],
v_cache: vec![Vec::new(); LAYERS],
}
}
fn embed(table: &[f32], dim: usize, tok: i64, acc: &mut [f32]) {
if tok == -1 {
return;
}
let tok = tok.max(0) as usize;
for (a, v) in acc.iter_mut().zip(&table[tok * dim..(tok + 1) * dim]) {
*a += v;
}
}
fn temporal_step(&self, x: &mut Vec<f32>, st: &mut MoshiState) -> Vec<f32> {
let pos = st.offset;
let mut h = vec![0f32; DIM];
let mut qkv = vec![0f32; 3 * DIM];
let mut attn_out = vec![0f32; DIM];
let mut proj = vec![0f32; DIM];
let mut gate = vec![0f32; 2 * HIDDEN];
let mut ffn = vec![0f32; DIM];
for (l, layer) in self.layers.iter().enumerate() {
rms_norm(x, &layer.norm1, &mut h);
layer.in_proj.mv(&h, &mut qkv);
let (q, kvv) = qkv.split_at_mut(DIM);
let (k, v) = kvv.split_at_mut(DIM);
rope_row(q, pos, HEAD_DIM);
rope_row(k, pos, HEAD_DIM);
st.k_cache[l].push(k.to_vec());
st.v_cache[l].push(v.to_vec());
if st.k_cache[l].len() > CONTEXT {
st.k_cache[l].remove(0);
st.v_cache[l].remove(0);
}
attend(q, &st.k_cache[l], &st.v_cache[l], HEADS, HEAD_DIM, &mut attn_out);
layer.out_proj.mv(&attn_out, &mut proj);
for (xi, p) in x.iter_mut().zip(&proj) {
*xi += p;
}
rms_norm(x, &layer.norm2, &mut h);
layer.gate_in.mv(&h, &mut gate);
let (a, b) = gate.split_at(HIDDEN);
let gated: Vec<f32> = a.iter().zip(b).map(|(g, m)| silu(*g) * m).collect();
layer.gate_out.mv(&gated, &mut ffn);
for (xi, f) in x.iter_mut().zip(&ffn) {
*xi += f;
}
}
let mut out = vec![0f32; DIM];
rms_norm(x, &self.out_norm, &mut out);
out
}
fn depformer_cycle(&self, transformer_out: &[f32], text_token: u32) -> (Vec<Vec<f32>>, [u32; DEP_Q]) {
let mut k_cache: Vec<Vec<Vec<f32>>> = vec![Vec::with_capacity(DEP_Q); DLAYERS];
let mut v_cache: Vec<Vec<Vec<f32>>> = vec![Vec::with_capacity(DEP_Q); DLAYERS];
let mut all_logits = Vec::with_capacity(DEP_Q);
let mut tokens = [0u32; DEP_Q];
let mut prev: i64 = text_token as i64;
for cb in 0..DEP_Q {
let mut x = vec![0f32; DDIM];
self.dep_in[cb].mv(transformer_out, &mut x);
let table = if cb == 0 { &self.dep_text_emb } else { &self.dep_emb[cb - 1] };
Self::embed(table, DDIM, prev, &mut x);
let mut h = vec![0f32; DDIM];
let mut qkv = vec![0f32; 3 * DDIM];
let mut attn_out = vec![0f32; DDIM];
let mut proj = vec![0f32; DDIM];
let mut gate = vec![0f32; 2 * DHIDDEN];
let mut ffn = vec![0f32; DDIM];
for (l, layer) in self.dep_layers.iter().enumerate() {
rms_norm(&x, &layer.norm1, &mut h);
layer.in_projs[cb].mv(&h, &mut qkv);
let (q, kvv) = qkv.split_at(DDIM);
let (k, v) = kvv.split_at(DDIM);
k_cache[l].push(k.to_vec());
v_cache[l].push(v.to_vec());
attend(q, &k_cache[l], &v_cache[l], DHEADS, DHEAD_DIM, &mut attn_out);
layer.out_projs[cb].mv(&attn_out, &mut proj);
for (xi, p) in x.iter_mut().zip(&proj) {
*xi += p;
}
rms_norm(&x, &layer.norm2, &mut h);
layer.gate_in[cb].mv(&h, &mut gate);
let (a, b) = gate.split_at(DHIDDEN);
let gated: Vec<f32> = a.iter().zip(b).map(|(g, m)| silu(*g) * m).collect();
layer.gate_out[cb].mv(&gated, &mut ffn);
for (xi, f) in x.iter_mut().zip(&ffn) {
*xi += f;
}
}
let mut logits = vec![0f32; CARD as usize];
self.linears[cb].mv(&x, &mut logits);
let tok = self.pick(&logits, false);
all_logits.push(logits);
tokens[cb] = tok;
prev = tok as i64;
}
(all_logits, tokens)
}
pub fn step(&self, st: &mut MoshiState, user_codes: &[u32; DEP_Q]) -> StepTrace {
self.step_ext(st, user_codes, None)
}
pub fn step_forced(
&self,
st: &mut MoshiState,
forced_text: u32,
forced_agent: &[u32; DEP_Q],
user_codes: &[u32; DEP_Q],
temporal: &mut dyn FnMut(&[f32], usize),
) {
for (k, &tok) in user_codes.iter().enumerate() {
let s = DEP_Q + 1 + k;
st.cache[s][(st.offset + DELAYS[s]) % CT] = tok as i64;
}
let mut input = [0i64; N_STREAMS];
for s in 0..N_STREAMS {
input[s] = if st.offset <= DELAYS[s] {
if s == 0 { TEXT_CARD as i64 } else { CARD as i64 }
} else {
st.cache[s][st.offset % CT]
};
}
let mut x = vec![0f32; DIM];
for (k, table) in self.emb.iter().enumerate() {
Self::embed(table, DIM, input[1 + k], &mut x);
}
Self::embed(&self.text_emb, DIM, input[0], &mut x);
temporal(&x, st.offset);
st.offset += 1;
let pos = st.offset % CT;
st.cache[0][pos] = forced_text as i64;
for (cb, &tok) in forced_agent.iter().enumerate() {
st.cache[1 + cb][pos] = tok as i64;
}
}
pub fn step_full_ext(
&self,
st: &mut MoshiState,
user_codes: &[u32; DEP_Q],
run: &mut dyn FnMut(&[f32], usize) -> (u32, [u32; DEP_Q]),
) -> (u32, [u32; DEP_Q], Option<[i32; DEP_Q + 1]>) {
for (k, &tok) in user_codes.iter().enumerate() {
let s = DEP_Q + 1 + k;
st.cache[s][(st.offset + DELAYS[s]) % CT] = tok as i64;
}
let mut input = [0i64; N_STREAMS];
for s in 0..N_STREAMS {
input[s] = if st.offset <= DELAYS[s] {
if s == 0 { TEXT_CARD as i64 } else { CARD as i64 }
} else {
st.cache[s][st.offset % CT]
};
}
let mut x = vec![0f32; DIM];
for (k, table) in self.emb.iter().enumerate() {
Self::embed(table, DIM, input[1 + k], &mut x);
}
Self::embed(&self.text_emb, DIM, input[0], &mut x);
let (text_token, audio_tokens) = run(&x, st.offset);
st.offset += 1;
let pos = st.offset % CT;
st.cache[0][pos] = text_token as i64;
for (cb, &tok) in audio_tokens.iter().enumerate() {
st.cache[1 + cb][pos] = tok as i64;
}
let out = (st.offset > MAX_DELAY).then(|| {
let mut o = [0i32; DEP_Q + 1];
for (s, oo) in o.iter_mut().enumerate() {
*oo = st.cache[s][(st.offset - MAX_DELAY + DELAYS[s]) % CT] as i32;
}
o
});
(text_token, audio_tokens, out)
}
pub fn step_fast(
&self,
st: &mut MoshiState,
user_codes: &[u32; DEP_Q],
temporal: &mut dyn FnMut(&[f32], usize) -> (Vec<f32>, u32),
) -> (u32, [u32; DEP_Q], Option<[i32; DEP_Q + 1]>) {
for (k, &tok) in user_codes.iter().enumerate() {
let s = DEP_Q + 1 + k;
st.cache[s][(st.offset + DELAYS[s]) % CT] = tok as i64;
}
let mut input = [0i64; N_STREAMS];
for s in 0..N_STREAMS {
input[s] = if st.offset <= DELAYS[s] {
if s == 0 { TEXT_CARD as i64 } else { CARD as i64 }
} else {
st.cache[s][st.offset % CT]
};
}
let mut x = vec![0f32; DIM];
for (k, table) in self.emb.iter().enumerate() {
Self::embed(table, DIM, input[1 + k], &mut x);
}
Self::embed(&self.text_emb, DIM, input[0], &mut x);
let (transformer_out, text_token) = temporal(&x, st.offset);
let (_logits, audio_tokens) = self.depformer_cycle(&transformer_out, text_token);
st.offset += 1;
let pos = st.offset % CT;
st.cache[0][pos] = text_token as i64;
for (cb, &tok) in audio_tokens.iter().enumerate() {
st.cache[1 + cb][pos] = tok as i64;
}
let out = (st.offset > MAX_DELAY).then(|| {
let mut o = [0i32; DEP_Q + 1];
for (s, oo) in o.iter_mut().enumerate() {
*oo = st.cache[s][(st.offset - MAX_DELAY + DELAYS[s]) % CT] as i32;
}
o
});
(text_token, audio_tokens, out)
}
pub fn step_ext(
&self,
st: &mut MoshiState,
user_codes: &[u32; DEP_Q],
temporal: Option<&mut dyn FnMut(&[f32], usize) -> (Vec<f32>, Vec<f32>)>,
) -> StepTrace {
for (k, &tok) in user_codes.iter().enumerate() {
let s = DEP_Q + 1 + k;
st.cache[s][(st.offset + DELAYS[s]) % CT] = tok as i64;
}
let mut input = [0i64; N_STREAMS];
for s in 0..N_STREAMS {
input[s] = if st.offset <= DELAYS[s] {
if s == 0 { TEXT_CARD as i64 } else { CARD as i64 }
} else {
st.cache[s][st.offset % CT]
};
}
let mut x = vec![0f32; DIM];
for (k, table) in self.emb.iter().enumerate() {
Self::embed(table, DIM, input[1 + k], &mut x);
}
Self::embed(&self.text_emb, DIM, input[0], &mut x);
let emb_sum = x.clone();
let timing = std::env::var_os("MOSHI_TIME").is_some();
let t0 = std::time::Instant::now();
let (transformer_out, text_logits) = match temporal {
Some(f) => f(&x, st.offset),
None => {
let out = self.temporal_step(&mut x, st);
let mut logits = vec![0f32; TEXT_CARD as usize];
self.text_linear.mv(&out, &mut logits);
(out, logits)
}
};
let t1 = t0; let text_token = self.pick(&text_logits, true);
let t2 = std::time::Instant::now();
let (audio_logits, audio_tokens) = self.depformer_cycle(&transformer_out, text_token);
if timing {
eprintln!(
" temporal {:.0} ms | text head {:.1} ms | depformer {:.0} ms",
t1.duration_since(t0).as_secs_f64() * 1e3,
t2.duration_since(t1).as_secs_f64() * 1e3,
t2.elapsed().as_secs_f64() * 1e3,
);
}
st.offset += 1;
let pos = st.offset % CT;
st.cache[0][pos] = text_token as i64;
for (cb, &tok) in audio_tokens.iter().enumerate() {
st.cache[1 + cb][pos] = tok as i64;
}
let out = (st.offset > MAX_DELAY).then(|| {
let mut o = [0i32; DEP_Q + 1];
for (s, oo) in o.iter_mut().enumerate() {
*oo = st.cache[s][(st.offset - MAX_DELAY + DELAYS[s]) % CT] as i32;
}
o
});
StepTrace {
input_tokens: input,
emb_sum,
transformer_out,
text_logits,
audio_logits,
text_token,
audio_tokens,
out,
}
}
}