aboutsummaryrefslogtreecommitdiffstats
path: root/crates/asist-asr/src
diff options
context:
space:
mode:
Diffstat (limited to 'crates/asist-asr/src')
-rw-r--r--crates/asist-asr/src/lib.rs401
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á");
+ }
+}