use crate::model::{CandidateEvidence, Error, FeatureRow, IdentifyEvidence};
use std::collections::{BTreeMap, BTreeSet};
const PRIOR_STRENGTH: f64 = 5.0;
const NUMERICAL_FEATURES: usize = 20;
const CATEGORICAL_FEATURES: usize = 4;
const NUMERICAL_FLOORS: [f64; NUMERICAL_FEATURES] = [
1.0, 1.0, 4.0, 25.0, 0.0001, 0.01, 0.25, 1.0, 1.0, 1.0, 0.01, 0.01, 0.25, 0.25, 0.01, 0.25,
0.01, 0.25, 0.25, 0.25,
];
#[derive(Clone, Debug)]
pub(crate) struct LabeledRow {
pub(crate) speaker_id: String,
pub(crate) row: FeatureRow,
}
#[derive(Default)]
struct NumericalStats {
count: usize,
mean: f64,
m2: f64,
}
impl NumericalStats {
fn add(&mut self, value: f64) {
self.count += 1;
let delta = value - self.mean;
self.mean += delta / self.count as f64;
let delta_from_new_mean = value - self.mean;
self.m2 += delta * delta_from_new_mean;
}
}
struct SpeakerStats<'a> {
speaker_id: &'a str,
numerical: [NumericalStats; NUMERICAL_FEATURES],
categorical: [BTreeMap<&'a str, usize>; CATEGORICAL_FEATURES],
count: usize,
}
impl<'a> SpeakerStats<'a> {
fn new(speaker_id: &'a str) -> Self {
Self {
speaker_id,
numerical: std::array::from_fn(|_| NumericalStats::default()),
categorical: std::array::from_fn(|_| BTreeMap::new()),
count: 0,
}
}
fn add(&mut self, row: &'a FeatureRow) {
self.count += 1;
for (stats, value) in self.numerical.iter_mut().zip(row.numerical_values()) {
stats.add(value);
}
for (counts, value) in self.categorical.iter_mut().zip(row.categorical_values()) {
*counts.entry(value).or_default() += 1;
}
}
}
pub(crate) fn score(
observations: &[LabeledRow],
query: &FeatureRow,
) -> Result<Option<IdentifyEvidence>, Error> {
if observations.is_empty() {
return Ok(None);
}
let mut population_numerical: [NumericalStats; NUMERICAL_FEATURES] =
std::array::from_fn(|_| NumericalStats::default());
let mut population_categorical: [BTreeMap<&str, usize>; CATEGORICAL_FEATURES] =
std::array::from_fn(|_| BTreeMap::new());
let mut speakers: BTreeMap<&str, SpeakerStats<'_>> = BTreeMap::new();
for observation in observations {
for (stats, value) in population_numerical
.iter_mut()
.zip(observation.row.numerical_values())
{
stats.add(value);
}
for (counts, value) in population_categorical
.iter_mut()
.zip(observation.row.categorical_values())
{
*counts.entry(value).or_default() += 1;
}
speakers
.entry(&observation.speaker_id)
.or_insert_with(|| SpeakerStats::new(&observation.speaker_id))
.add(&observation.row);
}
let population_variances: [f64; NUMERICAL_FEATURES] = std::array::from_fn(|index| {
let stats = &population_numerical[index];
(stats.m2 / stats.count as f64).max(NUMERICAL_FLOORS[index])
});
let query_numerical = query.numerical_values();
let query_categorical = query.categorical_values();
let mut background_cost = 0.0;
for index in 0..NUMERICAL_FEATURES {
background_cost += gaussian_cost(
query_numerical[index],
population_numerical[index].mean,
population_variances[index],
);
}
for index in 0..CATEGORICAL_FEATURES {
let probability = population_probability(
&population_categorical[index],
observations.len(),
query_categorical[index],
);
background_cost += categorical_cost(probability);
}
ensure_finite_cost(background_cost)?;
let mut candidates = Vec::with_capacity(speakers.len());
for speaker in speakers.values() {
let mut cost = 0.0;
for index in 0..NUMERICAL_FEATURES {
let denominator = (speaker.count - 1) as f64 + PRIOR_STRENGTH;
let variance = (speaker.numerical[index].m2
+ PRIOR_STRENGTH * population_variances[index])
/ denominator;
cost += gaussian_cost(
query_numerical[index],
speaker.numerical[index].mean,
variance,
);
}
for index in 0..CATEGORICAL_FEATURES {
let population_probability = population_probability(
&population_categorical[index],
observations.len(),
query_categorical[index],
);
let speaker_count = speaker.categorical[index]
.get(query_categorical[index])
.copied()
.unwrap_or(0);
let probability = (speaker_count as f64 + PRIOR_STRENGTH * population_probability)
/ (speaker.count as f64 + PRIOR_STRENGTH);
cost += categorical_cost(probability);
}
ensure_finite_cost(cost)?;
candidates.push(CandidateEvidence {
speaker_id: speaker.speaker_id.to_owned(),
cost,
});
}
candidates.sort_by(|left, right| {
left.cost
.total_cmp(&right.cost)
.then_with(|| left.speaker_id.cmp(&right.speaker_id))
});
let best = candidates.remove(0);
let runner_up = candidates.into_iter().next();
let absolute_gap = background_cost - best.cost;
let runner_up_gap = runner_up
.as_ref()
.map(|candidate| candidate.cost - best.cost);
let confidence_score = runner_up_gap.map_or(absolute_gap, |gap| absolute_gap.min(gap));
ensure_finite_cost(absolute_gap)?;
if let Some(gap) = runner_up_gap {
ensure_finite_cost(gap)?;
}
ensure_finite_cost(confidence_score)?;
Ok(Some(IdentifyEvidence {
best,
runner_up,
background_population_cost: background_cost,
absolute_gap,
runner_up_gap,
confidence_score,
}))
}
fn gaussian_cost(value: f64, mean: f64, variance: f64) -> f64 {
let difference = value - mean;
difference * difference / variance + variance.ln()
}
fn categorical_cost(probability: f64) -> f64 {
-2.0 * probability.ln()
}
fn population_probability(
counts: &BTreeMap<&str, usize>,
population_count: usize,
query: &str,
) -> f64 {
let categories: BTreeSet<&str> = counts.keys().copied().chain([query]).collect();
let query_count = counts.get(query).copied().unwrap_or(0);
(query_count as f64 + 1.0) / (population_count + categories.len()) as f64
}
fn ensure_finite_cost(cost: f64) -> Result<(), Error> {
if cost.is_finite() {
Ok(())
} else {
Err(Error::validation(
"row",
"feature magnitudes exceed the finite scoring range",
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Cefr;
use crate::model::tests::sample_row;
fn labeled(speaker: &str, row: FeatureRow) -> LabeledRow {
LabeledRow {
speaker_id: speaker.to_owned(),
row,
}
}
#[test]
fn numerical_evidence_selects_the_nearest_profile() {
let mut a1 = sample_row();
a1.perceived_age = 30.0;
let mut a2 = a1.clone();
a2.perceived_age = 32.0;
let mut b1 = a1.clone();
b1.perceived_age = 70.0;
let mut b2 = b1.clone();
b2.perceived_age = 72.0;
let mut query = a1.clone();
query.perceived_age = 31.0;
let evidence = score(
&[
labeled("speaker-a", a1),
labeled("speaker-a", a2),
labeled("speaker-b", b1),
labeled("speaker-b", b2),
],
&query,
)
.unwrap()
.unwrap();
assert_eq!(evidence.best.speaker_id, "speaker-a");
assert_eq!(evidence.runner_up.as_ref().unwrap().speaker_id, "speaker-b");
assert!(evidence.runner_up_gap.unwrap() > 0.0);
assert_eq!(
evidence.confidence_score,
evidence.absolute_gap.min(evidence.runner_up_gap.unwrap())
);
}
#[test]
fn cefr_ordinal_proximity_selects_the_nearest_profile() {
let mut near = sample_row();
near.cefr = Cefr::A2;
let mut far = near.clone();
far.cefr = Cefr::C2;
let mut query = near.clone();
query.cefr = Cefr::B1;
assert_eq!(near.categorical_values(), far.categorical_values());
assert_eq!(near.categorical_values(), query.categorical_values());
let evidence = score(
&[
labeled("zzz-near", near.clone()),
labeled("zzz-near", near),
labeled("aaa-far", far.clone()),
labeled("aaa-far", far),
],
&query,
)
.unwrap()
.unwrap();
assert_eq!(evidence.best.speaker_id, "zzz-near");
assert_eq!(evidence.runner_up.as_ref().unwrap().speaker_id, "aaa-far");
assert!(evidence.runner_up_gap.unwrap() > 0.0);
}
#[test]
fn laplace_categorical_evidence_selects_matching_profile() {
let mut alpha = sample_row();
alpha.accent_variety = "alpha".to_owned();
alpha.rhotic_realization = "alpha-r".to_owned();
let mut beta = alpha.clone();
beta.accent_variety = "beta".to_owned();
beta.rhotic_realization = "beta-r".to_owned();
let query = alpha.clone();
let evidence = score(
&[
labeled("alpha-speaker", alpha.clone()),
labeled("alpha-speaker", alpha),
labeled("beta-speaker", beta.clone()),
labeled("beta-speaker", beta),
],
&query,
)
.unwrap()
.unwrap();
assert_eq!(evidence.best.speaker_id, "alpha-speaker");
assert!(evidence.runner_up_gap.unwrap() > 0.0);
}
#[test]
fn includes_an_unseen_query_category_in_population_smoothing() {
let mut known = sample_row();
known.accent_variety = "known".to_owned();
let mut query = known.clone();
query.accent_variety = "unseen".to_owned();
let evidence = score(&[labeled("only", known)], &query).unwrap().unwrap();
assert_eq!(evidence.best.speaker_id, "only");
assert!(evidence.background_population_cost.is_finite());
assert!(evidence.confidence_score.is_finite());
}
#[test]
fn scores_many_profiles_deterministically() {
let mut observations = Vec::new();
for speaker in 0..200 {
for sample in 0..3 {
let mut row = sample_row();
row.perceived_age = 20.0 + speaker as f64 / 2.0 + sample as f64 / 10.0;
row.median_f0_hz = 100.0 + speaker as f64;
row.accent_variety = format!("group-{}", speaker % 8);
observations.push(labeled(&format!("speaker-{speaker:03}"), row));
}
}
let query = observations[369].row.clone();
let first = score(&observations, &query).unwrap().unwrap();
let second = score(&observations, &query).unwrap().unwrap();
assert_eq!(first, second);
assert_eq!(first.best.speaker_id, observations[369].speaker_id);
assert!(first.runner_up.is_some());
}
}