use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use anyhow::{Context, Result};
use cpal::traits::{DeviceTrait, HostTrait, StreamTrait};
use inferencelayer::mimi::{FRAME_SIZE, Mimi};
use inferencelayer::moshi_lm::{MoshiLm, Sampling};
const SR: usize = 24000;
struct Hook {
phrase: String,
cmd: String,
last_fire: std::time::Instant,
}
fn main() -> Result<()> {
let mut args = std::env::args().skip(1);
let mut gpu = true;
let mut greedy_gpu = false;
let mut seed = 299_792_458u64;
let mut hooks: Vec<Hook> = Vec::new();
let mut list = false;
while let Some(a) = args.next() {
match a.as_str() {
"--cpu" => gpu = false,
"--greedy-gpu" => greedy_gpu = true,
"--seed" => seed = args.next().context("--seed")?.parse()?,
"--list-devices" => list = true,
"--action" => {
let spec = args.next().context("--action")?;
let (phrase, cmd) = spec.split_once('=').context("--action phrase=command")?;
hooks.push(Hook {
phrase: phrase.to_lowercase(),
cmd: cmd.to_string(),
last_fire: std::time::Instant::now() - std::time::Duration::from_secs(60),
});
}
other => anyhow::bail!("unknown arg {other}"),
}
}
let host = cpal::default_host();
if list {
for d in host.input_devices()? {
println!("in : {}", d.name().unwrap_or_default());
}
for d in host.output_devices()? {
println!("out: {}", d.name().unwrap_or_default());
}
return Ok(());
}
let cache = std::path::PathBuf::from(std::env::var("HOME")?).join(".cache/inferencelayer");
eprintln!("loading Mimi (CPU) …");
let mimi = Mimi::load(&cache.join("mimi"))?;
eprintln!("loading Moshi 7B embeddings+depformer (CPU i8) …");
let mut lm = MoshiLm::load_i8(&cache.join("moshi-lm"))?;
lm.set_sampling(Sampling::default(), seed);
let gpu_lm = if gpu {
eprintln!("loading temporal onto wgpu (q8 FFN + q4 attention) …");
let ctx = inferencelayer::GpuCtx::new().context("no GPU — rerun with --cpu")?;
let w = inferencelayer::Weights::load(&ctx, cache.join("moshi-lm"))?;
let g = inferencelayer::Lfm2Gpu::new(&ctx, w);
let dep = if greedy_gpu {
eprintln!("loading depformer onto wgpu (q8, fused greedy step) …");
Some(inferencelayer::moshi_gpu::MoshiDepGpu::load(&ctx, cache.join("moshi-lm"), &g)?)
} else {
None
};
Some((ctx, g, dep))
} else {
eprintln!("CPU temporal (i8): expect ~4× slower than realtime — for smoke tests only");
None
};
let pieces = load_spm_pieces()?;
let in_dev = host.default_input_device().context("no input device")?;
let out_dev = host.default_output_device().context("no output device")?;
let in_cfg = in_dev.default_input_config()?;
let out_cfg = out_dev.default_output_config()?;
anyhow::ensure!(
in_cfg.sample_format() == cpal::SampleFormat::F32
&& out_cfg.sample_format() == cpal::SampleFormat::F32,
"demo expects f32 audio devices (got {:?}/{:?})",
in_cfg.sample_format(),
out_cfg.sample_format()
);
let (in_rate, in_ch) = (in_cfg.sample_rate().0 as usize, in_cfg.channels() as usize);
let (out_rate, out_ch) = (out_cfg.sample_rate().0 as usize, out_cfg.channels() as usize);
eprintln!(
"mic: {} @{in_rate} Hz ×{in_ch} | speaker: {} @{out_rate} Hz ×{out_ch}",
in_dev.name().unwrap_or_default(),
out_dev.name().unwrap_or_default()
);
let in_q: Arc<Mutex<VecDeque<f32>>> = Arc::default(); let out_q: Arc<Mutex<VecDeque<f32>>> = Arc::default();
let qi = in_q.clone();
let in_stream = in_dev.build_input_stream(
&in_cfg.into(),
move |data: &[f32], _: &_| {
let mut q = qi.lock().unwrap();
for fr in data.chunks_exact(in_ch) {
q.push_back(fr.iter().sum::<f32>() / in_ch as f32);
}
let cap = in_rate * 2;
while q.len() > cap {
q.pop_front();
}
},
|e| eprintln!("mic error: {e}"),
None,
)?;
let qo = out_q.clone();
let out_stream = out_dev.build_output_stream(
&out_cfg.into(),
move |data: &mut [f32], _: &_| {
let mut q = qo.lock().unwrap();
for fr in data.chunks_exact_mut(out_ch) {
let s = q.pop_front().unwrap_or(0.0);
fr.fill(s);
}
},
|e| eprintln!("speaker error: {e}"),
None,
)?;
in_stream.play()?;
out_stream.play()?;
let per_frame_in = FRAME_SIZE * in_rate / SR; let mut st = lm.state();
let (tok_tx, tok_rx) = std::sync::mpsc::channel::<[u32; inferencelayer::moshi_lm::DEP_Q]>();
let dec_out = out_q.clone();
let dec_mimi = std::sync::Arc::new(mimi);
let dec_handle = {
let mimi = dec_mimi.clone();
std::thread::spawn(move || {
let mut dec = mimi.stream();
while let Ok(codes) = tok_rx.recv() {
let voice = dec.decode_frame(&codes);
dec_out
.lock()
.unwrap()
.extend(resample(&voice, FRAME_SIZE * out_rate / SR));
}
})
};
let mimi = dec_mimi;
let mut enc = mimi.stream();
let mut transcript = String::new();
let mut frames = 0u64;
let mut slow = 0u32;
out_q.lock().unwrap().extend(std::iter::repeat_n(0.0f32, out_rate * 32 / 100));
let _ = &dec_handle;
eprintln!("\n=== live — talk now (Ctrl-C to quit) ===\n");
loop {
let chunk: Vec<f32> = loop {
{
let mut q = in_q.lock().unwrap();
if q.len() >= per_frame_in {
break q.drain(..per_frame_in).collect();
}
}
std::thread::sleep(std::time::Duration::from_millis(2));
};
let t0 = std::time::Instant::now();
let frame = resample(&chunk, FRAME_SIZE);
let user = enc.encode_frame(&frame);
let (text_token, audio_tokens) = match &gpu_lm {
Some((ctx, g, Some(dep))) => {
let (tt, at, _) = lm.step_full_ext(
&mut st,
&user,
&mut |emb: &[f32], pos: usize| {
inferencelayer::moshi_gpu::moshi_full_step(g, dep, ctx, emb, pos)
.expect("full gpu step")
},
);
(tt, at)
}
Some((ctx, g, None)) => {
let tr = lm.step_ext(
&mut st,
&user,
Some(&mut |emb: &[f32], pos: usize| {
let logits = g.forward_from_embeds(ctx, emb, pos).expect("gpu forward");
let hidden = g.read_hnorm_for_tests(ctx).expect("hnorm");
(hidden, logits)
}),
);
(tr.text_token, tr.audio_tokens)
}
None => {
let tr = lm.step(&mut st, &user);
(tr.text_token, tr.audio_tokens)
}
};
tok_tx.send(audio_tokens).ok();
let tok = text_token as usize;
if tok != 0 && tok != 3
&& let Some(p) = pieces.get(tok) {
let w = p.replace('\u{2581}', " ");
print!("{w}");
use std::io::Write;
std::io::stdout().flush().ok();
transcript.push_str(&w);
if transcript.len() > 400 {
let cut = transcript.len() - 400;
transcript.drain(..cut);
}
run_hooks(&mut hooks, &transcript);
}
frames += 1;
let ms = t0.elapsed().as_secs_f64() * 1e3;
if ms > 80.0 {
slow += 1;
if slow % 25 == 1 {
eprintln!("\n[warn] frame took {ms:.0} ms (>80); speaker buffer absorbs occasional spikes");
}
}
if frames.is_multiple_of(125) {
eprintln!("\n[{}s] {} frames, {} over budget", frames / 125 * 10, frames, slow);
}
}
}
fn run_hooks(hooks: &mut [Hook], transcript: &str) {
let recent = transcript.to_lowercase();
for h in hooks.iter_mut() {
if h.last_fire.elapsed().as_secs() >= 10 && recent.contains(&h.phrase) {
h.last_fire = std::time::Instant::now();
eprintln!("\n[action \"{}\"] → {}", h.phrase, h.cmd);
let cmd = h.cmd.clone();
let text = transcript.to_string();
std::thread::spawn(move || {
match std::process::Command::new("sh")
.arg("-c")
.arg(&cmd)
.env("MOSHI_TEXT", text)
.output()
{
Ok(o) => eprintln!("[action out] {}", String::from_utf8_lossy(&o.stdout).trim()),
Err(e) => eprintln!("[action err] {e}"),
}
});
}
}
}
fn resample(x: &[f32], n_out: usize) -> Vec<f32> {
if x.len() == n_out {
return x.to_vec();
}
let ratio = x.len() as f64 / n_out as f64;
(0..n_out)
.map(|i| {
let p = i as f64 * ratio;
let (a, t) = (p as usize, p.fract() as f32);
let b = (a + 1).min(x.len() - 1);
x[a] * (1.0 - t) + x[b] * t
})
.collect()
}
fn load_spm_pieces() -> Result<Vec<String>> {
inferencelayer::moshi_lm::spm_pieces(&inferencelayer::moshi_lm::find_spm_model()?)
}