use kcode_diag_gmm::{
DiagGmm, FitConfig as GmmFitConfig, fit as fit_gmm, log_likelihood, means, responsibilities,
with_means,
};
use kcode_speaker_types::{FeatureMask, FeatureVector, Key, LabeledSample};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::fmt;
const SNAPSHOT_VERSION: u8 = 1;
const MAX_ITERATIONS: u16 = 200;
const RELATIVE_TOLERANCE: f64 = 1e-8;
pub const GEMINI_11: FeatureMask = match FeatureMask::from_bits(
(1 << 0)
| (1 << 1)
| (1 << 2)
| (1 << 4)
| (1 << 5)
| (1 << 8)
| (1 << 11)
| (1 << 12)
| (1 << 13)
| (1 << 17)
| (1 << 22),
) {
Ok(mask) => mask,
Err(_) => panic!("invalid frozen mask"),
};
pub const GEMINI_20: FeatureMask =
match FeatureMask::from_bits(((1_u64 << 18) - 1) | (1 << 19) | (1 << 22)) {
Ok(mask) => mask,
Err(_) => panic!("invalid frozen mask"),
};
pub const ALL_35: FeatureMask = match FeatureMask::from_bits((1_u64 << 35) - 1) {
Ok(mask) => mask,
Err(_) => panic!("invalid frozen mask"),
};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ModelConfig {
pub mask: FeatureMask,
pub components: u8,
pub relevance: f64,
pub variance_floor: f64,
pub absolute_threshold: f64,
pub margin_threshold: f64,
}
pub struct FitInput<'a> {
pub cohort_id: &'a Key,
pub samples: &'a [LabeledSample],
pub config: ModelConfig,
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ModelSnapshot {
version: u8,
cohort_id: Key,
config: ModelConfig,
selected_indices: Vec<u8>,
normalizer_mean: Vec<f64>,
normalizer_std: Vec<f64>,
ubm: DiagGmm,
speakers: Vec<SpeakerModel>,
training_sample_count: usize,
manifest_sha256: [u8; 32],
}
#[derive(Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct SpeakerModel {
speaker_id: Key,
sample_count: usize,
occupancy: Vec<f64>,
first_moments: Vec<Vec<f64>>,
adapted_means: Vec<Vec<f64>>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct CandidateScore {
pub speaker_id: Key,
pub llr: f64,
}
#[derive(Clone, Debug, PartialEq)]
pub enum Decision {
Known { speaker_id: Key },
Unknown,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Identification {
pub decision: Decision,
pub best: CandidateScore,
pub runner_up: Option<CandidateScore>,
pub absolute_pass: bool,
pub margin_pass: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ModelError {
InvalidConfig,
InvalidSamples,
CohortMismatch,
DuplicateSampleId,
ZeroVariance,
Nonconverged,
Numerical,
Gmm,
Serialization,
MalformedArtifact,
}
impl fmt::Display for ModelError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let text = match self {
Self::InvalidConfig => "invalid model configuration",
Self::InvalidSamples => "invalid training samples",
Self::CohortMismatch => "cohort does not match",
Self::DuplicateSampleId => "duplicate training sample ID",
Self::ZeroVariance => "selected training feature has zero variance",
Self::Nonconverged => "diagonal GMM did not converge",
Self::Numerical => "nonfinite model computation",
Self::Gmm => "diagonal GMM operation failed",
Self::Serialization => "snapshot serialization failed",
Self::MalformedArtifact => "malformed or incompatible snapshot artifact",
};
f.write_str(text)
}
}
impl std::error::Error for ModelError {}
pub fn fit(input: FitInput<'_>) -> Result<ModelSnapshot, ModelError> {
validate_config(&input.config, input.samples.len())?;
if input.samples.is_empty() {
return Err(ModelError::InvalidSamples);
}
let mut samples: Vec<&LabeledSample> = input.samples.iter().collect();
samples.sort_by(|a, b| a.sample_id.as_ref().cmp(b.sample_id.as_ref()));
for sample in &samples {
if sample.cohort_id.as_ref() != input.cohort_id.as_ref() {
return Err(ModelError::CohortMismatch);
}
}
if samples
.windows(2)
.any(|pair| pair[0].sample_id.as_ref() == pair[1].sample_id.as_ref())
{
return Err(ModelError::DuplicateSampleId);
}
let selected = mask_indices(input.config.mask);
let raw_rows: Vec<Vec<f64>> = samples
.iter()
.map(|sample| {
selected
.iter()
.map(|&index| sample.features.as_ref()[index as usize] as f64)
.collect()
})
.collect();
let (normalizer_mean, normalizer_std) = fit_normalizer(&raw_rows)?;
let rows: Vec<Vec<f64>> = raw_rows
.iter()
.map(|row| normalize(row, &normalizer_mean, &normalizer_std))
.collect::<Result<_, _>>()?;
let ubm = fit_ubm_with_limits(
&rows,
input.config.components,
input.config.variance_floor,
MAX_ITERATIONS,
RELATIVE_TOLERANCE,
)?;
let ubm_means = means(&ubm);
let mut speaker_ids: Vec<Key> = samples
.iter()
.map(|sample| sample.speaker_id.clone())
.collect();
speaker_ids.sort_by(|a, b| a.as_ref().cmp(b.as_ref()));
speaker_ids.dedup_by(|a, b| a.as_ref() == b.as_ref());
if speaker_ids.is_empty() {
return Err(ModelError::InvalidSamples);
}
let components = input.config.components as usize;
let dimension = selected.len();
let mut speakers: Vec<SpeakerModel> = speaker_ids
.into_iter()
.map(|speaker_id| SpeakerModel {
speaker_id,
sample_count: 0,
occupancy: vec![0.0; components],
first_moments: vec![vec![0.0; dimension]; components],
adapted_means: Vec::new(),
})
.collect();
for (sample, row) in samples.iter().zip(&rows) {
let speaker_index = speakers
.binary_search_by(|speaker| speaker.speaker_id.as_ref().cmp(sample.speaker_id.as_ref()))
.map_err(|_| ModelError::InvalidSamples)?;
let gamma = responsibilities(&ubm, row).map_err(|_| ModelError::Gmm)?;
let speaker = &mut speakers[speaker_index];
speaker.sample_count += 1;
for (component, responsibility) in gamma[..components].iter().copied().enumerate() {
if !responsibility.is_finite() || responsibility < 0.0 {
return Err(ModelError::Numerical);
}
speaker.occupancy[component] += responsibility;
for (feature, value) in row.iter().enumerate() {
speaker.first_moments[component][feature] += responsibility * value;
}
}
}
for speaker in &mut speakers {
speaker.adapted_means = map_means(
ubm_means,
&speaker.occupancy,
&speaker.first_moments,
input.config.relevance,
)?;
}
let manifest_sha256 = manifest(&samples);
let model = ModelSnapshot {
version: SNAPSHOT_VERSION,
cohort_id: input.cohort_id.clone(),
config: input.config,
selected_indices: selected,
normalizer_mean,
normalizer_std,
ubm,
speakers,
training_sample_count: samples.len(),
manifest_sha256,
};
validate_snapshot(&model)?;
Ok(model)
}
pub fn identify(
model: &ModelSnapshot,
cohort_id: &Key,
features: &FeatureVector,
) -> Result<Identification, ModelError> {
validate_snapshot(model)?;
if cohort_id.as_ref() != model.cohort_id.as_ref() {
return Err(ModelError::CohortMismatch);
}
let raw: Vec<f64> = model
.selected_indices
.iter()
.map(|&index| features.as_ref()[index as usize] as f64)
.collect();
let row = normalize(&raw, &model.normalizer_mean, &model.normalizer_std)?;
let ubm_score = log_likelihood(&model.ubm, &row).map_err(|_| ModelError::Gmm)?;
if !ubm_score.is_finite() {
return Err(ModelError::Numerical);
}
let mut scores = Vec::with_capacity(model.speakers.len());
for speaker in &model.speakers {
let adapted =
with_means(&model.ubm, speaker.adapted_means.clone()).map_err(|_| ModelError::Gmm)?;
let score = log_likelihood(&adapted, &row).map_err(|_| ModelError::Gmm)? - ubm_score;
if !score.is_finite() {
return Err(ModelError::Numerical);
}
scores.push(CandidateScore {
speaker_id: speaker.speaker_id.clone(),
llr: score,
});
}
scores.sort_by(|a, b| {
b.llr
.total_cmp(&a.llr)
.then_with(|| a.speaker_id.as_ref().cmp(b.speaker_id.as_ref()))
});
let best = scores.remove(0);
let runner_up = scores.into_iter().next();
let absolute_pass = best.llr >= model.config.absolute_threshold;
let margin_pass = match &runner_up {
Some(runner) => best.llr - runner.llr >= model.config.margin_threshold,
None => true,
};
let decision = if absolute_pass && margin_pass {
Decision::Known {
speaker_id: best.speaker_id.clone(),
}
} else {
Decision::Unknown
};
Ok(Identification {
decision,
best,
runner_up,
absolute_pass,
margin_pass,
})
}
pub fn encode(model: &ModelSnapshot) -> Result<Vec<u8>, ModelError> {
validate_snapshot(model)?;
serde_json::to_vec(model).map_err(|_| ModelError::Serialization)
}
pub fn decode(bytes: &[u8]) -> Result<ModelSnapshot, ModelError> {
let model: ModelSnapshot =
serde_json::from_slice(bytes).map_err(|_| ModelError::MalformedArtifact)?;
validate_snapshot(&model).map_err(|_| ModelError::MalformedArtifact)?;
Ok(model)
}
fn validate_config(config: &ModelConfig, sample_count: usize) -> Result<(), ModelError> {
if config.components == 0
|| config.components as usize > sample_count
|| !config.relevance.is_finite()
|| config.relevance <= 0.0
|| !config.variance_floor.is_finite()
|| config.variance_floor <= 0.0
|| !config.absolute_threshold.is_finite()
|| !config.margin_threshold.is_finite()
{
return Err(ModelError::InvalidConfig);
}
Ok(())
}
fn mask_indices(mask: FeatureMask) -> Vec<u8> {
let bits = u64::from(mask);
(0_u8..35)
.filter(|index| bits & (1_u64 << index) != 0)
.collect()
}
fn fit_normalizer(rows: &[Vec<f64>]) -> Result<(Vec<f64>, Vec<f64>), ModelError> {
let dimension = rows.first().map_or(0, Vec::len);
if rows.is_empty() || dimension == 0 || rows.iter().any(|row| row.len() != dimension) {
return Err(ModelError::InvalidSamples);
}
let count = rows.len() as f64;
let mut mean = vec![0.0; dimension];
for row in rows {
for (sum, value) in mean.iter_mut().zip(row) {
*sum += value;
}
}
for value in &mut mean {
*value /= count;
}
let mut std = vec![0.0; dimension];
for row in rows {
for feature in 0..dimension {
let delta = row[feature] - mean[feature];
std[feature] += delta * delta;
}
}
for value in &mut std {
*value = (*value / count).sqrt();
if !value.is_finite() || *value <= 0.0 {
return Err(ModelError::ZeroVariance);
}
}
if mean.iter().any(|value| !value.is_finite()) {
return Err(ModelError::Numerical);
}
Ok((mean, std))
}
fn normalize(row: &[f64], mean: &[f64], std: &[f64]) -> Result<Vec<f64>, ModelError> {
if row.len() != mean.len() || mean.len() != std.len() {
return Err(ModelError::Numerical);
}
row.iter()
.zip(mean.iter().zip(std))
.map(|(value, (center, scale))| {
let normalized = (value - center) / scale;
normalized
.is_finite()
.then_some(normalized)
.ok_or(ModelError::Numerical)
})
.collect()
}
fn fit_ubm_with_limits(
rows: &[Vec<f64>],
components: u8,
variance_floor: f64,
max_iterations: u16,
relative_tolerance: f64,
) -> Result<DiagGmm, ModelError> {
let result = fit_gmm(
rows,
GmmFitConfig {
components,
max_iterations,
relative_tolerance,
variance_floor,
},
)
.map_err(|_| ModelError::Gmm)?;
if !result.converged {
return Err(ModelError::Nonconverged);
}
Ok(result.model)
}
fn map_means(
ubm_means: &[Vec<f64>],
occupancy: &[f64],
first_moments: &[Vec<f64>],
relevance: f64,
) -> Result<Vec<Vec<f64>>, ModelError> {
if occupancy.len() != ubm_means.len() || first_moments.len() != ubm_means.len() {
return Err(ModelError::Numerical);
}
let mut adapted = ubm_means.to_vec();
for component in 0..ubm_means.len() {
let n = occupancy[component];
if !n.is_finite() || n < 0.0 || first_moments[component].len() != ubm_means[component].len()
{
return Err(ModelError::Numerical);
}
if n > 0.0 {
let alpha = n / (n + relevance);
for feature in 0..ubm_means[component].len() {
let empirical = first_moments[component][feature] / n;
let value = alpha * empirical + (1.0 - alpha) * ubm_means[component][feature];
if !value.is_finite() {
return Err(ModelError::Numerical);
}
adapted[component][feature] = value;
}
}
}
Ok(adapted)
}
fn manifest(samples: &[&LabeledSample]) -> [u8; 32] {
let mut digest = Sha256::new();
for sample in samples {
let bytes = sample.sample_id.as_ref().as_bytes();
digest.update((bytes.len() as u64).to_be_bytes());
digest.update(bytes);
}
digest.finalize().into()
}
fn validate_snapshot(model: &ModelSnapshot) -> Result<(), ModelError> {
if model.version != SNAPSHOT_VERSION || model.training_sample_count == 0 {
return Err(ModelError::MalformedArtifact);
}
validate_config(&model.config, model.training_sample_count)
.map_err(|_| ModelError::MalformedArtifact)?;
let expected_indices = mask_indices(model.config.mask);
if model.selected_indices != expected_indices
|| model.normalizer_mean.len() != expected_indices.len()
|| model.normalizer_std.len() != expected_indices.len()
|| model.normalizer_mean.iter().any(|value| !value.is_finite())
|| model
.normalizer_std
.iter()
.any(|value| !value.is_finite() || *value <= 0.0)
{
return Err(ModelError::MalformedArtifact);
}
let dimension = expected_indices.len();
let components = model.config.components as usize;
let ubm_means = means(&model.ubm);
if ubm_means.len() != components
|| ubm_means
.iter()
.any(|row| row.len() != dimension || row.iter().any(|value| !value.is_finite()))
|| log_likelihood(&model.ubm, &vec![0.0; dimension]).is_err()
|| model.speakers.is_empty()
|| model.speakers.len() > model.training_sample_count
{
return Err(ModelError::MalformedArtifact);
}
let mut counted_samples = 0_usize;
for (index, speaker) in model.speakers.iter().enumerate() {
if speaker.sample_count == 0
|| index > 0
&& model.speakers[index - 1].speaker_id.as_ref() >= speaker.speaker_id.as_ref()
|| speaker.occupancy.len() != components
|| speaker.first_moments.len() != components
|| speaker.adapted_means.len() != components
{
return Err(ModelError::MalformedArtifact);
}
counted_samples = counted_samples
.checked_add(speaker.sample_count)
.ok_or(ModelError::MalformedArtifact)?;
let occupancy_sum: f64 = speaker.occupancy.iter().sum();
if !occupancy_sum.is_finite()
|| (occupancy_sum - speaker.sample_count as f64).abs()
> 1e-8 * (speaker.sample_count as f64).max(1.0)
{
return Err(ModelError::MalformedArtifact);
}
let expected = map_means(
ubm_means,
&speaker.occupancy,
&speaker.first_moments,
model.config.relevance,
)
.map_err(|_| ModelError::MalformedArtifact)?;
for (component, occupancy) in speaker.occupancy.iter().enumerate() {
if !occupancy.is_finite()
|| *occupancy < 0.0
|| speaker.first_moments[component].len() != dimension
|| speaker.adapted_means[component].len() != dimension
{
return Err(ModelError::MalformedArtifact);
}
for (feature, first) in speaker.first_moments[component].iter().enumerate() {
let actual = speaker.adapted_means[component][feature];
if !first.is_finite()
|| !actual.is_finite()
|| actual.to_bits() != expected[component][feature].to_bits()
{
return Err(ModelError::MalformedArtifact);
}
}
}
let adapted = with_means(&model.ubm, speaker.adapted_means.clone())
.map_err(|_| ModelError::MalformedArtifact)?;
if log_likelihood(&adapted, &vec![0.0; dimension]).is_err() {
return Err(ModelError::MalformedArtifact);
}
}
if counted_samples != model.training_sample_count {
return Err(ModelError::MalformedArtifact);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use kcode_speaker_types::{ObjectId, RecordingKind, SegmentRef};
use serde::de::DeserializeOwned;
use serde_json::{Map, Value, json};
fn key(value: &str) -> Key {
Key::parse(value).unwrap()
}
fn vector(seed: u8) -> FeatureVector {
let mut values = [0_u8; 35];
for (index, value) in values.iter_mut().enumerate() {
*value = (seed as usize + index * 3) as u8 % 90;
}
FeatureVector::new(values).unwrap()
}
fn missing_field(error: &serde_json::Error) -> Option<String> {
let message = error.to_string();
let start = message.find("missing field `")? + "missing field `".len();
let end = message[start..].find('`')? + start;
Some(message[start..end].to_owned())
}
fn build_missing<T, F>(mut fields: Map<String, Value>, mut value_for: F) -> T
where
T: DeserializeOwned,
F: FnMut(&str) -> Value,
{
loop {
match serde_json::from_value(Value::Object(fields.clone())) {
Ok(value) => return value,
Err(error) => {
let field = missing_field(&error)
.unwrap_or_else(|| panic!("invalid test fixture: {error}"));
assert!(
!fields.contains_key(&field),
"invalid value for test fixture field {field}: {error}"
);
fields.insert(field.clone(), value_for(&field));
}
}
}
}
fn segment_value() -> Value {
let segment: SegmentRef = build_missing(Map::new(), |field| match field {
"ordinal" | "start_ms" => json!(0),
"segment_count" => json!(1),
"end_ms" => json!(100),
name if name.contains("object") => {
serde_json::to_value(ObjectId::parse("AAAAAAAA").unwrap()).unwrap()
}
name if name.ends_with("_id") => serde_json::to_value(key("recording")).unwrap(),
name if name.contains("kind") => {
serde_json::to_value(RecordingKind::VoiceNote).unwrap()
}
_ => json!(0),
});
segment.validate().unwrap();
serde_json::to_value(segment).unwrap()
}
fn sample(id: &str, cohort: &str, speaker: &str, seed: u8) -> LabeledSample {
let features = vector(seed);
let mut fields = Map::new();
fields.insert("sample_id".into(), serde_json::to_value(key(id)).unwrap());
fields.insert(
"cohort_id".into(),
serde_json::to_value(key(cohort)).unwrap(),
);
fields.insert(
"speaker_id".into(),
serde_json::to_value(key(speaker)).unwrap(),
);
fields.insert("features".into(), serde_json::to_value(&features).unwrap());
build_missing(fields, |field| match field {
"recording_kind" => serde_json::to_value(RecordingKind::VoiceNote).unwrap(),
"primary_language" => serde_json::to_value(key("eng")).unwrap(),
"segment" | "segment_ref" => segment_value(),
name if name.contains("object") => {
serde_json::to_value(ObjectId::parse("AAAAAAAA").unwrap()).unwrap()
}
name if name.ends_with("_id") => serde_json::to_value(key(name)).unwrap(),
name if name.contains("confirmed") => json!(true),
name if name.ends_with("_count") => json!(1),
name if name.ends_with("_ms") => json!(100),
_ => json!(0),
})
}
fn config(mask: FeatureMask) -> ModelConfig {
ModelConfig {
mask,
components: 1,
relevance: 2.0,
variance_floor: 0.01,
absolute_threshold: -1e9,
margin_threshold: -1e9,
}
}
fn training() -> Vec<LabeledSample> {
vec![
sample("AAAAAAA1", "cohort", "alice", 10),
sample("AAAAAAA2", "cohort", "alice", 14),
sample("AAAAAAA3", "cohort", "bob", 60),
sample("AAAAAAA4", "cohort", "bob", 66),
]
}
#[test]
fn masks_match_frozen_sets() {
let bits11 = [0, 1, 2, 4, 5, 8, 11, 12, 13, 17, 22]
.into_iter()
.fold(0_u64, |bits, index| bits | (1 << index));
let bits20 = (0..=17)
.chain([19, 22])
.fold(0_u64, |bits, index| bits | (1 << index));
assert_eq!(u64::from(GEMINI_11), bits11);
assert_eq!(u64::from(GEMINI_20), bits20);
assert_eq!(u64::from(ALL_35), (1_u64 << 35) - 1);
}
#[test]
fn hand_computed_map_and_unoccupied_component() {
let ubm = vec![vec![1.0, -1.0], vec![7.0, 8.0]];
let n = vec![2.0, 0.0];
let f = vec![vec![6.0, 2.0], vec![0.0, 0.0]];
let adapted = map_means(&ubm, &n, &f, 2.0).unwrap();
assert_eq!(adapted[0], vec![2.0, 0.0]);
assert_eq!(adapted[1], ubm[1]);
}
#[test]
fn k1_llr_is_shared_covariance_mahalanobis_difference() {
let samples = training();
let cohort = key("cohort");
let mask = FeatureMask::from_bits(1).unwrap();
let model = fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: config(mask),
})
.unwrap();
let probe = vector(12);
let result = identify(&model, &cohort, &probe).unwrap();
let speaker = model
.speakers
.iter()
.find(|speaker| speaker.speaker_id.as_ref() == result.best.speaker_id.as_ref())
.unwrap();
let x = (probe.as_ref()[0] as f64 - model.normalizer_mean[0]) / model.normalizer_std[0];
let ubm_mean = means(&model.ubm)[0][0];
let value = serde_json::to_value(&model.ubm).unwrap();
let variance = value["variances"][0][0].as_f64().unwrap();
let adapted_mean = speaker.adapted_means[0][0];
let expected =
-0.5 * ((x - adapted_mean).powi(2) / variance - (x - ubm_mean).powi(2) / variance);
assert!((result.best.llr - expected).abs() < 1e-12);
}
#[test]
fn threshold_equality_and_one_speaker_rules() {
let samples = training();
let cohort = key("cohort");
let mut model = fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: config(GEMINI_11),
})
.unwrap();
let probe = vector(12);
let initial = identify(&model, &cohort, &probe).unwrap();
let margin = initial.best.llr - initial.runner_up.as_ref().unwrap().llr;
model.config.absolute_threshold = initial.best.llr;
model.config.margin_threshold = margin;
let equality = identify(&model, &cohort, &probe).unwrap();
assert!(equality.absolute_pass);
assert!(equality.margin_pass);
assert!(matches!(equality.decision, Decision::Known { .. }));
let one_speaker = &samples[..2];
let one = fit(FitInput {
cohort_id: &cohort,
samples: one_speaker,
config: config(GEMINI_11),
})
.unwrap();
let result = identify(&one, &cohort, &probe).unwrap();
assert!(result.runner_up.is_none());
assert!(result.margin_pass);
let mut blocked = one;
blocked.config.absolute_threshold = result.best.llr + 1e-12;
assert!(matches!(
identify(&blocked, &cohort, &probe).unwrap().decision,
Decision::Unknown
));
}
#[test]
fn cohort_mismatch_is_typed() {
let samples = training();
let cohort = key("cohort");
let model = fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: config(GEMINI_11),
})
.unwrap();
assert_eq!(
identify(&model, &key("different"), &vector(12))
.err()
.unwrap(),
ModelError::CohortMismatch
);
}
#[test]
fn manifest_changes_and_input_order_does_not() {
let samples = training();
let cohort = key("cohort");
let first = fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: config(GEMINI_11),
})
.unwrap();
let mut reversed = training();
reversed.reverse();
let second = fit(FitInput {
cohort_id: &cohort,
samples: &reversed,
config: config(GEMINI_11),
})
.unwrap();
assert_eq!(encode(&first).unwrap(), encode(&second).unwrap());
let mut changed = training();
changed[0].sample_id = key("BBBBBBB1");
let third = fit(FitInput {
cohort_id: &cohort,
samples: &changed,
config: config(GEMINI_11),
})
.unwrap();
assert_ne!(first.manifest_sha256, third.manifest_sha256);
assert_ne!(encode(&first).unwrap(), encode(&third).unwrap());
}
#[test]
fn serialization_is_deterministic_and_rejects_corruption() {
let samples = training();
let cohort = key("cohort");
let model = fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: config(GEMINI_11),
})
.unwrap();
let bytes = encode(&model).unwrap();
let decoded = decode(&bytes).unwrap();
assert_eq!(bytes, encode(&decoded).unwrap());
assert_eq!(encode(&model).unwrap(), encode(&model).unwrap());
assert_eq!(decode(b"{}").err().unwrap(), ModelError::MalformedArtifact);
let mut value: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
value["version"] = serde_json::json!(2);
assert_eq!(
decode(&serde_json::to_vec(&value).unwrap()).err().unwrap(),
ModelError::MalformedArtifact
);
}
#[test]
fn rejects_bad_samples_and_configuration() {
let cohort = key("cohort");
let mut mixed = training();
mixed[0].cohort_id = key("other");
assert_eq!(
fit(FitInput {
cohort_id: &cohort,
samples: &mixed,
config: config(GEMINI_11),
})
.err()
.unwrap(),
ModelError::CohortMismatch
);
let mut duplicate = training();
duplicate[1].sample_id = duplicate[0].sample_id.clone();
assert_eq!(
fit(FitInput {
cohort_id: &cohort,
samples: &duplicate,
config: config(GEMINI_11),
})
.err()
.unwrap(),
ModelError::DuplicateSampleId
);
let same = vec![
sample("CCCCCCC1", "cohort", "alice", 10),
sample("CCCCCCC2", "cohort", "alice", 10),
];
assert_eq!(
fit(FitInput {
cohort_id: &cohort,
samples: &same,
config: config(GEMINI_11),
})
.err()
.unwrap(),
ModelError::ZeroVariance
);
let samples = training();
for bad in [
ModelConfig {
components: 0,
..config(GEMINI_11)
},
ModelConfig {
components: 5,
..config(GEMINI_11)
},
ModelConfig {
relevance: 0.0,
..config(GEMINI_11)
},
ModelConfig {
variance_floor: f64::NAN,
..config(GEMINI_11)
},
ModelConfig {
absolute_threshold: f64::INFINITY,
..config(GEMINI_11)
},
] {
assert_eq!(
fit(FitInput {
cohort_id: &cohort,
samples: &samples,
config: bad,
})
.err()
.unwrap(),
ModelError::InvalidConfig
);
}
}
#[test]
fn rejects_nonconverged_gmm() {
let rows = vec![
vec![-5.0],
vec![-4.0],
vec![-1.0],
vec![2.0],
vec![3.0],
vec![10.0],
];
assert_eq!(
fit_ubm_with_limits(&rows, 2, 0.01, 1, f64::MIN_POSITIVE)
.err()
.unwrap(),
ModelError::Nonconverged
);
}
}