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}