native_whisperx/config/
diarization.rs1use 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}