1use 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
55pub 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(¢roid);
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}