diff options
Diffstat (limited to 'crates/asist-asr/src/lib.rs')
| -rw-r--r-- | crates/asist-asr/src/lib.rs | 401 |
1 files changed, 401 insertions, 0 deletions
diff --git a/crates/asist-asr/src/lib.rs b/crates/asist-asr/src/lib.rs new file mode 100644 index 0000000..8212ceb --- /dev/null +++ b/crates/asist-asr/src/lib.rs @@ -0,0 +1,401 @@ +//! Reconocimiento de voz sobre Canary (ONNX). +//! +//! El hecho que da forma a este crate: **decodificar una ventana cuesta más +//! que grabarla**. En esta máquina, 6 s de ventana tardan cerca de 800 ms en +//! decodificarse y sólo avanzan 400 ms de audio. Cualquier diseño que procese +//! todas las ventanas en orden se va quedando atrás del hablante sin límite. +//! +//! La salida es tirar trabajo: el hilo de decodificación vacía su cola entera, +//! se queda sólo con la ventana más reciente y descarta el resto. La latencia +//! queda acotada por una decodificación en lugar de crecer sin freno, y lo que +//! se pierde son transcripciones provisionales que iban a ser sobreescritas. + +use std::time::{Duration, Instant}; + +use crossbeam_channel::{Receiver, Sender}; + +use asist_core::config::AsrConfig; +use asist_core::error::{Error, Result}; +use asist_core::event::TurnId; + +pub use canary_rs::{Canary, CanarySession, ExecutionConfig, ExecutionProvider, StreamConfig}; + +/// Trabajo que llega al decodificador. +#[derive(Debug)] +pub enum AsrJob { + /// Audio nuevo para la ventana deslizante. + Window { + turn: TurnId, + samples: Vec<f32>, + at: Instant, + }, + /// Intervención cerrada, para transcribir entera. + Utterance { + turn: TurnId, + samples: Vec<f32>, + at: Instant, + }, + /// La intervención no tenía nada: reinicia el estado de la ventana. + Reset { turn: TurnId }, +} + +/// Lo que el decodificador devuelve. +#[derive(Debug, Clone)] +pub enum AsrResult { + Partial { + turn: TurnId, + committed: String, + volatile: String, + /// Ventanas descartadas por obsoletas antes de esta. + dropped: usize, + decode: Duration, + }, + Final { + turn: TurnId, + text: String, + audio_secs: f32, + decode: Duration, + /// Instante en que el usuario dejó de hablar. Es el origen del que se + /// mide la latencia percibida, y no puede tomarse aquí: para cuando la + /// transcripción está lista ya ha pasado casi un segundo. + spoken_at: Instant, + }, + Empty { + turn: TurnId, + }, + Error { + turn: TurnId, + message: String, + }, +} + +/// Motor cargado y listo para decodificar. +pub struct Recognizer { + model: Canary, + config: AsrConfig, +} + +impl Recognizer { + /// Carga el modelo desde `config.model_dir`. + pub fn load(config: &AsrConfig) -> Result<Self> { + if !config.model_dir.is_dir() { + return Err(Error::Asr(format!( + "no existe la carpeta del modelo: {}. Ejecuta scripts/bootstrap.sh", + config.model_dir.display() + ))); + } + let started = Instant::now(); + let model = Canary::from_pretrained( + config.model_dir.to_string_lossy().as_ref(), + Some(execution_config(config)), + ) + .map_err(|e| Error::Asr(format!("no se pudo cargar Canary: {e}")))?; + + tracing::info!( + target: "asr", + carpeta = %config.model_dir.display(), + proveedor = %config.execution_provider, + ms = started.elapsed().as_millis(), + "modelo cargado" + ); + Ok(Self { + model, + config: config.clone(), + }) + } + + /// Sesión suelta para transcribir de una vez (pruebas y comprobaciones). + pub fn transcribe(&self, samples: &[f32], sample_rate: u32) -> Result<String> { + let mut session = self.model.session(); + session + .transcribe_samples( + samples, + sample_rate as usize, + 1, + &self.config.source_lang, + &self.config.target_lang, + ) + .map(|r| normalize(&r.text)) + .map_err(|e| Error::Asr(e.to_string())) + } + + /// Bucle del decodificador. Se ejecuta en su propio hilo hasta que se + /// cierre `jobs`. + pub fn run(self, jobs: Receiver<AsrJob>, out: Sender<AsrResult>) { + let mut stream = match self.model.stream( + self.config.source_lang.clone(), + self.config.target_lang.clone(), + stream_config(&self.config), + ) { + Ok(stream) => stream, + Err(err) => { + let _ = out.send(AsrResult::Error { + turn: TurnId::default(), + message: format!("no se pudo abrir el flujo: {err}"), + }); + return; + } + }; + let mut session = self.model.session(); + let mut committed = String::new(); + let step_samples = (self.config.step * 16_000.0).max(1.0) as usize; + + while let Some(batch) = drain(&jobs) { + // Todo el audio que se quedó detrás de un cierre de intervención + // pertenece a un turno que ya se va a transcribir entero: gastar + // una ventana en él sería trabajo condenado a sobreescribirse. + let mut pending: Vec<f32> = Vec::new(); + let mut pending_turn = TurnId::default(); + let mut newest = Instant::now(); + let mut dropped = 0usize; + + for job in batch { + match job { + AsrJob::Window { turn, samples, at } => { + pending_turn = turn; + pending.extend_from_slice(&samples); + newest = at; + } + AsrJob::Reset { turn } => { + dropped += pending.len() / step_samples; + pending.clear(); + stream.reset(); + committed.clear(); + if out.send(AsrResult::Empty { turn }).is_err() { + return; + } + } + AsrJob::Utterance { turn, samples, at } => { + dropped += pending.len() / step_samples; + pending.clear(); + stream.reset(); + committed.clear(); + let result = decode_final(&mut session, &self.config, turn, &samples, at); + if out.send(result).is_err() { + return; + } + } + } + } + + if pending.is_empty() || !self.config.partials { + continue; + } + // Sólo sobrevive la ventana más nueva; lo anterior ya no describe + // lo que se está diciendo ahora. + dropped += (pending.len() / step_samples).saturating_sub(1); + + let started = Instant::now(); + let chunks = match stream.push_samples(&pending, 16_000, 1) { + Ok(chunks) => chunks, + Err(err) => { + let _ = out.send(AsrResult::Error { + turn: pending_turn, + message: err.to_string(), + }); + continue; + } + }; + let Some(chunk) = chunks.last() else { continue }; + if is_degenerate(&chunk.result.text) { + continue; + } + + append_delta(&mut committed, chunk.delta_text.trim()); + let volatile = volatile_tail(&committed, chunk.result.text.trim()); + let decode = started.elapsed(); + tracing::debug!( + target: "asr", + turno = pending_turn.0, + ms = decode.as_millis(), + retraso_ms = newest.elapsed().as_millis(), + descartadas = dropped, + "ventana" + ); + if out + .send(AsrResult::Partial { + turn: pending_turn, + committed: committed.clone(), + volatile, + dropped, + decode, + }) + .is_err() + { + return; + } + } + } +} + +fn decode_final( + session: &mut CanarySession, + config: &AsrConfig, + turn: TurnId, + samples: &[f32], + at: Instant, +) -> AsrResult { + let started = Instant::now(); + let audio_secs = samples.len() as f32 / 16_000.0; + match session.transcribe_samples(samples, 16_000, 1, &config.source_lang, &config.target_lang) { + Ok(result) => { + let text = normalize(&result.text); + let decode = started.elapsed(); + tracing::debug!( + target: "asr", + turno = turn.0, + ms = decode.as_millis(), + audio_s = audio_secs, + rtf = decode.as_secs_f32() / audio_secs.max(0.001), + retraso_ms = at.elapsed().as_millis(), + "transcripción final" + ); + if text.is_empty() { + AsrResult::Empty { turn } + } else { + AsrResult::Final { + turn, + text, + audio_secs, + decode, + spoken_at: at, + } + } + } + Err(err) => AsrResult::Error { + turn, + message: err.to_string(), + }, + } +} + +fn execution_config(config: &AsrConfig) -> ExecutionConfig { + let provider = match config.execution_provider.trim().to_lowercase().as_str() { + "cuda" => ExecutionProvider::Cuda, + "tensorrt" => ExecutionProvider::TensorRT, + "rocm" => ExecutionProvider::ROCm, + "openvino" => ExecutionProvider::OpenVINO, + "webgpu" => ExecutionProvider::WebGPU, + "coreml" => ExecutionProvider::CoreML, + "directml" => ExecutionProvider::DirectML, + _ => ExecutionProvider::Cpu, + }; + ExecutionConfig::new() + .with_execution_provider(provider) + .with_threads(config.inter_threads, config.intra_threads) +} + +fn stream_config(config: &AsrConfig) -> StreamConfig { + StreamConfig::new() + .with_window_duration(config.window) + .with_step_duration(config.step) + .with_emit_partial(true) + .with_pad_partial(false) + .with_stability_window(config.stability) + // El motivo de todo el diseño: nunca arrastrar una cola de ventanas viejas. + .with_max_windows_per_push(1) +} + +/// Recoge de golpe todo lo encolado, bloqueando hasta que llegue lo primero. +fn drain(jobs: &Receiver<AsrJob>) -> Option<Vec<AsrJob>> { + let mut batch = vec![jobs.recv().ok()?]; + while let Ok(job) = jobs.try_recv() { + batch.push(job); + } + Some(batch) +} + +/// Une el texto estable con el fragmento nuevo respetando los espacios. +pub fn append_delta(committed: &mut String, delta: &str) { + if delta.is_empty() { + return; + } + let needs_space = !committed.is_empty() + && !committed.ends_with(' ') + && !delta.starts_with(|c: char| c.is_ascii_punctuation()); + if needs_space { + committed.push(' '); + } + committed.push_str(delta); +} + +/// Parte de la ventana actual que aún no es estable. +pub fn volatile_tail(committed: &str, window: &str) -> String { + // La ventana repite el final de lo ya fijado; lo interesante es lo que + // sobra por detrás. + match window.rfind(last_words(committed, 3).as_str()) { + Some(idx) if !committed.is_empty() => window[idx + last_words(committed, 3).len()..] + .trim_start() + .to_string(), + _ => window.to_string(), + } +} + +fn last_words(text: &str, n: usize) -> String { + let words: Vec<&str> = text.split_whitespace().collect(); + words[words.len().saturating_sub(n)..].join(" ") +} + +/// Canary a veces se atasca repitiendo un token cuando la ventana es casi +/// silencio. Mostrarlo confundiría más que ayudar. +pub fn is_degenerate(text: &str) -> bool { + let words: Vec<&str> = text.split_whitespace().collect(); + if words.len() < 6 { + return false; + } + let distinct: std::collections::HashSet<&&str> = words.iter().collect(); + distinct.len() * 4 <= words.len() +} + +/// Arregla el espaciado de la puntuación que a veces deja el decodificador. +pub fn normalize(text: &str) -> String { + let mut out = String::with_capacity(text.len()); + for c in text.trim().chars() { + if matches!(c, ',' | '.' | '!' | '?' | ';' | ':') { + while out.ends_with(' ') { + out.pop(); + } + } + out.push(c); + } + out +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn el_delta_se_une_con_espacio_salvo_ante_puntuacion() { + let mut s = String::from("hola"); + append_delta(&mut s, "mundo"); + assert_eq!(s, "hola mundo"); + append_delta(&mut s, ","); + assert_eq!(s, "hola mundo,"); + append_delta(&mut s, ""); + assert_eq!(s, "hola mundo,"); + } + + #[test] + fn la_cola_volatil_es_lo_que_sobra_tras_lo_fijado() { + assert_eq!(volatile_tail("hola qué tal", "hola qué tal estás"), "estás"); + } + + #[test] + fn sin_texto_fijado_toda_la_ventana_es_volatil() { + assert_eq!(volatile_tail("", "hola qué tal"), "hola qué tal"); + } + + #[test] + fn se_detecta_la_repeticion_degenerada() { + assert!(is_degenerate("sí sí sí sí sí sí sí sí")); + assert!(!is_degenerate("hola qué tal estás hoy amigo")); + assert!(!is_degenerate("sí sí"), "una frase corta no es degenerada"); + } + + #[test] + fn la_puntuacion_pierde_el_espacio_de_delante() { + assert_eq!(normalize("hola , qué tal ?"), "hola, qué tal?"); + assert_eq!(normalize(" ya está "), "ya está"); + } +} |