inferencelayer 0.2.4

Kortexya's engine-native inference layer — LLM generation + embedding/encoder family on wgpu (WGSL kernels, any adapter) with a pure-Rust CPU fallback
Documentation
//! `moshi-live` — a LIVE full-duplex conversation with Moshi on this engine: microphone in,
//! Moshi's voice out of the speaker, rolling transcript on the console. No turns — talk over
//! it, interrupt it, it hears you while it speaks (that's the architecture, not plumbing).
//!
//! ```text
//! moshi-live [--cpu] [--greedy-gpu] [--seed N] [--action "phrase=shell command"]... [--list-devices]
//! ```
//!
//! Action hooks (demo-grade): each `--action` watches the LIVE transcript for a phrase
//! (case-insensitive substring); on match the shell command runs async with `MOSHI_TEXT` set
//! to the recent transcript, and its output prints to the console. Moshi does NOT consume the
//! result mid-conversation — feeding results back in is the MoshiRAG-style trigger-token
//! work (native tool calling), a later milestone.
//!
//! The loop must sustain 12.5 Hz: one 80 ms frame = Mimi encode + 7B step (wgpu) + Mimi
//! decode, measured ~75 ms/frame on an M4 Max. A ~4-frame speaker pre-buffer rides p95
//! spikes; sustained overruns print a warning rather than crashing.

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,
            // The fused full-GPU LM step (temporal + depformer, one submit): lowest latency,
            // GREEDY sampling only — the default sampled path sounds more natural.
            "--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(());
    }

    // ---- models ------------------------------------------------------------------------------
    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()?;

    // ---- audio devices (f32 streams; device rate → linear resample to/from 24 kHz) ----------
    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(); // mono @ device rate
    let out_q: Arc<Mutex<VecDeque<f32>>> = Arc::default(); // mono @ device rate

    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);
            }
            // never let a stall grow the buffer beyond ~2 s
            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()?;

    // ---- the live loop -----------------------------------------------------------------------
    let per_frame_in = FRAME_SIZE * in_rate / SR; // device samples per 80 ms
    let mut st = lm.state();
    // Mimi DECODE runs on its own thread: it depends only on the frame's generated codes, so
    // it overlaps the next frame's encode+LM — critical path drops from ~80 ms to ~60 ms.
    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;
    // speaker pre-buffer: ~4 frames of silence so p95 model spikes don't underrun
    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 {
        // block until one frame of mic audio is available
        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}"),
                }
            });
        }
    }
}

/// Linear resample to exactly `n_out` samples — chunk-local (demo-grade).
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()?)
}