Skip to main content

native_whisperx/config/
diarization.rs

1//! Diarization configuration for native and delegated speaker assignment.
2
3use std::path::PathBuf;
4
5use serde::{Deserialize, Serialize};
6
7use crate::speaker_directory::SpeakerDirectorySelection;
8
9use super::defaults::default_true;
10use super::ConfigSelection;
11
12#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
13#[serde(rename_all = "camelCase")]
14pub struct DiarizationConfig {
15    #[serde(default)]
16    pub enabled: bool,
17    #[serde(default = "default_diarization_model_id")]
18    pub model_id: String,
19    #[serde(default, skip_serializing_if = "ConfigSelection::is_explicit")]
20    pub model_selection: ConfigSelection,
21    #[serde(default)]
22    pub hf_token: Option<String>,
23    #[serde(default)]
24    pub hf_token_env: Option<String>,
25    #[serde(default)]
26    pub return_speaker_embeddings: bool,
27    #[serde(default)]
28    pub model_bundle: Option<PathBuf>,
29    #[serde(default)]
30    pub manifest_file: Option<String>,
31    #[serde(default)]
32    pub segmentation_model_file: Option<String>,
33    #[serde(default)]
34    pub embedding_model_file: Option<String>,
35    #[serde(default)]
36    pub plda_transform_file: Option<String>,
37    #[serde(default)]
38    pub plda_model_file: Option<String>,
39    #[serde(default)]
40    pub clustering_config_file: Option<String>,
41    #[serde(default)]
42    pub speaker_embedding_model_bundle: Option<PathBuf>,
43    #[serde(default)]
44    pub speaker_embedding_model_file: Option<String>,
45    #[serde(default)]
46    pub speaker_embedding_dimension: Option<usize>,
47    #[serde(default)]
48    pub speaker_embedding_sample_rate: Option<u32>,
49    #[serde(default)]
50    pub min_speakers: Option<usize>,
51    #[serde(default)]
52    pub max_speakers: Option<usize>,
53    #[serde(default)]
54    pub assignment_policy: AssignmentPolicy,
55    #[serde(default)]
56    pub speaker_directory: SpeakerDirectorySelection,
57    #[serde(default)]
58    pub disable_speaker_library: bool,
59    #[serde(default = "default_true")]
60    pub save_draft_speakers: bool,
61    #[serde(default = "default_true")]
62    pub use_draft_speakers: bool,
63}
64
65impl Default for DiarizationConfig {
66    fn default() -> Self {
67        Self {
68            enabled: false,
69            model_id: default_diarization_model_id(),
70            model_selection: ConfigSelection::Explicit,
71            hf_token: None,
72            hf_token_env: None,
73            return_speaker_embeddings: false,
74            model_bundle: None,
75            manifest_file: None,
76            segmentation_model_file: None,
77            embedding_model_file: None,
78            plda_transform_file: None,
79            plda_model_file: None,
80            clustering_config_file: None,
81            speaker_embedding_model_bundle: None,
82            speaker_embedding_model_file: None,
83            speaker_embedding_dimension: None,
84            speaker_embedding_sample_rate: None,
85            min_speakers: None,
86            max_speakers: None,
87            assignment_policy: AssignmentPolicy::Majority,
88            speaker_directory: SpeakerDirectorySelection::default(),
89            disable_speaker_library: false,
90            save_draft_speakers: true,
91            use_draft_speakers: true,
92        }
93    }
94}
95
96#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
97#[serde(rename_all = "camelCase")]
98pub enum AssignmentPolicy {
99    #[default]
100    Majority,
101    NearestStart,
102    StrictContained,
103}
104
105fn default_diarization_model_id() -> String {
106    "native-spectral-speaker-baseline".to_string()
107}
108
109pub(crate) fn is_pyannote_diarization_model(model_id: &str) -> bool {
110    model_id
111        .trim()
112        .to_ascii_lowercase()
113        .starts_with("pyannote/")
114}
115
116#[cfg(test)]
117mod tests {
118    use super::*;
119    use crate::ConfigSelection;
120
121    #[test]
122    fn diarization_config_serializes_automatic_model_selection_separately_from_explicit_model() {
123        let automatic = DiarizationConfig {
124            enabled: true,
125            model_selection: ConfigSelection::Automatic,
126            model_id: "pyannote/speaker-diarization-community-1".to_string(),
127            ..DiarizationConfig::default()
128        };
129
130        let json = serde_json::to_value(&automatic).expect("serialize diarization config");
131
132        assert_eq!(json["modelSelection"], "automatic");
133        assert_eq!(json["modelId"], "pyannote/speaker-diarization-community-1");
134
135        let explicit = DiarizationConfig {
136            enabled: true,
137            model_id: "pyannote/speaker-diarization-community-1".to_string(),
138            model_bundle: Some(PathBuf::from("/models/pyannote-diarization")),
139            ..DiarizationConfig::default()
140        };
141        let json = serde_json::to_value(&explicit).expect("serialize explicit diarization config");
142
143        assert!(json.get("modelSelection").is_none());
144        assert_eq!(json["modelId"], "pyannote/speaker-diarization-community-1");
145        assert_eq!(json["modelBundle"], "/models/pyannote-diarization");
146
147        let decoded: DiarizationConfig = serde_json::from_value(serde_json::json!({
148            "enabled": true,
149            "modelId": "pyannote/speaker-diarization-community-1"
150        }))
151        .expect("deserialize existing diarization config shape");
152        assert_eq!(decoded.model_selection, ConfigSelection::Explicit);
153        assert_eq!(decoded.model_id, "pyannote/speaker-diarization-community-1");
154
155        let decoded: DiarizationConfig = serde_json::from_value(serde_json::json!({
156            "enabled": true,
157            "modelSelection": "automatic",
158            "modelId": "pyannote/speaker-diarization-community-1"
159        }))
160        .expect("deserialize automatic diarization config");
161        assert_eq!(decoded.model_selection, ConfigSelection::Automatic);
162        assert_eq!(decoded.model_id, "pyannote/speaker-diarization-community-1");
163    }
164}