kcode-speech-classification 0.1.1

Typed open-set speaker classification with SQLite-backed training and an NDJSON Unix-socket service.
Documentation
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());
    }
}