use super::measures::SampleRate;
pub const DEFAULT_AHC_THRESHOLD: f32 = 0.45;
#[derive(Debug, Clone, Copy)]
pub struct ClusterConfig {
pub threshold: f32,
pub max_speakers: usize,
pub min_cluster_size: usize,
pub min_cluster_secs: f64,
}
impl Default for ClusterConfig {
fn default() -> Self {
Self {
threshold: DEFAULT_AHC_THRESHOLD,
max_speakers: 64,
min_cluster_size: 2,
min_cluster_secs: 0.0,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct WindowConfig {
pub window_secs: f32,
pub hop_secs: f32,
pub sample_rate: SampleRate,
}
impl Default for WindowConfig {
fn default() -> Self {
Self {
window_secs: 1.5,
hop_secs: 0.75,
sample_rate: SampleRate::default(),
}
}
}
impl WindowConfig {
pub fn window_samples(&self) -> usize {
(self.window_secs * self.sample_rate.get() as f32) as usize
}
pub fn hop_samples(&self) -> usize {
(self.hop_secs * self.sample_rate.get() as f32) as usize
}
}
#[derive(Debug, Clone, Copy)]
pub struct SpeechFilterConfig {
pub min_speech_secs: f32,
pub max_gap_secs: f32,
}
impl Default for SpeechFilterConfig {
fn default() -> Self {
Self {
min_speech_secs: 0.25,
max_gap_secs: 0.5,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct DiarizationConfig {
pub cluster: ClusterConfig,
pub window: WindowConfig,
pub speech_filter: SpeechFilterConfig,
pub max_duration_secs: f32,
}
impl Default for DiarizationConfig {
fn default() -> Self {
Self {
cluster: ClusterConfig::default(),
window: WindowConfig::default(),
speech_filter: SpeechFilterConfig::default(),
max_duration_secs: 3600.0,
}
}
}
impl DiarizationConfig {
pub fn window_samples(&self) -> usize {
self.window.window_samples()
}
pub fn hop_samples(&self) -> usize {
self.window.hop_samples()
}
pub fn validate(&self) -> Result<(), ConfigError> {
let window_secs = self.window.window_secs;
if !(window_secs.is_finite() && window_secs > 0.0) {
return Err(ConfigError::InvalidWindowSecs(window_secs));
}
let hop_secs = self.window.hop_secs;
if !(hop_secs.is_finite() && hop_secs > 0.0) {
return Err(ConfigError::InvalidHopSecs(hop_secs));
}
if hop_secs > window_secs {
return Err(ConfigError::HopExceedsWindow {
hop_secs,
window_secs,
});
}
if self.window_samples() == 0 || self.hop_samples() == 0 {
return Err(ConfigError::SubSampleWindow {
window_secs,
hop_secs,
});
}
let threshold = self.cluster.threshold;
if !(-1.0..=1.0).contains(&threshold) {
return Err(ConfigError::InvalidThreshold(threshold));
}
let min_speech_secs = self.speech_filter.min_speech_secs;
if !(min_speech_secs.is_finite() && min_speech_secs >= 0.0) {
return Err(ConfigError::InvalidMinSpeechSecs(min_speech_secs));
}
let max_gap_secs = self.speech_filter.max_gap_secs;
if !(max_gap_secs.is_finite() && max_gap_secs >= 0.0) {
return Err(ConfigError::InvalidMaxGapSecs(max_gap_secs));
}
let max_duration_secs = self.max_duration_secs;
if !(max_duration_secs.is_finite() && max_duration_secs > 0.0) {
return Err(ConfigError::InvalidMaxDurationSecs(max_duration_secs));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum ConfigError {
#[error("window.window_secs must be finite and > 0, got {0}")]
InvalidWindowSecs(f32),
#[error("window.hop_secs must be finite and > 0, got {0}")]
InvalidHopSecs(f32),
#[error("window.hop_secs ({hop_secs}) must be <= window.window_secs ({window_secs})")]
HopExceedsWindow { hop_secs: f32, window_secs: f32 },
#[error(
"window geometry too small: window_secs={window_secs}, hop_secs={hop_secs} quantize to zero samples"
)]
SubSampleWindow { window_secs: f32, hop_secs: f32 },
#[error("cluster.threshold must be in [-1.0, 1.0], got {0}")]
InvalidThreshold(f32),
#[error("speech_filter.min_speech_secs must be finite and >= 0, got {0}")]
InvalidMinSpeechSecs(f32),
#[error("speech_filter.max_gap_secs must be finite and >= 0, got {0}")]
InvalidMaxGapSecs(f32),
#[error("max_duration_secs must be finite and > 0, got {0}")]
InvalidMaxDurationSecs(f32),
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_config_is_valid() {
DiarizationConfig::default().validate().unwrap();
}
#[test]
fn non_positive_or_non_finite_window_secs_rejected() {
for bad in [0.0f32, -1.0, f32::NAN, f32::INFINITY] {
let config = DiarizationConfig {
window: WindowConfig {
window_secs: bad,
..WindowConfig::default()
},
..DiarizationConfig::default()
};
assert!(
matches!(config.validate(), Err(ConfigError::InvalidWindowSecs(_))),
"window_secs={bad}"
);
}
}
#[test]
fn non_positive_or_non_finite_hop_secs_rejected() {
for bad in [0.0f32, -0.5, f32::NAN] {
let config = DiarizationConfig {
window: WindowConfig {
hop_secs: bad,
..WindowConfig::default()
},
..DiarizationConfig::default()
};
assert!(
matches!(config.validate(), Err(ConfigError::InvalidHopSecs(_))),
"hop_secs={bad}"
);
}
}
#[test]
fn hop_greater_than_window_rejected() {
let config = DiarizationConfig {
window: WindowConfig {
window_secs: 1.0,
hop_secs: 1.5,
..WindowConfig::default()
},
..DiarizationConfig::default()
};
assert!(matches!(
config.validate(),
Err(ConfigError::HopExceedsWindow {
hop_secs: 1.5,
window_secs: 1.0,
})
));
}
#[test]
fn sub_sample_window_geometry_rejected() {
let config = DiarizationConfig {
window: WindowConfig {
window_secs: 1e-9,
hop_secs: 1e-9,
..WindowConfig::default()
},
..DiarizationConfig::default()
};
assert!(matches!(
config.validate(),
Err(ConfigError::SubSampleWindow { .. })
));
}
#[test]
fn out_of_range_threshold_rejected() {
for bad in [-1.5f32, 1.5, f32::NAN] {
let config = DiarizationConfig {
cluster: ClusterConfig {
threshold: bad,
..ClusterConfig::default()
},
..DiarizationConfig::default()
};
assert!(
matches!(config.validate(), Err(ConfigError::InvalidThreshold(_))),
"threshold={bad}"
);
}
for ok in [-1.0f32, 1.0] {
let config = DiarizationConfig {
cluster: ClusterConfig {
threshold: ok,
..ClusterConfig::default()
},
..DiarizationConfig::default()
};
assert!(config.validate().is_ok(), "threshold={ok}");
}
}
#[test]
fn negative_speech_filter_values_rejected() {
let config = DiarizationConfig {
speech_filter: SpeechFilterConfig {
min_speech_secs: -0.1,
..SpeechFilterConfig::default()
},
..DiarizationConfig::default()
};
assert!(matches!(
config.validate(),
Err(ConfigError::InvalidMinSpeechSecs(_))
));
let config = DiarizationConfig {
speech_filter: SpeechFilterConfig {
max_gap_secs: f32::NAN,
..SpeechFilterConfig::default()
},
..DiarizationConfig::default()
};
assert!(matches!(
config.validate(),
Err(ConfigError::InvalidMaxGapSecs(_))
));
}
#[test]
fn non_positive_max_duration_rejected() {
for bad in [0.0f32, -10.0, f32::NAN] {
let config = DiarizationConfig {
max_duration_secs: bad,
..DiarizationConfig::default()
};
assert!(
matches!(
config.validate(),
Err(ConfigError::InvalidMaxDurationSecs(_))
),
"max_duration_secs={bad}"
);
}
}
}