Files
rustytorch/crates/models/rtx-csm/src/converse.rs
T
2026-05-07 16:30:04 +00:00

473 lines
18 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Composable orchestrator for the LLM → TTS half of the conversational
//! stack. Takes a prompt + chat history, streams LLM tokens, buffers them
//! into sentences, and dispatches each completed sentence to CSM TTS as
//! soon as it's ready. PCM is emitted as the model produces it.
//!
//! ## Why sentence buffering?
//!
//! CSM is a sentence-level TTS — its prosody is best when fed a complete
//! sentence at a time, not token-by-token. The LLM streams tokens; we
//! accumulate until a terminal punctuation mark (`.`, `!`, `?`) or newline,
//! then flush the sentence to TTS. This trades a small chunk of buffering
//! latency (typically <1s for short sentences) for natural prosody.
//!
//! ## Pipeline
//!
//! ```text
//! prompt + history
//! │
//! ▼
//! LlmClient.generate_stream → token chunks (k1, k2, ...)
//! │
//! ▼ buffer until terminal punctuation
//! Sentence "Hello."
//! │
//! ▼
//! Generator.generate(text) → 24 kHz PCM (Vec<f32>)
//! │
//! ▼ on_audio callback fires once per sentence
//! caller (e.g. WebSocket sink, file writer)
//! ```
//!
//! When the next implementation lands a streaming TTS variant that takes
//! token-level input directly (no sentence boundary needed), the
//! orchestrator can be swapped to that path. Today's CSM is the bottleneck.
//!
//! ## Half-duplex; full duplex (6c.2) is deferred
//!
//! This orchestrator handles the LLM → TTS direction only. To get full
//! voice-in/voice-out, pair this with the Phase 6a STT once its word
//! emission is fixed: `audio_in → STT → user_text → Converse → audio_out`.
use crate::GenerateOptions;
use crate::audio_io;
use crate::error::{CsmError, Result};
use crate::generator::Generator;
use crate::llm_client::{ChatMessage, GenConfig, LlmClient};
use crate::post::PostProcess;
use futures_util::StreamExt;
use std::path::Path;
/// One unit emitted by [`Converse::run`] each time a sentence completes.
#[derive(Debug, Clone)]
pub struct Utterance {
pub text: String,
/// 24 kHz mono PCM, post-processed and (if a watermarker is installed
/// on the Generator) watermarked.
pub audio: Vec<f32>,
/// Wall-clock time-to-first-audio for THIS sentence, in milliseconds.
/// (Time from sentence-buffer flush to PCM produced.)
pub tts_latency_ms: u128,
}
/// How aggressively to flush sentences. `Punctuation` waits for `. ! ?`
/// or newline; `Eager` flushes more often for lower latency at the cost
/// of less natural prosody.
#[derive(Debug, Clone, Copy)]
pub enum FlushPolicy {
Punctuation,
/// Flush at every comma, semicolon, or punctuation mark — useful for
/// long-running monologues where you want low first-audio latency.
Eager,
}
#[derive(Debug, Clone)]
pub struct ConverseOptions {
pub generate: GenerateOptions,
pub speaker: u32,
pub flush: FlushPolicy,
/// Skip sentences shorter than this (after trim). Avoids dispatching
/// tiny "k." or single-character fragments to TTS.
pub min_sentence_chars: usize,
}
impl Default for ConverseOptions {
fn default() -> Self {
Self {
generate: GenerateOptions {
max_audio_ms: 6_000,
..GenerateOptions::default()
},
speaker: 0,
flush: FlushPolicy::Punctuation,
min_sentence_chars: 2,
}
}
}
/// Locate the byte index AFTER the first sentence-boundary character in
/// the buffer per policy. Returns `Some(end)` so the caller can do
/// `buf.drain(..end)` to extract the completed sentence (including the
/// punctuation mark and any trailing whitespace up to the boundary).
/// Returns `None` if no boundary is present yet.
fn find_first_boundary(buf: &str, policy: FlushPolicy) -> Option<usize> {
let is_boundary = |c: char| match policy {
FlushPolicy::Punctuation => matches!(c, '.' | '!' | '?' | '\n'),
FlushPolicy::Eager => matches!(c, '.' | '!' | '?' | '\n' | ',' | ';' | ':'),
};
for (i, c) in buf.char_indices() {
if is_boundary(c) {
return Some(i + c.len_utf8());
}
}
None
}
/// Buffer contains a flushable boundary anywhere (used by the
/// streaming-path tests). For the actual flush logic we use
/// [`find_first_boundary`] to get the byte index of the first boundary.
#[cfg(test)]
fn should_flush(buf: &str, policy: FlushPolicy) -> bool {
find_first_boundary(buf, policy).is_some()
}
/// Optional hook invoked with `&mut Generator` before each sentence's
/// synthesis call. Use this to apply per-sentence steering (e.g. emotion
/// shifts mid-reply). Sync — apply_steering is non-async — so the hook
/// fits cleanly between LLM token consumption and the synthesize call.
pub type PreSentenceHook = Box<dyn FnMut(&mut Generator, &str) -> Result<()> + Send + 'static>;
pub struct Converse<'a, L: LlmClient> {
llm: &'a L,
generator: &'a mut Generator,
post: PostProcess,
/// Persistent speaker-prompt context prepended to every sentence's
/// generation call. Without this, CSM-1B has no voice anchor and
/// drifts between speakers / pitches across turns. Typically a
/// single `Segment` carrying a 515s reference clip with its
/// transcript pins the voice for the whole session.
context: Vec<crate::prompt::Segment>,
/// Hook fired before each sentence's `synthesize_streaming` call.
/// Receives mutable access to the underlying Generator and the
/// sentence text, so callers can apply per-sentence steering or
/// other state changes without touching this crate's internals.
pre_sentence_hook: Option<PreSentenceHook>,
}
impl<'a, L: LlmClient> Converse<'a, L> {
pub fn new(llm: &'a L, generator: &'a mut Generator) -> Self {
Self {
llm,
generator,
post: PostProcess::default(),
context: Vec::new(),
pre_sentence_hook: None,
}
}
/// Install a hook that fires before each sentence's synth call,
/// with mutable access to the underlying Generator + the sentence
/// text. Use this to apply per-sentence steering (e.g. emotion
/// shifts mid-reply).
pub fn with_pre_sentence_hook(mut self, hook: PreSentenceHook) -> Self {
self.pre_sentence_hook = Some(hook);
self
}
/// Override post-processing (HPF + declick + LUFS). Pass
/// [`PostProcess::disabled`] to skip.
pub fn with_post(mut self, post: PostProcess) -> Self {
self.post = post;
self
}
/// Set persistent speaker-prompt context for the session.
/// Each segment provides a (speaker, text, audio) triple the model
/// uses as a voice anchor. Without this, CSM-1B picks a different
/// speaker per turn.
pub fn with_context(mut self, ctx: Vec<crate::prompt::Segment>) -> Self {
self.context = ctx;
self
}
/// Run the full pipeline: stream LLM tokens, flush sentences to TTS,
/// invoke `on_utterance` for each completed (text, audio) pair as
/// they're produced. Returns the full assistant message text once
/// the LLM stream ends.
pub async fn run<F>(
&mut self,
messages: Vec<ChatMessage>,
gen_cfg: GenConfig,
opts: ConverseOptions,
mut on_utterance: F,
) -> Result<String>
where
F: FnMut(&Utterance) -> Result<()>,
{
// Phase tracing: TTFT (time-to-first-token) for the LLM, and per-
// sentence llm_buffer (time accumulating tokens until a boundary)
// and tts_gen (time inside synthesize). All emitted at info level
// with a stable "conv-phase:" prefix so callers can filter or grep.
let run_t = std::time::Instant::now();
let mut stream = self.llm.generate_stream(messages, gen_cfg).await?;
let mut buf = String::new();
let mut full = String::new();
let mut first_token_logged = false;
let mut sentence_idx: usize = 0;
let mut sentence_t = std::time::Instant::now();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if !first_token_logged {
first_token_logged = true;
tracing::info!(
"conv-phase: ttft={}ms (LLM first token)",
run_t.elapsed().as_millis()
);
sentence_t = std::time::Instant::now();
}
full.push_str(&chunk);
buf.push_str(&chunk);
// Flush as many complete sentences as the buffer contains.
while let Some(end) = find_first_boundary(&buf, opts.flush) {
let sentence: String = buf.drain(..end).collect();
let trimmed = sentence.trim();
if trimmed.chars().count() < opts.min_sentence_chars {
continue;
}
let llm_buffer_ms = sentence_t.elapsed().as_millis();
if let Some(hook) = self.pre_sentence_hook.as_mut() {
hook(self.generator, trimmed)?;
}
let utt = self.synthesize(trimmed, &opts)?;
tracing::info!(
"conv-phase: sentence[{sentence_idx}] llm_buffer={llm_buffer_ms}ms \
tts_gen={}ms chars={}",
utt.tts_latency_ms,
trimmed.chars().count(),
);
sentence_idx += 1;
sentence_t = std::time::Instant::now();
on_utterance(&utt)?;
}
}
// Trailing text without terminal punctuation — flush as one
// final sentence so callers don't lose the tail.
let trailing = buf.trim().to_string();
if trailing.chars().count() >= opts.min_sentence_chars {
let llm_buffer_ms = sentence_t.elapsed().as_millis();
if let Some(hook) = self.pre_sentence_hook.as_mut() {
hook(self.generator, &trailing)?;
}
let utt = self.synthesize(&trailing, &opts)?;
tracing::info!(
"conv-phase: sentence[{sentence_idx}] (trailing) llm_buffer={llm_buffer_ms}ms \
tts_gen={}ms chars={}",
utt.tts_latency_ms,
trailing.chars().count(),
);
on_utterance(&utt)?;
}
Ok(full)
}
fn synthesize(&mut self, sentence: &str, opts: &ConverseOptions) -> Result<Utterance> {
let t = std::time::Instant::now();
let mut pcm = self
.generator
.generate(sentence, opts.speaker, &self.context, opts.generate.clone())
.map_err(|e| CsmError::Config(format!("converse generate: {e}")))?;
self.post
.apply(&mut pcm, self.generator.config.sample_rate)
.map_err(|e| CsmError::Config(format!("converse post: {e}")))?;
if let Some(wm) = self.generator.watermarker.as_ref() {
pcm = wm.embed(&pcm)?;
}
Ok(Utterance {
text: sentence.to_string(),
audio: pcm,
tts_latency_ms: t.elapsed().as_millis(),
})
}
/// Streaming variant of [`Self::run`]: for each sentence, the
/// underlying TTS uses [`Generator::generate_streaming`] with
/// `chunk_frames`-sized chunks (4 frames ≈ 320 ms at 12.5 Hz).
/// `on_chunk` fires for each chunk as it's produced — this is the
/// path that delivers low first-audio latency. `on_sentence` fires
/// once per completed sentence with the accumulated audio (useful
/// for /metrics + transcript bookkeeping; the audio has already
/// been streamed out via `on_chunk`).
///
/// **No post-process or watermark in this mode.** Both operations
/// require full-sentence context (LUFS normalize over the whole
/// utterance, AudioSeal needs ~1s of audio for a stable embed) and
/// can't run on partial chunks without re-sending audio that has
/// already left the server. Callers needing those features should
/// use [`Self::run`] instead. If a watermarker is installed on the
/// underlying generator, this method returns an error to make the
/// tradeoff explicit.
pub async fn run_streaming<C, U>(
&mut self,
messages: Vec<ChatMessage>,
gen_cfg: GenConfig,
opts: ConverseOptions,
chunk_frames: usize,
mut on_chunk: C,
mut on_sentence: U,
) -> Result<String>
where
C: FnMut(&[f32]) -> Result<()>,
U: FnMut(&Utterance) -> Result<()>,
{
if self.generator.watermarker.is_some() {
return Err(CsmError::Config(
"run_streaming: watermarker installed but streaming mode \
cannot embed (needs full-sentence context). Use run() \
instead, or remove the watermarker."
.into(),
));
}
let run_t = std::time::Instant::now();
let mut stream = self.llm.generate_stream(messages, gen_cfg).await?;
let mut buf = String::new();
let mut full = String::new();
let mut first_token_logged = false;
let mut sentence_idx: usize = 0;
let mut sentence_t = std::time::Instant::now();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if !first_token_logged {
first_token_logged = true;
tracing::info!(
"conv-phase: ttft={}ms (LLM first token, streaming)",
run_t.elapsed().as_millis()
);
sentence_t = std::time::Instant::now();
}
full.push_str(&chunk);
buf.push_str(&chunk);
while let Some(end) = find_first_boundary(&buf, opts.flush) {
let sentence: String = buf.drain(..end).collect();
let trimmed = sentence.trim();
if trimmed.chars().count() < opts.min_sentence_chars {
continue;
}
let llm_buffer_ms = sentence_t.elapsed().as_millis();
if let Some(hook) = self.pre_sentence_hook.as_mut() {
hook(self.generator, trimmed)?;
}
let utt = self.synthesize_streaming(trimmed, &opts, chunk_frames, &mut on_chunk)?;
tracing::info!(
"conv-phase: sentence[{sentence_idx}] (streaming) \
llm_buffer={llm_buffer_ms}ms tts_gen={}ms chars={}",
utt.tts_latency_ms,
trimmed.chars().count(),
);
sentence_idx += 1;
sentence_t = std::time::Instant::now();
on_sentence(&utt)?;
}
}
let trailing = buf.trim().to_string();
if trailing.chars().count() >= opts.min_sentence_chars {
let llm_buffer_ms = sentence_t.elapsed().as_millis();
if let Some(hook) = self.pre_sentence_hook.as_mut() {
hook(self.generator, &trailing)?;
}
let utt = self.synthesize_streaming(&trailing, &opts, chunk_frames, &mut on_chunk)?;
tracing::info!(
"conv-phase: sentence[{sentence_idx}] (streaming, trailing) \
llm_buffer={llm_buffer_ms}ms tts_gen={}ms chars={}",
utt.tts_latency_ms,
trailing.chars().count(),
);
on_sentence(&utt)?;
}
Ok(full)
}
fn synthesize_streaming<C>(
&mut self,
sentence: &str,
opts: &ConverseOptions,
chunk_frames: usize,
on_chunk: &mut C,
) -> Result<Utterance>
where
C: FnMut(&[f32]) -> Result<()>,
{
let t = std::time::Instant::now();
// generate_streaming is synchronous. The chunk callback fires
// per-chunk; chunks are forwarded via whatever channel the
// caller wired up. For chunks to actually flush to the wire as
// they're decoded (rather than batched at sentence boundary),
// the caller must arrange for the *consumer* to live on a
// different task than this `run_streaming` future — e.g., spawn
// it via `tokio::spawn` rather than `tokio::join!`. Without
// that, chunks pile up in the channel until this sync block
// returns control to the runtime.
let pcm = self
.generator
.generate_streaming(
sentence,
opts.speaker,
&self.context,
opts.generate.clone(),
chunk_frames,
|chunk| on_chunk(chunk),
)
.map_err(|e| CsmError::Config(format!("converse generate_streaming: {e}")))?;
Ok(Utterance {
text: sentence.to_string(),
audio: pcm,
tts_latency_ms: t.elapsed().as_millis(),
})
}
}
/// Convenience: write all utterances concatenated into a single 24 kHz WAV.
pub fn write_concatenated_wav<I>(utterances: I, out: &Path) -> Result<()>
where
I: IntoIterator<Item = Utterance>,
{
let mut all: Vec<f32> = Vec::new();
for u in utterances {
all.extend(u.audio);
}
audio_io::write_wav_24k_mono(out, &all)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn flush_punctuation_modes() {
assert!(should_flush("Hello.", FlushPolicy::Punctuation));
assert!(should_flush("Hello. World", FlushPolicy::Punctuation));
assert!(should_flush("Wait!", FlushPolicy::Punctuation));
assert!(should_flush("What?", FlushPolicy::Punctuation));
assert!(should_flush("Line\n", FlushPolicy::Punctuation));
assert!(!should_flush("Hello", FlushPolicy::Punctuation));
assert!(!should_flush("Hello, world", FlushPolicy::Punctuation));
assert!(should_flush("Hello,", FlushPolicy::Eager));
assert!(should_flush("Hello, world", FlushPolicy::Eager));
assert!(should_flush("then;", FlushPolicy::Eager));
assert!(!should_flush("Hello world", FlushPolicy::Eager));
}
#[test]
fn flush_empty_buffer() {
assert!(!should_flush("", FlushPolicy::Punctuation));
assert!(!should_flush("", FlushPolicy::Eager));
}
#[test]
fn find_boundary_returns_index_after_punctuation() {
// "Hello. World" → boundary at index 6 (after the .).
assert_eq!(
find_first_boundary("Hello. World", FlushPolicy::Punctuation),
Some(6)
);
assert_eq!(find_first_boundary("Hello", FlushPolicy::Punctuation), None);
// Eager: comma at index 5, returns 6.
assert_eq!(
find_first_boundary("Hello, world", FlushPolicy::Eager),
Some(6)
);
}
}