Skip to main content

kcode_speaker_model/
lib.rs

1use kcode_diag_gmm::{
2    DiagGmm, FitConfig as GmmFitConfig, fit as fit_gmm, log_likelihood, means, responsibilities,
3    with_means,
4};
5use kcode_speaker_types::{FeatureMask, FeatureVector, Key, LabeledSample};
6use serde::{Deserialize, Serialize};
7use sha2::{Digest, Sha256};
8use std::fmt;
9
10const SNAPSHOT_VERSION: u8 = 1;
11const MAX_ITERATIONS: u16 = 200;
12const RELATIVE_TOLERANCE: f64 = 1e-8;
13
14pub const GEMINI_11: FeatureMask = match FeatureMask::from_bits(
15    (1 << 0)
16        | (1 << 1)
17        | (1 << 2)
18        | (1 << 4)
19        | (1 << 5)
20        | (1 << 8)
21        | (1 << 11)
22        | (1 << 12)
23        | (1 << 13)
24        | (1 << 17)
25        | (1 << 22),
26) {
27    Ok(mask) => mask,
28    Err(_) => panic!("invalid frozen mask"),
29};
30
31pub const GEMINI_20: FeatureMask =
32    match FeatureMask::from_bits(((1_u64 << 18) - 1) | (1 << 19) | (1 << 22)) {
33        Ok(mask) => mask,
34        Err(_) => panic!("invalid frozen mask"),
35    };
36
37pub const ALL_35: FeatureMask = match FeatureMask::from_bits((1_u64 << 35) - 1) {
38    Ok(mask) => mask,
39    Err(_) => panic!("invalid frozen mask"),
40};
41
42#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
43#[serde(deny_unknown_fields)]
44pub struct ModelConfig {
45    pub mask: FeatureMask,
46    pub components: u8,
47    pub relevance: f64,
48    pub variance_floor: f64,
49    pub absolute_threshold: f64,
50    pub margin_threshold: f64,
51}
52
53pub struct FitInput<'a> {
54    pub cohort_id: &'a Key,
55    pub samples: &'a [LabeledSample],
56    pub config: ModelConfig,
57}
58
59#[derive(Serialize, Deserialize)]
60#[serde(deny_unknown_fields)]
61pub struct ModelSnapshot {
62    version: u8,
63    cohort_id: Key,
64    config: ModelConfig,
65    selected_indices: Vec<u8>,
66    normalizer_mean: Vec<f64>,
67    normalizer_std: Vec<f64>,
68    ubm: DiagGmm,
69    speakers: Vec<SpeakerModel>,
70    training_sample_count: usize,
71    manifest_sha256: [u8; 32],
72}
73
74#[derive(Serialize, Deserialize)]
75#[serde(deny_unknown_fields)]
76struct SpeakerModel {
77    speaker_id: Key,
78    sample_count: usize,
79    occupancy: Vec<f64>,
80    first_moments: Vec<Vec<f64>>,
81    adapted_means: Vec<Vec<f64>>,
82}
83
84#[derive(Clone, Debug, PartialEq)]
85pub struct CandidateScore {
86    pub speaker_id: Key,
87    pub llr: f64,
88}
89
90#[derive(Clone, Debug, PartialEq)]
91pub enum Decision {
92    Known { speaker_id: Key },
93    Unknown,
94}
95
96#[derive(Clone, Debug, PartialEq)]
97pub struct Identification {
98    pub decision: Decision,
99    pub best: CandidateScore,
100    pub runner_up: Option<CandidateScore>,
101    pub absolute_pass: bool,
102    pub margin_pass: bool,
103}
104
105#[derive(Clone, Debug, PartialEq, Eq)]
106pub enum ModelError {
107    InvalidConfig,
108    InvalidSamples,
109    CohortMismatch,
110    DuplicateSampleId,
111    ZeroVariance,
112    Nonconverged,
113    Numerical,
114    Gmm,
115    Serialization,
116    MalformedArtifact,
117}
118
119impl fmt::Display for ModelError {
120    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
121        let text = match self {
122            Self::InvalidConfig => "invalid model configuration",
123            Self::InvalidSamples => "invalid training samples",
124            Self::CohortMismatch => "cohort does not match",
125            Self::DuplicateSampleId => "duplicate training sample ID",
126            Self::ZeroVariance => "selected training feature has zero variance",
127            Self::Nonconverged => "diagonal GMM did not converge",
128            Self::Numerical => "nonfinite model computation",
129            Self::Gmm => "diagonal GMM operation failed",
130            Self::Serialization => "snapshot serialization failed",
131            Self::MalformedArtifact => "malformed or incompatible snapshot artifact",
132        };
133        f.write_str(text)
134    }
135}
136
137impl std::error::Error for ModelError {}
138
139pub fn fit(input: FitInput<'_>) -> Result<ModelSnapshot, ModelError> {
140    validate_config(&input.config, input.samples.len())?;
141    if input.samples.is_empty() {
142        return Err(ModelError::InvalidSamples);
143    }
144
145    let mut samples: Vec<&LabeledSample> = input.samples.iter().collect();
146    samples.sort_by(|a, b| a.sample_id.as_ref().cmp(b.sample_id.as_ref()));
147    for sample in &samples {
148        if sample.cohort_id.as_ref() != input.cohort_id.as_ref() {
149            return Err(ModelError::CohortMismatch);
150        }
151    }
152    if samples
153        .windows(2)
154        .any(|pair| pair[0].sample_id.as_ref() == pair[1].sample_id.as_ref())
155    {
156        return Err(ModelError::DuplicateSampleId);
157    }
158
159    let selected = mask_indices(input.config.mask);
160    let raw_rows: Vec<Vec<f64>> = samples
161        .iter()
162        .map(|sample| {
163            selected
164                .iter()
165                .map(|&index| sample.features.as_ref()[index as usize] as f64)
166                .collect()
167        })
168        .collect();
169    let (normalizer_mean, normalizer_std) = fit_normalizer(&raw_rows)?;
170    let rows: Vec<Vec<f64>> = raw_rows
171        .iter()
172        .map(|row| normalize(row, &normalizer_mean, &normalizer_std))
173        .collect::<Result<_, _>>()?;
174
175    let ubm = fit_ubm_with_limits(
176        &rows,
177        input.config.components,
178        input.config.variance_floor,
179        MAX_ITERATIONS,
180        RELATIVE_TOLERANCE,
181    )?;
182    let ubm_means = means(&ubm);
183
184    let mut speaker_ids: Vec<Key> = samples
185        .iter()
186        .map(|sample| sample.speaker_id.clone())
187        .collect();
188    speaker_ids.sort_by(|a, b| a.as_ref().cmp(b.as_ref()));
189    speaker_ids.dedup_by(|a, b| a.as_ref() == b.as_ref());
190    if speaker_ids.is_empty() {
191        return Err(ModelError::InvalidSamples);
192    }
193
194    let components = input.config.components as usize;
195    let dimension = selected.len();
196    let mut speakers: Vec<SpeakerModel> = speaker_ids
197        .into_iter()
198        .map(|speaker_id| SpeakerModel {
199            speaker_id,
200            sample_count: 0,
201            occupancy: vec![0.0; components],
202            first_moments: vec![vec![0.0; dimension]; components],
203            adapted_means: Vec::new(),
204        })
205        .collect();
206
207    for (sample, row) in samples.iter().zip(&rows) {
208        let speaker_index = speakers
209            .binary_search_by(|speaker| speaker.speaker_id.as_ref().cmp(sample.speaker_id.as_ref()))
210            .map_err(|_| ModelError::InvalidSamples)?;
211        let gamma = responsibilities(&ubm, row).map_err(|_| ModelError::Gmm)?;
212        let speaker = &mut speakers[speaker_index];
213        speaker.sample_count += 1;
214        for (component, responsibility) in gamma[..components].iter().copied().enumerate() {
215            if !responsibility.is_finite() || responsibility < 0.0 {
216                return Err(ModelError::Numerical);
217            }
218            speaker.occupancy[component] += responsibility;
219            for (feature, value) in row.iter().enumerate() {
220                speaker.first_moments[component][feature] += responsibility * value;
221            }
222        }
223    }
224
225    for speaker in &mut speakers {
226        speaker.adapted_means = map_means(
227            ubm_means,
228            &speaker.occupancy,
229            &speaker.first_moments,
230            input.config.relevance,
231        )?;
232    }
233
234    let manifest_sha256 = manifest(&samples);
235    let model = ModelSnapshot {
236        version: SNAPSHOT_VERSION,
237        cohort_id: input.cohort_id.clone(),
238        config: input.config,
239        selected_indices: selected,
240        normalizer_mean,
241        normalizer_std,
242        ubm,
243        speakers,
244        training_sample_count: samples.len(),
245        manifest_sha256,
246    };
247    validate_snapshot(&model)?;
248    Ok(model)
249}
250
251pub fn identify(
252    model: &ModelSnapshot,
253    cohort_id: &Key,
254    features: &FeatureVector,
255) -> Result<Identification, ModelError> {
256    validate_snapshot(model)?;
257    if cohort_id.as_ref() != model.cohort_id.as_ref() {
258        return Err(ModelError::CohortMismatch);
259    }
260
261    let raw: Vec<f64> = model
262        .selected_indices
263        .iter()
264        .map(|&index| features.as_ref()[index as usize] as f64)
265        .collect();
266    let row = normalize(&raw, &model.normalizer_mean, &model.normalizer_std)?;
267    let ubm_score = log_likelihood(&model.ubm, &row).map_err(|_| ModelError::Gmm)?;
268    if !ubm_score.is_finite() {
269        return Err(ModelError::Numerical);
270    }
271
272    let mut scores = Vec::with_capacity(model.speakers.len());
273    for speaker in &model.speakers {
274        let adapted =
275            with_means(&model.ubm, speaker.adapted_means.clone()).map_err(|_| ModelError::Gmm)?;
276        let score = log_likelihood(&adapted, &row).map_err(|_| ModelError::Gmm)? - ubm_score;
277        if !score.is_finite() {
278            return Err(ModelError::Numerical);
279        }
280        scores.push(CandidateScore {
281            speaker_id: speaker.speaker_id.clone(),
282            llr: score,
283        });
284    }
285    scores.sort_by(|a, b| {
286        b.llr
287            .total_cmp(&a.llr)
288            .then_with(|| a.speaker_id.as_ref().cmp(b.speaker_id.as_ref()))
289    });
290
291    let best = scores.remove(0);
292    let runner_up = scores.into_iter().next();
293    let absolute_pass = best.llr >= model.config.absolute_threshold;
294    let margin_pass = match &runner_up {
295        Some(runner) => best.llr - runner.llr >= model.config.margin_threshold,
296        None => true,
297    };
298    let decision = if absolute_pass && margin_pass {
299        Decision::Known {
300            speaker_id: best.speaker_id.clone(),
301        }
302    } else {
303        Decision::Unknown
304    };
305
306    Ok(Identification {
307        decision,
308        best,
309        runner_up,
310        absolute_pass,
311        margin_pass,
312    })
313}
314
315pub fn encode(model: &ModelSnapshot) -> Result<Vec<u8>, ModelError> {
316    validate_snapshot(model)?;
317    serde_json::to_vec(model).map_err(|_| ModelError::Serialization)
318}
319
320pub fn decode(bytes: &[u8]) -> Result<ModelSnapshot, ModelError> {
321    let model: ModelSnapshot =
322        serde_json::from_slice(bytes).map_err(|_| ModelError::MalformedArtifact)?;
323    validate_snapshot(&model).map_err(|_| ModelError::MalformedArtifact)?;
324    Ok(model)
325}
326
327fn validate_config(config: &ModelConfig, sample_count: usize) -> Result<(), ModelError> {
328    if config.components == 0
329        || config.components as usize > sample_count
330        || !config.relevance.is_finite()
331        || config.relevance <= 0.0
332        || !config.variance_floor.is_finite()
333        || config.variance_floor <= 0.0
334        || !config.absolute_threshold.is_finite()
335        || !config.margin_threshold.is_finite()
336    {
337        return Err(ModelError::InvalidConfig);
338    }
339    Ok(())
340}
341
342fn mask_indices(mask: FeatureMask) -> Vec<u8> {
343    let bits = u64::from(mask);
344    (0_u8..35)
345        .filter(|index| bits & (1_u64 << index) != 0)
346        .collect()
347}
348
349fn fit_normalizer(rows: &[Vec<f64>]) -> Result<(Vec<f64>, Vec<f64>), ModelError> {
350    let dimension = rows.first().map_or(0, Vec::len);
351    if rows.is_empty() || dimension == 0 || rows.iter().any(|row| row.len() != dimension) {
352        return Err(ModelError::InvalidSamples);
353    }
354    let count = rows.len() as f64;
355    let mut mean = vec![0.0; dimension];
356    for row in rows {
357        for (sum, value) in mean.iter_mut().zip(row) {
358            *sum += value;
359        }
360    }
361    for value in &mut mean {
362        *value /= count;
363    }
364
365    let mut std = vec![0.0; dimension];
366    for row in rows {
367        for feature in 0..dimension {
368            let delta = row[feature] - mean[feature];
369            std[feature] += delta * delta;
370        }
371    }
372    for value in &mut std {
373        *value = (*value / count).sqrt();
374        if !value.is_finite() || *value <= 0.0 {
375            return Err(ModelError::ZeroVariance);
376        }
377    }
378    if mean.iter().any(|value| !value.is_finite()) {
379        return Err(ModelError::Numerical);
380    }
381    Ok((mean, std))
382}
383
384fn normalize(row: &[f64], mean: &[f64], std: &[f64]) -> Result<Vec<f64>, ModelError> {
385    if row.len() != mean.len() || mean.len() != std.len() {
386        return Err(ModelError::Numerical);
387    }
388    row.iter()
389        .zip(mean.iter().zip(std))
390        .map(|(value, (center, scale))| {
391            let normalized = (value - center) / scale;
392            normalized
393                .is_finite()
394                .then_some(normalized)
395                .ok_or(ModelError::Numerical)
396        })
397        .collect()
398}
399
400fn fit_ubm_with_limits(
401    rows: &[Vec<f64>],
402    components: u8,
403    variance_floor: f64,
404    max_iterations: u16,
405    relative_tolerance: f64,
406) -> Result<DiagGmm, ModelError> {
407    let result = fit_gmm(
408        rows,
409        GmmFitConfig {
410            components,
411            max_iterations,
412            relative_tolerance,
413            variance_floor,
414        },
415    )
416    .map_err(|_| ModelError::Gmm)?;
417    if !result.converged {
418        return Err(ModelError::Nonconverged);
419    }
420    Ok(result.model)
421}
422
423fn map_means(
424    ubm_means: &[Vec<f64>],
425    occupancy: &[f64],
426    first_moments: &[Vec<f64>],
427    relevance: f64,
428) -> Result<Vec<Vec<f64>>, ModelError> {
429    if occupancy.len() != ubm_means.len() || first_moments.len() != ubm_means.len() {
430        return Err(ModelError::Numerical);
431    }
432    let mut adapted = ubm_means.to_vec();
433    for component in 0..ubm_means.len() {
434        let n = occupancy[component];
435        if !n.is_finite() || n < 0.0 || first_moments[component].len() != ubm_means[component].len()
436        {
437            return Err(ModelError::Numerical);
438        }
439        if n > 0.0 {
440            let alpha = n / (n + relevance);
441            for feature in 0..ubm_means[component].len() {
442                let empirical = first_moments[component][feature] / n;
443                let value = alpha * empirical + (1.0 - alpha) * ubm_means[component][feature];
444                if !value.is_finite() {
445                    return Err(ModelError::Numerical);
446                }
447                adapted[component][feature] = value;
448            }
449        }
450    }
451    Ok(adapted)
452}
453
454fn manifest(samples: &[&LabeledSample]) -> [u8; 32] {
455    let mut digest = Sha256::new();
456    for sample in samples {
457        let bytes = sample.sample_id.as_ref().as_bytes();
458        digest.update((bytes.len() as u64).to_be_bytes());
459        digest.update(bytes);
460    }
461    digest.finalize().into()
462}
463
464fn validate_snapshot(model: &ModelSnapshot) -> Result<(), ModelError> {
465    if model.version != SNAPSHOT_VERSION || model.training_sample_count == 0 {
466        return Err(ModelError::MalformedArtifact);
467    }
468    validate_config(&model.config, model.training_sample_count)
469        .map_err(|_| ModelError::MalformedArtifact)?;
470    let expected_indices = mask_indices(model.config.mask);
471    if model.selected_indices != expected_indices
472        || model.normalizer_mean.len() != expected_indices.len()
473        || model.normalizer_std.len() != expected_indices.len()
474        || model.normalizer_mean.iter().any(|value| !value.is_finite())
475        || model
476            .normalizer_std
477            .iter()
478            .any(|value| !value.is_finite() || *value <= 0.0)
479    {
480        return Err(ModelError::MalformedArtifact);
481    }
482
483    let dimension = expected_indices.len();
484    let components = model.config.components as usize;
485    let ubm_means = means(&model.ubm);
486    if ubm_means.len() != components
487        || ubm_means
488            .iter()
489            .any(|row| row.len() != dimension || row.iter().any(|value| !value.is_finite()))
490        || log_likelihood(&model.ubm, &vec![0.0; dimension]).is_err()
491        || model.speakers.is_empty()
492        || model.speakers.len() > model.training_sample_count
493    {
494        return Err(ModelError::MalformedArtifact);
495    }
496
497    let mut counted_samples = 0_usize;
498    for (index, speaker) in model.speakers.iter().enumerate() {
499        if speaker.sample_count == 0
500            || index > 0
501                && model.speakers[index - 1].speaker_id.as_ref() >= speaker.speaker_id.as_ref()
502            || speaker.occupancy.len() != components
503            || speaker.first_moments.len() != components
504            || speaker.adapted_means.len() != components
505        {
506            return Err(ModelError::MalformedArtifact);
507        }
508        counted_samples = counted_samples
509            .checked_add(speaker.sample_count)
510            .ok_or(ModelError::MalformedArtifact)?;
511
512        let occupancy_sum: f64 = speaker.occupancy.iter().sum();
513        if !occupancy_sum.is_finite()
514            || (occupancy_sum - speaker.sample_count as f64).abs()
515                > 1e-8 * (speaker.sample_count as f64).max(1.0)
516        {
517            return Err(ModelError::MalformedArtifact);
518        }
519        let expected = map_means(
520            ubm_means,
521            &speaker.occupancy,
522            &speaker.first_moments,
523            model.config.relevance,
524        )
525        .map_err(|_| ModelError::MalformedArtifact)?;
526        for (component, occupancy) in speaker.occupancy.iter().enumerate() {
527            if !occupancy.is_finite()
528                || *occupancy < 0.0
529                || speaker.first_moments[component].len() != dimension
530                || speaker.adapted_means[component].len() != dimension
531            {
532                return Err(ModelError::MalformedArtifact);
533            }
534            for (feature, first) in speaker.first_moments[component].iter().enumerate() {
535                let actual = speaker.adapted_means[component][feature];
536                if !first.is_finite()
537                    || !actual.is_finite()
538                    || actual.to_bits() != expected[component][feature].to_bits()
539                {
540                    return Err(ModelError::MalformedArtifact);
541                }
542            }
543        }
544        let adapted = with_means(&model.ubm, speaker.adapted_means.clone())
545            .map_err(|_| ModelError::MalformedArtifact)?;
546        if log_likelihood(&adapted, &vec![0.0; dimension]).is_err() {
547            return Err(ModelError::MalformedArtifact);
548        }
549    }
550    if counted_samples != model.training_sample_count {
551        return Err(ModelError::MalformedArtifact);
552    }
553    Ok(())
554}
555
556#[cfg(test)]
557mod tests {
558    use super::*;
559    use kcode_speaker_types::{ObjectId, RecordingKind, SegmentRef};
560    use serde::de::DeserializeOwned;
561    use serde_json::{Map, Value, json};
562
563    fn key(value: &str) -> Key {
564        Key::parse(value).unwrap()
565    }
566
567    fn vector(seed: u8) -> FeatureVector {
568        let mut values = [0_u8; 35];
569        for (index, value) in values.iter_mut().enumerate() {
570            *value = (seed as usize + index * 3) as u8 % 90;
571        }
572        FeatureVector::new(values).unwrap()
573    }
574
575    fn missing_field(error: &serde_json::Error) -> Option<String> {
576        let message = error.to_string();
577        let start = message.find("missing field `")? + "missing field `".len();
578        let end = message[start..].find('`')? + start;
579        Some(message[start..end].to_owned())
580    }
581
582    fn build_missing<T, F>(mut fields: Map<String, Value>, mut value_for: F) -> T
583    where
584        T: DeserializeOwned,
585        F: FnMut(&str) -> Value,
586    {
587        loop {
588            match serde_json::from_value(Value::Object(fields.clone())) {
589                Ok(value) => return value,
590                Err(error) => {
591                    let field = missing_field(&error)
592                        .unwrap_or_else(|| panic!("invalid test fixture: {error}"));
593                    assert!(
594                        !fields.contains_key(&field),
595                        "invalid value for test fixture field {field}: {error}"
596                    );
597                    fields.insert(field.clone(), value_for(&field));
598                }
599            }
600        }
601    }
602
603    fn segment_value() -> Value {
604        let segment: SegmentRef = build_missing(Map::new(), |field| match field {
605            "ordinal" | "start_ms" => json!(0),
606            "segment_count" => json!(1),
607            "end_ms" => json!(100),
608            name if name.contains("object") => {
609                serde_json::to_value(ObjectId::parse("AAAAAAAA").unwrap()).unwrap()
610            }
611            name if name.ends_with("_id") => serde_json::to_value(key("recording")).unwrap(),
612            name if name.contains("kind") => {
613                serde_json::to_value(RecordingKind::VoiceNote).unwrap()
614            }
615            _ => json!(0),
616        });
617        segment.validate().unwrap();
618        serde_json::to_value(segment).unwrap()
619    }
620
621    fn sample(id: &str, cohort: &str, speaker: &str, seed: u8) -> LabeledSample {
622        let features = vector(seed);
623        let mut fields = Map::new();
624        fields.insert("sample_id".into(), serde_json::to_value(key(id)).unwrap());
625        fields.insert(
626            "cohort_id".into(),
627            serde_json::to_value(key(cohort)).unwrap(),
628        );
629        fields.insert(
630            "speaker_id".into(),
631            serde_json::to_value(key(speaker)).unwrap(),
632        );
633        fields.insert("features".into(), serde_json::to_value(&features).unwrap());
634        build_missing(fields, |field| match field {
635            "recording_kind" => serde_json::to_value(RecordingKind::VoiceNote).unwrap(),
636            "primary_language" => serde_json::to_value(key("eng")).unwrap(),
637            "segment" | "segment_ref" => segment_value(),
638            name if name.contains("object") => {
639                serde_json::to_value(ObjectId::parse("AAAAAAAA").unwrap()).unwrap()
640            }
641            name if name.ends_with("_id") => serde_json::to_value(key(name)).unwrap(),
642            name if name.contains("confirmed") => json!(true),
643            name if name.ends_with("_count") => json!(1),
644            name if name.ends_with("_ms") => json!(100),
645            _ => json!(0),
646        })
647    }
648
649    fn config(mask: FeatureMask) -> ModelConfig {
650        ModelConfig {
651            mask,
652            components: 1,
653            relevance: 2.0,
654            variance_floor: 0.01,
655            absolute_threshold: -1e9,
656            margin_threshold: -1e9,
657        }
658    }
659
660    fn training() -> Vec<LabeledSample> {
661        vec![
662            sample("AAAAAAA1", "cohort", "alice", 10),
663            sample("AAAAAAA2", "cohort", "alice", 14),
664            sample("AAAAAAA3", "cohort", "bob", 60),
665            sample("AAAAAAA4", "cohort", "bob", 66),
666        ]
667    }
668
669    #[test]
670    fn masks_match_frozen_sets() {
671        let bits11 = [0, 1, 2, 4, 5, 8, 11, 12, 13, 17, 22]
672            .into_iter()
673            .fold(0_u64, |bits, index| bits | (1 << index));
674        let bits20 = (0..=17)
675            .chain([19, 22])
676            .fold(0_u64, |bits, index| bits | (1 << index));
677        assert_eq!(u64::from(GEMINI_11), bits11);
678        assert_eq!(u64::from(GEMINI_20), bits20);
679        assert_eq!(u64::from(ALL_35), (1_u64 << 35) - 1);
680    }
681
682    #[test]
683    fn hand_computed_map_and_unoccupied_component() {
684        let ubm = vec![vec![1.0, -1.0], vec![7.0, 8.0]];
685        let n = vec![2.0, 0.0];
686        let f = vec![vec![6.0, 2.0], vec![0.0, 0.0]];
687        let adapted = map_means(&ubm, &n, &f, 2.0).unwrap();
688        assert_eq!(adapted[0], vec![2.0, 0.0]);
689        assert_eq!(adapted[1], ubm[1]);
690    }
691
692    #[test]
693    fn k1_llr_is_shared_covariance_mahalanobis_difference() {
694        let samples = training();
695        let cohort = key("cohort");
696        let mask = FeatureMask::from_bits(1).unwrap();
697        let model = fit(FitInput {
698            cohort_id: &cohort,
699            samples: &samples,
700            config: config(mask),
701        })
702        .unwrap();
703        let probe = vector(12);
704        let result = identify(&model, &cohort, &probe).unwrap();
705        let speaker = model
706            .speakers
707            .iter()
708            .find(|speaker| speaker.speaker_id.as_ref() == result.best.speaker_id.as_ref())
709            .unwrap();
710        let x = (probe.as_ref()[0] as f64 - model.normalizer_mean[0]) / model.normalizer_std[0];
711        let ubm_mean = means(&model.ubm)[0][0];
712        let value = serde_json::to_value(&model.ubm).unwrap();
713        let variance = value["variances"][0][0].as_f64().unwrap();
714        let adapted_mean = speaker.adapted_means[0][0];
715        let expected =
716            -0.5 * ((x - adapted_mean).powi(2) / variance - (x - ubm_mean).powi(2) / variance);
717        assert!((result.best.llr - expected).abs() < 1e-12);
718    }
719
720    #[test]
721    fn threshold_equality_and_one_speaker_rules() {
722        let samples = training();
723        let cohort = key("cohort");
724        let mut model = fit(FitInput {
725            cohort_id: &cohort,
726            samples: &samples,
727            config: config(GEMINI_11),
728        })
729        .unwrap();
730        let probe = vector(12);
731        let initial = identify(&model, &cohort, &probe).unwrap();
732        let margin = initial.best.llr - initial.runner_up.as_ref().unwrap().llr;
733        model.config.absolute_threshold = initial.best.llr;
734        model.config.margin_threshold = margin;
735        let equality = identify(&model, &cohort, &probe).unwrap();
736        assert!(equality.absolute_pass);
737        assert!(equality.margin_pass);
738        assert!(matches!(equality.decision, Decision::Known { .. }));
739
740        let one_speaker = &samples[..2];
741        let one = fit(FitInput {
742            cohort_id: &cohort,
743            samples: one_speaker,
744            config: config(GEMINI_11),
745        })
746        .unwrap();
747        let result = identify(&one, &cohort, &probe).unwrap();
748        assert!(result.runner_up.is_none());
749        assert!(result.margin_pass);
750
751        let mut blocked = one;
752        blocked.config.absolute_threshold = result.best.llr + 1e-12;
753        assert!(matches!(
754            identify(&blocked, &cohort, &probe).unwrap().decision,
755            Decision::Unknown
756        ));
757    }
758
759    #[test]
760    fn cohort_mismatch_is_typed() {
761        let samples = training();
762        let cohort = key("cohort");
763        let model = fit(FitInput {
764            cohort_id: &cohort,
765            samples: &samples,
766            config: config(GEMINI_11),
767        })
768        .unwrap();
769        assert_eq!(
770            identify(&model, &key("different"), &vector(12))
771                .err()
772                .unwrap(),
773            ModelError::CohortMismatch
774        );
775    }
776
777    #[test]
778    fn manifest_changes_and_input_order_does_not() {
779        let samples = training();
780        let cohort = key("cohort");
781        let first = fit(FitInput {
782            cohort_id: &cohort,
783            samples: &samples,
784            config: config(GEMINI_11),
785        })
786        .unwrap();
787        let mut reversed = training();
788        reversed.reverse();
789        let second = fit(FitInput {
790            cohort_id: &cohort,
791            samples: &reversed,
792            config: config(GEMINI_11),
793        })
794        .unwrap();
795        assert_eq!(encode(&first).unwrap(), encode(&second).unwrap());
796
797        let mut changed = training();
798        changed[0].sample_id = key("BBBBBBB1");
799        let third = fit(FitInput {
800            cohort_id: &cohort,
801            samples: &changed,
802            config: config(GEMINI_11),
803        })
804        .unwrap();
805        assert_ne!(first.manifest_sha256, third.manifest_sha256);
806        assert_ne!(encode(&first).unwrap(), encode(&third).unwrap());
807    }
808
809    #[test]
810    fn serialization_is_deterministic_and_rejects_corruption() {
811        let samples = training();
812        let cohort = key("cohort");
813        let model = fit(FitInput {
814            cohort_id: &cohort,
815            samples: &samples,
816            config: config(GEMINI_11),
817        })
818        .unwrap();
819        let bytes = encode(&model).unwrap();
820        let decoded = decode(&bytes).unwrap();
821        assert_eq!(bytes, encode(&decoded).unwrap());
822        assert_eq!(encode(&model).unwrap(), encode(&model).unwrap());
823        assert_eq!(decode(b"{}").err().unwrap(), ModelError::MalformedArtifact);
824
825        let mut value: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
826        value["version"] = serde_json::json!(2);
827        assert_eq!(
828            decode(&serde_json::to_vec(&value).unwrap()).err().unwrap(),
829            ModelError::MalformedArtifact
830        );
831    }
832
833    #[test]
834    fn rejects_bad_samples_and_configuration() {
835        let cohort = key("cohort");
836        let mut mixed = training();
837        mixed[0].cohort_id = key("other");
838        assert_eq!(
839            fit(FitInput {
840                cohort_id: &cohort,
841                samples: &mixed,
842                config: config(GEMINI_11),
843            })
844            .err()
845            .unwrap(),
846            ModelError::CohortMismatch
847        );
848
849        let mut duplicate = training();
850        duplicate[1].sample_id = duplicate[0].sample_id.clone();
851        assert_eq!(
852            fit(FitInput {
853                cohort_id: &cohort,
854                samples: &duplicate,
855                config: config(GEMINI_11),
856            })
857            .err()
858            .unwrap(),
859            ModelError::DuplicateSampleId
860        );
861
862        let same = vec![
863            sample("CCCCCCC1", "cohort", "alice", 10),
864            sample("CCCCCCC2", "cohort", "alice", 10),
865        ];
866        assert_eq!(
867            fit(FitInput {
868                cohort_id: &cohort,
869                samples: &same,
870                config: config(GEMINI_11),
871            })
872            .err()
873            .unwrap(),
874            ModelError::ZeroVariance
875        );
876
877        let samples = training();
878        for bad in [
879            ModelConfig {
880                components: 0,
881                ..config(GEMINI_11)
882            },
883            ModelConfig {
884                components: 5,
885                ..config(GEMINI_11)
886            },
887            ModelConfig {
888                relevance: 0.0,
889                ..config(GEMINI_11)
890            },
891            ModelConfig {
892                variance_floor: f64::NAN,
893                ..config(GEMINI_11)
894            },
895            ModelConfig {
896                absolute_threshold: f64::INFINITY,
897                ..config(GEMINI_11)
898            },
899        ] {
900            assert_eq!(
901                fit(FitInput {
902                    cohort_id: &cohort,
903                    samples: &samples,
904                    config: bad,
905                })
906                .err()
907                .unwrap(),
908                ModelError::InvalidConfig
909            );
910        }
911    }
912
913    #[test]
914    fn rejects_nonconverged_gmm() {
915        let rows = vec![
916            vec![-5.0],
917            vec![-4.0],
918            vec![-1.0],
919            vec![2.0],
920            vec![3.0],
921            vec![10.0],
922        ];
923        assert_eq!(
924            fit_ubm_with_limits(&rows, 2, 0.01, 1, f64::MIN_POSITIVE)
925                .err()
926                .unwrap(),
927            ModelError::Nonconverged
928        );
929    }
930}