Skip to main content

uqa_ml/
training.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7use serde::{Deserialize, Serialize};
8
9use crate::backend::{try_filled_vec, try_vec_with_capacity, MLError, MLResult};
10use crate::model::{DeepLayerSpec, DeepModel, GatingSpec};
11
12#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
13pub struct TrainingExample {
14    pub features: Vec<f64>,
15    pub label: usize,
16}
17
18#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
19pub struct TrainingSet {
20    pub examples: Vec<TrainingExample>,
21    #[serde(default, skip_serializing_if = "Option::is_none")]
22    pub class_count: Option<usize>,
23}
24
25#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
26pub struct LearnOptions {
27    #[serde(default)]
28    pub alpha: f64,
29    #[serde(default)]
30    pub gating: GatingSpec,
31}
32
33impl Default for LearnOptions {
34    fn default() -> Self {
35        Self {
36            alpha: 0.0,
37            gating: GatingSpec::None,
38        }
39    }
40}
41
42#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
43pub struct TrainingReport {
44    pub examples: usize,
45    pub feature_dimensions: usize,
46    pub class_count: usize,
47}
48
49#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
50pub struct DeepLearnOutput {
51    pub model: DeepModel,
52    pub report: TrainingReport,
53}
54
55/// Analytical nearest-centroid softmax training.
56///
57/// For each class this estimates the class centroid and emits a single
58/// dense-softmax classifier with logits equivalent to the linear part of
59/// negative squared Euclidean distance plus log-prior:
60///
61/// `logit_c(x) = centroid_c dot x - 0.5 * ||centroid_c||^2 + log prior_c`.
62pub fn deep_learn(training_set: &TrainingSet, options: &LearnOptions) -> MLResult<DeepLearnOutput> {
63    let (dims, class_count) = validate_training_shape(training_set, options)?;
64    let (counts, sums) = accumulate_training_classes(training_set, dims, class_count)?;
65    let (weights, bias) = classifier_parameters(training_set.examples.len(), &counts, &sums)?;
66
67    let model = DeepModel {
68        layers: vec![
69            DeepLayerSpec::Input { dimensions: dims },
70            DeepLayerSpec::Dense {
71                weights,
72                bias,
73                output_channels: class_count,
74                input_channels: dims,
75            },
76            DeepLayerSpec::Softmax,
77        ],
78        alpha: options.alpha,
79        gating: options.gating,
80    };
81    Ok(DeepLearnOutput {
82        model,
83        report: TrainingReport {
84            examples: training_set.examples.len(),
85            feature_dimensions: dims,
86            class_count,
87        },
88    })
89}
90
91fn validate_training_shape(
92    training_set: &TrainingSet,
93    options: &LearnOptions,
94) -> MLResult<(usize, usize)> {
95    let Some(first) = training_set.examples.first() else {
96        return Err(MLError::InvalidTrainingSet(
97            "deep_learn requires at least one training example".into(),
98        ));
99    };
100    let dims = first.features.len();
101    if dims == 0 {
102        return Err(MLError::InvalidTrainingSet(
103            "deep_learn requires non-empty feature vectors".into(),
104        ));
105    }
106    if training_set
107        .examples
108        .iter()
109        .any(|example| example.features.len() != dims)
110    {
111        return Err(MLError::InvalidTrainingSet(
112            "deep_learn requires all feature vectors to have the same dimension".into(),
113        ));
114    }
115    for (row, example) in training_set.examples.iter().enumerate() {
116        if let Some((column, value)) = example
117            .features
118            .iter()
119            .enumerate()
120            .find(|(_, value)| !value.is_finite())
121        {
122            return Err(MLError::InvalidTrainingSet(format!(
123                "training feature [{row}][{column}] must be finite, got {value}"
124            )));
125        }
126    }
127    if !options.alpha.is_finite() {
128        return Err(MLError::InvalidTrainingSet(format!(
129            "training alpha must be finite, got {}",
130            options.alpha
131        )));
132    }
133    let inferred_classes = training_set
134        .examples
135        .iter()
136        .map(|example| example.label)
137        .max()
138        .map(|label| {
139            label.checked_add(1).ok_or_else(|| {
140                MLError::InvalidTrainingSet("training label exceeds the usize range".into())
141            })
142        })
143        .transpose()?
144        .unwrap_or(0);
145    let class_count = training_set.class_count.unwrap_or(inferred_classes);
146    if class_count == 0 {
147        return Err(MLError::InvalidTrainingSet(
148            "deep_learn requires at least one class".into(),
149        ));
150    }
151    if class_count > training_set.examples.len() {
152        return Err(MLError::InvalidTrainingSet(format!(
153            "class_count={class_count} exceeds the number of training examples and necessarily contains an empty class"
154        )));
155    }
156    if training_set
157        .examples
158        .iter()
159        .any(|example| example.label >= class_count)
160    {
161        return Err(MLError::InvalidTrainingSet(format!(
162            "training label is outside class_count={class_count}"
163        )));
164    }
165
166    class_count.checked_mul(dims).ok_or_else(|| {
167        MLError::InvalidTrainingSet("training matrix dimensions overflow usize".into())
168    })?;
169    Ok((dims, class_count))
170}
171
172fn accumulate_training_classes(
173    training_set: &TrainingSet,
174    dims: usize,
175    class_count: usize,
176) -> MLResult<(Vec<usize>, Vec<Vec<f64>>)> {
177    let mut counts = try_filled_vec(class_count, 0usize, "training class counts")?;
178    let mut sums = try_vec_with_capacity(class_count, "training class feature sums")?;
179    for class in 0..class_count {
180        sums.push(try_filled_vec(
181            dims,
182            0.0f64,
183            &format!("training feature sums for class {class}"),
184        )?);
185    }
186    for example in &training_set.examples {
187        counts[example.label] = counts[example.label].checked_add(1).ok_or_else(|| {
188            MLError::InvalidTrainingSet(format!(
189                "training example count for class {} overflows usize",
190                example.label
191            ))
192        })?;
193        for (i, value) in example.features.iter().enumerate() {
194            let sum = sums[example.label][i] + value;
195            if !sum.is_finite() {
196                return Err(MLError::InvalidTrainingSet(format!(
197                    "training feature sum for class {}, column {i} is non-finite",
198                    example.label
199                )));
200            }
201            sums[example.label][i] = sum;
202        }
203    }
204    if let Some(empty_class) = counts.iter().position(|count| *count == 0) {
205        return Err(MLError::InvalidTrainingSet(format!(
206            "class {empty_class} has no training examples"
207        )));
208    }
209    Ok((counts, sums))
210}
211
212fn classifier_parameters(
213    example_count: usize,
214    counts: &[usize],
215    sums: &[Vec<f64>],
216) -> MLResult<(Vec<f64>, Vec<f64>)> {
217    let dims = sums.first().map_or(0, Vec::len);
218    let weight_count = counts.len().checked_mul(dims).ok_or_else(|| {
219        MLError::InvalidTrainingSet("trained weight count overflows usize".into())
220    })?;
221    let mut weights = try_vec_with_capacity(weight_count, "trained classifier weights")?;
222    let mut bias = try_vec_with_capacity(counts.len(), "trained classifier bias")?;
223    let total = usize_to_f64_exact(example_count, "training example count")?;
224    for class in 0..counts.len() {
225        let class_examples = usize_to_f64_exact(counts[class], "class example count")?;
226        let inv_count = 1.0 / class_examples;
227        let mut centroid = try_vec_with_capacity(dims, "training class centroid")?;
228        centroid.extend(sums[class].iter().map(|value| value * inv_count));
229        let norm_sq: f64 = centroid.iter().map(|value| value * value).sum();
230        if !norm_sq.is_finite() {
231            return Err(MLError::InvalidTrainingSet(format!(
232                "centroid norm for class {class} is non-finite"
233            )));
234        }
235        weights.extend_from_slice(&centroid);
236        let class_bias = -0.5 * norm_sq + (class_examples / total).ln();
237        if !class_bias.is_finite() {
238            return Err(MLError::InvalidTrainingSet(format!(
239                "trained bias for class {class} is non-finite"
240            )));
241        }
242        bias.push(class_bias);
243    }
244    Ok((weights, bias))
245}
246
247fn usize_to_f64_exact(value: usize, context: &str) -> MLResult<f64> {
248    const MAX_EXACT_INTEGER: u64 = 9_007_199_254_740_992;
249    let value = u64::try_from(value)
250        .map_err(|_| MLError::InvalidTrainingSet(format!("{context} exceeds the u64 bridge")))?;
251    if value > MAX_EXACT_INTEGER {
252        return Err(MLError::InvalidTrainingSet(format!(
253            "{context} exceeds f64's exact integer range"
254        )));
255    }
256    Ok(value as f64)
257}
258
259#[cfg(test)]
260mod tests {
261    use super::*;
262    use crate::backend::{CPUBackend, MLBackend};
263
264    #[test]
265    fn centroid_training_separates_two_classes() {
266        let training_set = TrainingSet {
267            examples: vec![
268                TrainingExample {
269                    features: vec![2.0, 0.0],
270                    label: 0,
271                },
272                TrainingExample {
273                    features: vec![3.0, 0.0],
274                    label: 0,
275                },
276                TrainingExample {
277                    features: vec![0.0, 2.0],
278                    label: 1,
279                },
280                TrainingExample {
281                    features: vec![0.0, 3.0],
282                    label: 1,
283                },
284            ],
285            class_count: None,
286        };
287        let output = deep_learn(&training_set, &LearnOptions::default()).unwrap();
288        assert_eq!(output.report.class_count, 2);
289
290        let backend = CPUBackend;
291        let (_, probs) = backend
292            .predict_features(&output.model, &[(1, vec![4.0, 0.0]), (2, vec![0.0, 4.0])])
293            .unwrap();
294        assert!(probs[&1][0] > probs[&1][1], "{probs:?}");
295        assert!(probs[&2][1] > probs[&2][0], "{probs:?}");
296    }
297
298    #[test]
299    fn training_rejects_non_finite_features_and_alpha() {
300        let non_finite_feature = TrainingSet {
301            examples: vec![TrainingExample {
302                features: vec![f64::INFINITY],
303                label: 0,
304            }],
305            class_count: Some(1),
306        };
307        let error = deep_learn(&non_finite_feature, &LearnOptions::default())
308            .expect_err("non-finite input must be rejected");
309        assert!(error.to_string().contains("must be finite"));
310
311        let valid = TrainingSet {
312            examples: vec![TrainingExample {
313                features: vec![1.0],
314                label: 0,
315            }],
316            class_count: Some(1),
317        };
318        let error = deep_learn(
319            &valid,
320            &LearnOptions {
321                alpha: f64::NAN,
322                ..LearnOptions::default()
323            },
324        )
325        .expect_err("non-finite alpha must be rejected");
326        assert!(error.to_string().contains("alpha must be finite"));
327    }
328
329    #[test]
330    fn impossible_class_counts_fail_before_allocation() {
331        let training_set = TrainingSet {
332            examples: vec![TrainingExample {
333                features: vec![1.0],
334                label: 0,
335            }],
336            class_count: Some(usize::MAX),
337        };
338        let error = deep_learn(&training_set, &LearnOptions::default())
339            .expect_err("an impossible class count must not trigger a huge allocation");
340        assert!(error
341            .to_string()
342            .contains("exceeds the number of training examples"));
343    }
344}