Skip to main content

native_whisperx/config/
asr.rs

1//! ASR provider, decode, and external WhisperX compatibility configuration.
2
3use std::path::PathBuf;
4
5use serde::{Deserialize, Serialize};
6
7use super::defaults::{
8    default_batch_chunks, default_external_whisperx_model, default_max_batch_size,
9    default_whisper_model_id, default_whisperx_command,
10};
11
12#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
13#[serde(rename_all = "camelCase")]
14pub struct AsrConfig {
15    #[serde(default)]
16    pub provider: AsrProvider,
17    #[serde(default)]
18    pub task: TranscriptionTask,
19    #[serde(default = "default_whisper_model_id")]
20    pub model_id: String,
21    #[serde(default)]
22    pub language: Option<String>,
23    #[serde(default)]
24    pub whisper_bundle: Option<PathBuf>,
25    #[serde(default)]
26    pub model_dir: Option<PathBuf>,
27    #[serde(default)]
28    pub model_cache_only: bool,
29    #[serde(default)]
30    pub device: DevicePreference,
31    #[serde(default)]
32    pub device_index: Option<String>,
33    #[serde(default)]
34    pub compute_type: Option<String>,
35    #[serde(default = "default_batch_chunks")]
36    pub batch_chunks: bool,
37    #[serde(default = "default_max_batch_size")]
38    pub max_batch_size: Option<usize>,
39    #[serde(default)]
40    pub decode: WhisperxDecodeConfig,
41    #[serde(default)]
42    pub external_whisperx: ExternalWhisperxConfig,
43}
44
45impl Default for AsrConfig {
46    fn default() -> Self {
47        Self {
48            provider: AsrProvider::Native,
49            task: TranscriptionTask::Transcribe,
50            model_id: default_whisper_model_id(),
51            language: None,
52            whisper_bundle: None,
53            model_dir: None,
54            model_cache_only: false,
55            device: DevicePreference::Auto,
56            device_index: None,
57            compute_type: None,
58            batch_chunks: true,
59            max_batch_size: Some(4),
60            decode: WhisperxDecodeConfig::default(),
61            external_whisperx: ExternalWhisperxConfig::default(),
62        }
63    }
64}
65
66#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
67#[serde(rename_all = "camelCase")]
68pub enum AsrProvider {
69    #[default]
70    Native,
71    ExternalWhisperX,
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
75#[serde(rename_all = "lowercase")]
76pub enum TranscriptionTask {
77    #[default]
78    Transcribe,
79    Translate,
80}
81
82impl TranscriptionTask {
83    pub fn as_whisperx_arg(self) -> &'static str {
84        match self {
85            Self::Transcribe => "transcribe",
86            Self::Translate => "translate",
87        }
88    }
89}
90
91#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
92#[serde(rename_all = "lowercase")]
93pub enum DevicePreference {
94    #[default]
95    Auto,
96    Cpu,
97    Cuda,
98}
99
100#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
101#[serde(rename_all = "camelCase")]
102pub struct WhisperxDecodeConfig {
103    #[serde(default)]
104    pub temperature: Vec<f32>,
105    #[serde(default)]
106    pub best_of: Option<usize>,
107    #[serde(default)]
108    pub beam_size: Option<usize>,
109    #[serde(default)]
110    pub patience: Option<f32>,
111    #[serde(default)]
112    pub length_penalty: Option<f32>,
113    #[serde(default)]
114    pub suppress_tokens: Option<String>,
115    #[serde(default)]
116    pub suppress_numerals: bool,
117    #[serde(default)]
118    pub initial_prompt: Option<String>,
119    #[serde(default)]
120    pub hotwords: Option<String>,
121    #[serde(default)]
122    pub condition_on_previous_text: Option<bool>,
123    #[serde(default)]
124    pub fp16: Option<bool>,
125    #[serde(default)]
126    pub compression_ratio_threshold: Option<f32>,
127    #[serde(default)]
128    pub logprob_threshold: Option<f32>,
129    #[serde(default)]
130    pub no_speech_threshold: Option<f32>,
131    #[serde(default)]
132    pub threads: Option<usize>,
133}
134
135#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
136#[serde(rename_all = "camelCase")]
137pub struct ExternalWhisperxConfig {
138    #[serde(default = "default_whisperx_command")]
139    pub command: PathBuf,
140    #[serde(default = "default_external_whisperx_model")]
141    pub model: String,
142    #[serde(default)]
143    pub compute_type: Option<String>,
144    #[serde(default)]
145    pub batch_size: Option<usize>,
146    #[serde(default)]
147    pub align_model: Option<String>,
148    #[serde(default)]
149    pub diarize: bool,
150    #[serde(default)]
151    pub min_speakers: Option<usize>,
152    #[serde(default)]
153    pub max_speakers: Option<usize>,
154    #[serde(default)]
155    pub hf_token_env: Option<String>,
156    #[serde(default)]
157    pub output_dir: Option<PathBuf>,
158    #[serde(default)]
159    pub timeout_seconds: Option<u64>,
160    #[serde(default)]
161    pub extra_args: Vec<String>,
162}
163
164impl Default for ExternalWhisperxConfig {
165    fn default() -> Self {
166        Self {
167            command: default_whisperx_command(),
168            model: default_external_whisperx_model(),
169            compute_type: None,
170            batch_size: None,
171            align_model: None,
172            diarize: false,
173            min_speakers: None,
174            max_speakers: None,
175            hf_token_env: None,
176            output_dir: None,
177            timeout_seconds: None,
178            extra_args: Vec::new(),
179        }
180    }
181}