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
//! `pocket-tts` — Kyutai's 100M CPU text-to-speech on the engine. Text + a voice → a 24 kHz WAV,
//! streamed frame by frame (80 ms each) so the first audio lands long before the last.
//!
//!   pocket-tts <model-dir> "Hello world." [out.wav]
//!   POCKET_TTS_VOICE=javert pocket-tts <model-dir> "Hello world."
//!   POCKET_TTS_TEMP=0.7 POCKET_TTS_SEED=42 pocket-tts <model-dir> "Hello world."
//!
//! `<model-dir>` is one language directory of the `kyutai/pocket-tts` (or
//! `…-without-voice-cloning`) repo: `model.safetensors`, `tokenizer.model`, `embeddings/*.
//! safetensors`. Voices are pre-computed KV caches, so `--voice` costs nothing to switch.

use std::path::PathBuf;
use std::time::Instant;

use anyhow::{Context, Result};
use inferencelayer::pocket_tts::{GenOpts, PocketTts};

fn env_parse<T: std::str::FromStr>(key: &str, default: T) -> T {
    std::env::var(key)
        .ok()
        .and_then(|v| v.parse().ok())
        .unwrap_or(default)
}

/// Read a 16-bit PCM WAV → (mono samples, sample rate).
fn read_wav(path: &PathBuf) -> Result<(Vec<f32>, u32)> {
    let b = std::fs::read(path).with_context(|| format!("read {}", path.display()))?;
    let fmt = b
        .windows(4)
        .position(|w| w == b"fmt ")
        .context("no `fmt ` chunk")?;
    let ch = u16::from_le_bytes(b[fmt + 10..fmt + 12].try_into()?) as usize;
    let sr = u32::from_le_bytes(b[fmt + 12..fmt + 16].try_into()?);
    let dp = b
        .windows(4)
        .position(|w| w == b"data")
        .context("no `data` chunk")?
        + 8;
    let n = u32::from_le_bytes(b[dp - 4..dp].try_into()?) as usize;
    let pcm: Vec<f32> = b[dp..dp + n]
        .chunks_exact(2)
        .map(|c| i16::from_le_bytes([c[0], c[1]]) as f32 / 32768.0)
        .collect();
    let mono = if ch > 1 {
        pcm.chunks(ch)
            .map(|f| f.iter().sum::<f32>() / ch as f32)
            .collect()
    } else {
        pcm
    };
    Ok((mono, sr))
}

/// 16-bit PCM mono WAV.
fn write_wav(path: &PathBuf, pcm: &[f32], rate: u32) -> Result<()> {
    let mut b = Vec::with_capacity(44 + pcm.len() * 2);
    let data_len = (pcm.len() * 2) as u32;
    b.extend(b"RIFF");
    b.extend((36 + data_len).to_le_bytes());
    b.extend(b"WAVEfmt ");
    b.extend(16u32.to_le_bytes());
    b.extend(1u16.to_le_bytes()); // PCM
    b.extend(1u16.to_le_bytes()); // mono
    b.extend(rate.to_le_bytes());
    b.extend((rate * 2).to_le_bytes());
    b.extend(2u16.to_le_bytes());
    b.extend(16u16.to_le_bytes());
    b.extend(b"data");
    b.extend(data_len.to_le_bytes());
    for v in pcm {
        b.extend(((v.clamp(-1.0, 1.0) * 32767.0) as i16).to_le_bytes());
    }
    std::fs::write(path, b).with_context(|| format!("write {}", path.display()))
}

fn main() -> Result<()> {
    let mut args = std::env::args().skip(1);
    let usage = "usage: pocket-tts <model-dir> \"text to speak\" [out.wav]";
    let dir = PathBuf::from(args.next().context(usage)?);
    let text = args.next().context(usage)?;
    let out = PathBuf::from(args.next().unwrap_or_else(|| "pocket_tts.wav".to_string()));

    let t0 = Instant::now();
    let tts = PocketTts::load(&dir)?;
    // A catalog name, or a path to a 16-bit PCM WAV to clone (gated checkpoint only).
    let voice_name = std::env::var("POCKET_TTS_VOICE").unwrap_or_else(|_| "alba".to_string());
    let voice = if voice_name.to_ascii_lowercase().ends_with(".wav") {
        let (pcm, sr) = read_wav(&PathBuf::from(&voice_name))?;
        eprintln!(
            "cloning from {voice_name} ({sr} Hz, {:.1}s)",
            pcm.len() as f32 / sr as f32
        );
        tts.clone_voice(&inferencelayer::chatterbox::resample_sinc(&pcm, sr, 24_000))?
    } else {
        tts.load_voice(dir.join(format!("embeddings/{voice_name}.safetensors")))?
    };
    eprintln!(
        "loaded {} ({} layers, d{}) + voice {voice_name} ({} positions) in {:?}",
        dir.display(),
        tts.config().num_layers,
        tts.config().d_model,
        voice.len(),
        t0.elapsed()
    );

    let opts = GenOpts {
        temperature: env_parse("POCKET_TTS_TEMP", 0.3),
        seed: env_parse("POCKET_TTS_SEED", GenOpts::default().seed),
        ..GenOpts::default()
    };

    let t0 = Instant::now();
    let mut pcm = Vec::new();
    let mut first_frame = None;
    tts.generate_streaming(&voice, &text, &opts, |frame| {
        first_frame.get_or_insert_with(|| t0.elapsed());
        pcm.extend_from_slice(frame);
    })?;
    let elapsed = t0.elapsed();

    let secs = pcm.len() as f32 / tts.sample_rate() as f32;
    eprintln!(
        "{:.2}s of audio in {:?} ({:.2}x real-time), first frame after {:?}",
        secs,
        elapsed,
        secs / elapsed.as_secs_f32(),
        first_frame.unwrap_or_default()
    );
    write_wav(&out, &pcm, tts.sample_rate() as u32)?;
    eprintln!("wrote {}", out.display());
    Ok(())
}