#[cfg(test)]
mod tests;
pub const DEFAULT_THRESHOLD: f64 = 0.6;
pub const DEFAULT_FA: f64 = 0.07;
pub const DEFAULT_FB: f64 = 0.8;
pub const DEFAULT_MAX_ITERS: usize = 20;
pub const DEFAULT_MIN_DURATION_OFF: f64 = 0.0;
#[cfg(feature = "serde")]
fn default_threshold() -> f64 {
DEFAULT_THRESHOLD
}
#[cfg(feature = "serde")]
fn default_fa() -> f64 {
DEFAULT_FA
}
#[cfg(feature = "serde")]
fn default_fb() -> f64 {
DEFAULT_FB
}
#[cfg(feature = "serde")]
fn default_max_iters() -> usize {
DEFAULT_MAX_ITERS
}
#[cfg(feature = "serde")]
fn default_min_duration_off() -> f64 {
DEFAULT_MIN_DURATION_OFF
}
#[cfg(feature = "serde")]
const NON_FINITE_FLOAT_MSG: &str = "non-finite float (NaN or infinity) is not representable in \
JSON and is rejected to keep the serde round trip lossless";
#[cfg(feature = "serde")]
const NEGATIVE_OR_NON_FINITE_MSG: &str = "min_duration_off must be a finite, non-negative float \
(seconds); NaN, infinity, and negative values are rejected";
#[cfg(feature = "serde")]
pub(crate) mod finite_f64 {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &f64, serializer: S) -> Result<S::Ok, S::Error> {
if !value.is_finite() {
return Err(serde::ser::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
let value = f64::deserialize(deserializer)?;
if !value.is_finite() {
return Err(serde::de::Error::custom(super::NON_FINITE_FLOAT_MSG));
}
Ok(value)
}
}
#[cfg(feature = "serde")]
pub(crate) mod finite_nonneg_f64 {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &f64, serializer: S) -> Result<S::Ok, S::Error> {
if !super::check_min_duration_off(*value) {
return Err(serde::ser::Error::custom(super::NEGATIVE_OR_NON_FINITE_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
let value = f64::deserialize(deserializer)?;
if !super::check_min_duration_off(value) {
return Err(serde::de::Error::custom(super::NEGATIVE_OR_NON_FINITE_MSG));
}
Ok(value)
}
}
#[inline]
const fn check_min_duration_off(v: f64) -> bool {
#[allow(clippy::eq_op)] let not_nan = !(v != v);
not_nan && v >= 0.0 && v != f64::INFINITY
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct OfflineOptions {
#[cfg_attr(
feature = "serde",
serde(default = "default_threshold", with = "finite_f64")
)]
threshold: f64,
#[cfg_attr(feature = "serde", serde(default = "default_fa", with = "finite_f64"))]
fa: f64,
#[cfg_attr(feature = "serde", serde(default = "default_fb", with = "finite_f64"))]
fb: f64,
#[cfg_attr(feature = "serde", serde(default = "default_max_iters"))]
max_iters: usize,
#[cfg_attr(
feature = "serde",
serde(default = "default_min_duration_off", with = "finite_nonneg_f64")
)]
min_duration_off: f64,
}
impl Default for OfflineOptions {
fn default() -> Self {
Self::new()
}
}
impl OfflineOptions {
pub const fn new() -> Self {
Self {
threshold: DEFAULT_THRESHOLD,
fa: DEFAULT_FA,
fb: DEFAULT_FB,
max_iters: DEFAULT_MAX_ITERS,
min_duration_off: DEFAULT_MIN_DURATION_OFF,
}
}
#[inline(always)]
pub const fn threshold(&self) -> f64 {
self.threshold
}
#[inline(always)]
pub const fn fa(&self) -> f64 {
self.fa
}
#[inline(always)]
pub const fn fb(&self) -> f64 {
self.fb
}
#[inline(always)]
pub const fn max_iters(&self) -> usize {
self.max_iters
}
#[inline(always)]
pub const fn min_duration_off(&self) -> f64 {
self.min_duration_off
}
#[must_use]
#[inline(always)]
pub const fn with_threshold(mut self, threshold: f64) -> Self {
self.set_threshold(threshold);
self
}
#[inline(always)]
pub const fn set_threshold(&mut self, threshold: f64) -> &mut Self {
self.threshold = threshold;
self
}
#[must_use]
#[inline(always)]
pub const fn with_fa(mut self, fa: f64) -> Self {
self.set_fa(fa);
self
}
#[inline(always)]
pub const fn set_fa(&mut self, fa: f64) -> &mut Self {
self.fa = fa;
self
}
#[must_use]
#[inline(always)]
pub const fn with_fb(mut self, fb: f64) -> Self {
self.set_fb(fb);
self
}
#[inline(always)]
pub const fn set_fb(&mut self, fb: f64) -> &mut Self {
self.fb = fb;
self
}
#[must_use]
#[inline(always)]
pub const fn with_max_iters(mut self, max_iters: usize) -> Self {
self.set_max_iters(max_iters);
self
}
#[inline(always)]
pub const fn set_max_iters(&mut self, max_iters: usize) -> &mut Self {
self.max_iters = max_iters;
self
}
#[must_use]
#[inline(always)]
pub const fn with_min_duration_off(mut self, min_duration_off: f64) -> Self {
self.set_min_duration_off(min_duration_off);
self
}
#[inline(always)]
pub const fn set_min_duration_off(&mut self, min_duration_off: f64) -> &mut Self {
assert!(
check_min_duration_off(min_duration_off),
"min_duration_off must be finite and >= 0"
);
self.min_duration_off = min_duration_off;
self
}
#[must_use]
pub(crate) fn apply_to<'a>(
&self,
input: diaric::offline::OfflineInput<'a>,
) -> diaric::offline::OfflineInput<'a> {
input
.with_threshold(self.threshold)
.with_fa(self.fa)
.with_fb(self.fb)
.with_max_iters(self.max_iters)
.with_min_duration_off(self.min_duration_off)
}
}
pub const DEFAULT_SPEAKER_THRESHOLD: f32 = 0.65;
pub const DEFAULT_EMBEDDING_THRESHOLD: f32 = 0.45;
pub const DEFAULT_MIN_SPEECH_DURATION: f32 = 1.0;
#[cfg(feature = "serde")]
fn default_speaker_threshold() -> f32 {
DEFAULT_SPEAKER_THRESHOLD
}
#[cfg(feature = "serde")]
fn default_embedding_threshold() -> f32 {
DEFAULT_EMBEDDING_THRESHOLD
}
#[cfg(feature = "serde")]
fn default_min_speech_duration() -> f32 {
DEFAULT_MIN_SPEECH_DURATION
}
#[inline]
#[allow(clippy::manual_range_contains)] const fn check_online_threshold(v: f32) -> bool {
v >= 0.0 && v <= 2.0
}
#[inline]
const fn check_min_speech_duration(v: f32) -> bool {
#[allow(clippy::eq_op)] let not_nan = !(v != v);
not_nan && v >= 0.0 && v != f32::INFINITY
}
#[cfg(feature = "serde")]
const ONLINE_THRESHOLD_MSG: &str = "online cluster threshold must be a finite cosine distance in \
[0.0, 2.0]; NaN, infinity, and out-of-range values are rejected";
#[cfg(feature = "serde")]
const ONLINE_DURATION_MSG: &str = "min_speech_duration must be a finite, non-negative float \
(seconds); NaN, infinity, and negative values are rejected";
#[cfg(feature = "serde")]
pub(crate) mod finite_threshold_f32 {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &f32, serializer: S) -> Result<S::Ok, S::Error> {
if !super::check_online_threshold(*value) {
return Err(serde::ser::Error::custom(super::ONLINE_THRESHOLD_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f32, D::Error> {
let value = f32::deserialize(deserializer)?;
if !super::check_online_threshold(value) {
return Err(serde::de::Error::custom(super::ONLINE_THRESHOLD_MSG));
}
Ok(value)
}
}
#[cfg(feature = "serde")]
pub(crate) mod finite_nonneg_f32 {
use serde::{Deserialize, Deserializer, Serialize, Serializer};
pub(crate) fn serialize<S: Serializer>(value: &f32, serializer: S) -> Result<S::Ok, S::Error> {
if !super::check_min_speech_duration(*value) {
return Err(serde::ser::Error::custom(super::ONLINE_DURATION_MSG));
}
value.serialize(serializer)
}
pub(crate) fn deserialize<'de, D: Deserializer<'de>>(deserializer: D) -> Result<f32, D::Error> {
let value = f32::deserialize(deserializer)?;
if !super::check_min_speech_duration(value) {
return Err(serde::de::Error::custom(super::ONLINE_DURATION_MSG));
}
Ok(value)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct OnlineOptions {
#[cfg_attr(
feature = "serde",
serde(default = "default_speaker_threshold", with = "finite_threshold_f32")
)]
speaker_threshold: f32,
#[cfg_attr(
feature = "serde",
serde(default = "default_embedding_threshold", with = "finite_threshold_f32")
)]
embedding_threshold: f32,
#[cfg_attr(
feature = "serde",
serde(default = "default_min_speech_duration", with = "finite_nonneg_f32")
)]
min_speech_duration: f32,
}
impl Default for OnlineOptions {
fn default() -> Self {
Self::new()
}
}
impl OnlineOptions {
pub const fn new() -> Self {
Self {
speaker_threshold: DEFAULT_SPEAKER_THRESHOLD,
embedding_threshold: DEFAULT_EMBEDDING_THRESHOLD,
min_speech_duration: DEFAULT_MIN_SPEECH_DURATION,
}
}
#[must_use]
pub const fn from_clustering_threshold(base: f32) -> Self {
Self::new()
.with_speaker_threshold(base * 1.2)
.with_embedding_threshold(base * 0.8)
}
#[inline(always)]
pub const fn speaker_threshold(&self) -> f32 {
self.speaker_threshold
}
#[inline(always)]
pub const fn embedding_threshold(&self) -> f32 {
self.embedding_threshold
}
#[inline(always)]
pub const fn min_speech_duration(&self) -> f32 {
self.min_speech_duration
}
#[must_use]
#[inline(always)]
pub const fn with_speaker_threshold(mut self, speaker_threshold: f32) -> Self {
self.set_speaker_threshold(speaker_threshold);
self
}
#[inline(always)]
pub const fn set_speaker_threshold(&mut self, speaker_threshold: f32) -> &mut Self {
assert!(
check_online_threshold(speaker_threshold),
"speaker_threshold must be a finite cosine distance in [0.0, 2.0]"
);
self.speaker_threshold = speaker_threshold;
self
}
#[must_use]
#[inline(always)]
pub const fn with_embedding_threshold(mut self, embedding_threshold: f32) -> Self {
self.set_embedding_threshold(embedding_threshold);
self
}
#[inline(always)]
pub const fn set_embedding_threshold(&mut self, embedding_threshold: f32) -> &mut Self {
assert!(
check_online_threshold(embedding_threshold),
"embedding_threshold must be a finite cosine distance in [0.0, 2.0]"
);
self.embedding_threshold = embedding_threshold;
self
}
#[must_use]
#[inline(always)]
pub const fn with_min_speech_duration(mut self, min_speech_duration: f32) -> Self {
self.set_min_speech_duration(min_speech_duration);
self
}
#[inline(always)]
pub const fn set_min_speech_duration(&mut self, min_speech_duration: f32) -> &mut Self {
assert!(
check_min_speech_duration(min_speech_duration),
"min_speech_duration must be finite and >= 0"
);
self.min_speech_duration = min_speech_duration;
self
}
#[must_use]
pub fn to_dia_options(&self) -> diaric::cluster::online::OnlineClusterOptions {
diaric::cluster::online::OnlineClusterOptions::new()
.with_speaker_threshold(self.speaker_threshold)
.with_embedding_threshold(self.embedding_threshold)
.with_min_speech_duration(self.min_speech_duration)
}
}
macro_rules! define_cluster_backend {
(
$( #[$enum_meta:meta] )*
pub enum ClusterBackend {
$(
$( #[$variant_meta:meta] )*
$variant:ident ( $payload:ty ) => $spelling:literal
),+ $(,)?
}
) => {
$( #[$enum_meta] )*
pub enum ClusterBackend {
$(
$( #[$variant_meta] )*
$variant($payload),
)+
}
impl ClusterBackend {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
$( Self::$variant(..) => $spelling, )+
}
}
}
impl core::str::FromStr for ClusterBackend {
type Err = ParseClusterBackendError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
$( $spelling => Self::$variant(<$payload>::default()), )+
_ => return Err(ParseClusterBackendError(())),
})
}
}
#[cfg(test)]
pub(crate) const CLUSTER_BACKEND_SPELLINGS: &[&str] = &[ $( $spelling ),+ ];
};
}
define_cluster_backend! {
#[derive(Debug, Clone, Copy, PartialEq, derive_more::Display)]
#[display("{}", self.as_str())]
#[non_exhaustive]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum ClusterBackend {
Offline(OfflineOptions) => "offline",
Online(OnlineOptions) => "online",
}
}
impl Default for ClusterBackend {
fn default() -> Self {
Self::Offline(OfflineOptions::new())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown cluster backend name")]
pub struct ParseClusterBackendError(());