1use crate::matrix::Matrix;
2use rand::{Rng, SeedableRng, rngs::StdRng};
3use serde::{Deserialize, Deserializer, Serialize, de};
4use std::{error::Error, fmt};
5
6fn sigmoid(x: &mut Matrix) {
7 x.apply(|x| 1.0 / (1.0 + (-x).exp()))
8}
9
10fn sigmoid_derivative(x: &mut Matrix) {
11 x.apply(|x| x * (1.0 - x))
12}
13
14fn tanh(x: &mut Matrix) {
15 x.apply(|x| x.tanh())
16}
17
18fn tanh_derivative(x: &mut Matrix) {
19 x.apply(|x| 1.0 - x.powi(2))
20}
21
22fn linear(_: &mut Matrix) {}
23
24fn linear_derivative(x: &mut Matrix) {
25 x.apply(|_| 1.0)
26}
27
28#[derive(Clone, Debug, Serialize, Deserialize, Default)]
29pub enum ActivationFunction {
30 #[default]
31 Sigmoid,
32 Tanh,
33 Linear,
34}
35
36impl ActivationFunction {
37 fn apply(&self, x: &mut Matrix) {
38 match self {
39 ActivationFunction::Sigmoid => sigmoid(x),
40 ActivationFunction::Tanh => tanh(x),
41 ActivationFunction::Linear => linear(x),
42 }
43 }
44
45 fn derivative(&self, x: &mut Matrix) {
46 match self {
47 ActivationFunction::Sigmoid => sigmoid_derivative(x),
48 ActivationFunction::Tanh => tanh_derivative(x),
49 ActivationFunction::Linear => linear_derivative(x),
50 }
51 }
52}
53
54#[derive(Clone, Debug, PartialEq, Eq)]
55pub enum NeuralNetworkError {
56 InvalidLayerCount { got: usize },
57 InvalidLayerSize { layer_index: usize, size: usize },
58 InputLengthMismatch { expected: usize, got: usize },
59 TargetLengthMismatch { expected: usize, got: usize },
60}
61
62impl fmt::Display for NeuralNetworkError {
63 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64 match self {
65 NeuralNetworkError::InvalidLayerCount { got } => {
66 write!(
67 f,
68 "invalid layer count: expected at least 2 layers (input and output), got {got}"
69 )
70 }
71 NeuralNetworkError::InvalidLayerSize { layer_index, size } => {
72 write!(
73 f,
74 "invalid layer size at index {layer_index}: expected a positive size, got {size}"
75 )
76 }
77 NeuralNetworkError::InputLengthMismatch { expected, got } => {
78 write!(f, "input length mismatch: expected {expected}, got {got}")
79 }
80 NeuralNetworkError::TargetLengthMismatch { expected, got } => {
81 write!(f, "target length mismatch: expected {expected}, got {got}")
82 }
83 }
84 }
85}
86
87impl Error for NeuralNetworkError {}
88
89#[derive(Clone, Debug, Default, Serialize)]
91pub struct NeuralNetwork {
92 layer_sizes: Vec<usize>,
93 weights: Vec<Matrix>,
94 biases: Vec<Matrix>,
95 learning_rate: f64,
96 activation_function: ActivationFunction,
97}
98
99#[derive(Deserialize)]
100struct NeuralNetworkRepr {
101 layer_sizes: Vec<usize>,
102 weights: Vec<Matrix>,
103 biases: Vec<Matrix>,
104 learning_rate: f64,
105 activation_function: ActivationFunction,
106}
107
108impl<'de> Deserialize<'de> for NeuralNetwork {
109 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
110 where
111 D: Deserializer<'de>,
112 {
113 let repr = NeuralNetworkRepr::deserialize(deserializer)?;
114
115 if repr.layer_sizes.len() < 2 {
116 return Err(de::Error::custom(format!(
117 "invalid layer count: expected at least 2 layers, got {}",
118 repr.layer_sizes.len()
119 )));
120 }
121
122 for (layer_index, &size) in repr.layer_sizes.iter().enumerate() {
123 if size == 0 {
124 return Err(de::Error::custom(format!(
125 "invalid layer size at index {layer_index}: expected a positive size, got {size}"
126 )));
127 }
128 }
129
130 let expected_parameter_layers = repr.layer_sizes.len() - 1;
131 if repr.weights.len() != expected_parameter_layers {
132 return Err(de::Error::custom(format!(
133 "weight layer count mismatch: expected {}, got {}",
134 expected_parameter_layers,
135 repr.weights.len()
136 )));
137 }
138 if repr.biases.len() != expected_parameter_layers {
139 return Err(de::Error::custom(format!(
140 "bias layer count mismatch: expected {}, got {}",
141 expected_parameter_layers,
142 repr.biases.len()
143 )));
144 }
145
146 for (layer_index, pair) in repr.layer_sizes.windows(2).enumerate() {
147 let fan_in = pair[0];
148 let fan_out = pair[1];
149 let weight = &repr.weights[layer_index];
150 let bias = &repr.biases[layer_index];
151
152 if weight.rows() != fan_out || weight.cols() != fan_in {
153 return Err(de::Error::custom(format!(
154 "weight shape mismatch at layer {layer_index}: expected {}x{}, got {}x{}",
155 fan_out,
156 fan_in,
157 weight.rows(),
158 weight.cols()
159 )));
160 }
161
162 if bias.rows() != fan_out || bias.cols() != 1 {
163 return Err(de::Error::custom(format!(
164 "bias shape mismatch at layer {layer_index}: expected {}x1, got {}x{}",
165 fan_out,
166 bias.rows(),
167 bias.cols()
168 )));
169 }
170 }
171
172 Ok(Self {
173 layer_sizes: repr.layer_sizes,
174 weights: repr.weights,
175 biases: repr.biases,
176 learning_rate: repr.learning_rate,
177 activation_function: repr.activation_function,
178 })
179 }
180}
181
182impl NeuralNetwork {
183 fn input_size(&self) -> usize {
184 self.layer_sizes.first().copied().unwrap_or(0)
185 }
186
187 fn output_size(&self) -> usize {
188 self.layer_sizes.last().copied().unwrap_or(0)
189 }
190
191 fn validate_input_len(&self, actual: usize) -> Result<(), NeuralNetworkError> {
192 if actual == self.input_size() {
193 Ok(())
194 } else {
195 Err(NeuralNetworkError::InputLengthMismatch {
196 expected: self.input_size(),
197 got: actual,
198 })
199 }
200 }
201
202 fn validate_target_len(&self, actual: usize) -> Result<(), NeuralNetworkError> {
203 if actual == self.output_size() {
204 Ok(())
205 } else {
206 Err(NeuralNetworkError::TargetLengthMismatch {
207 expected: self.output_size(),
208 got: actual,
209 })
210 }
211 }
212
213 pub fn new(
219 layer_sizes: Vec<usize>,
220 rng: Option<&mut StdRng>,
221 ) -> Result<Self, NeuralNetworkError> {
222 if layer_sizes.len() < 2 {
223 return Err(NeuralNetworkError::InvalidLayerCount {
224 got: layer_sizes.len(),
225 });
226 }
227
228 for (layer_index, &size) in layer_sizes.iter().enumerate() {
229 if size == 0 {
230 return Err(NeuralNetworkError::InvalidLayerSize { layer_index, size });
231 }
232 }
233
234 let rng = match rng {
235 Some(rng) => rng,
236 None => &mut StdRng::from_os_rng(),
237 };
238
239 let mut weights = Vec::with_capacity(layer_sizes.len() - 1);
240 let mut biases = Vec::with_capacity(layer_sizes.len() - 1);
241
242 for pair in layer_sizes.windows(2) {
243 let fan_in = pair[0];
244 let fan_out = pair[1];
245 let limit = (6.0 / (fan_in + fan_out) as f64).sqrt();
246
247 weights.push(Matrix::random_range(rng, fan_out, fan_in, -limit, limit));
248 biases.push(Matrix::new(fan_out, 1));
249 }
250
251 Ok(NeuralNetwork {
252 layer_sizes,
253 weights,
254 biases,
255 learning_rate: 0.01,
256 activation_function: ActivationFunction::default(),
257 })
258 }
259
260 pub fn layer_sizes(&self) -> &[usize] {
262 &self.layer_sizes
263 }
264
265 pub fn learning_rate(&self) -> f64 {
267 self.learning_rate
268 }
269
270 pub fn set_learning_rate(&mut self, learning_rate: f64) {
272 self.learning_rate = learning_rate;
273 }
274
275 pub fn activation_function(&self) -> &ActivationFunction {
277 &self.activation_function
278 }
279
280 pub fn set_activation_function(&mut self, activation_function: ActivationFunction) {
282 self.activation_function = activation_function;
283 }
284
285 pub fn predict(&self, input: Vec<f64>) -> Result<Vec<f64>, NeuralNetworkError> {
287 self.validate_input_len(input.len())?;
288
289 let mut activation = Matrix::from_col_vec(input);
290
291 for (weights, biases) in self.weights.iter().zip(self.biases.iter()) {
292 let mut layer_input = weights * &activation;
293 layer_input += biases;
294 self.activation_function.apply(&mut layer_input);
295 activation = layer_input;
296 }
297
298 Ok(activation.col(0))
299 }
300
301 pub fn train(&mut self, input: Vec<f64>, target: Vec<f64>) -> Result<(), NeuralNetworkError> {
303 self.validate_input_len(input.len())?;
304 self.validate_target_len(target.len())?;
305
306 let mut activations = Vec::with_capacity(self.layer_sizes.len());
307 let mut activation = Matrix::from_col_vec(input);
308 activations.push(activation.clone());
309
310 for (weights, biases) in self.weights.iter().zip(self.biases.iter()) {
311 let mut layer_input = weights * &activation;
312 layer_input += biases;
313 self.activation_function.apply(&mut layer_input);
314 activation = layer_input;
315 activations.push(activation.clone());
316 }
317
318 let target = Matrix::from_col_vec(target);
319 let mut errors = target;
320 errors -= activations
321 .last()
322 .expect("output layer activation is missing");
323
324 for layer_idx in (0..self.weights.len()).rev() {
325 let mut deltas = activations[layer_idx + 1].clone();
326 self.activation_function.derivative(&mut deltas);
327 deltas.hadamard_product(&errors);
328
329 let mut gradients = deltas.clone();
330 gradients *= self.learning_rate;
331
332 let prev_activation_t = activations[layer_idx].transpose();
333 let weight_deltas = &gradients * &prev_activation_t;
334
335 let weights_transposed = self.weights[layer_idx].transpose();
336 let next_errors = &weights_transposed * &deltas;
337
338 self.weights[layer_idx] += &weight_deltas;
339 self.biases[layer_idx] += &gradients;
340
341 errors = next_errors;
342 }
343
344 Ok(())
345 }
346
347 pub fn mutate(&mut self, rng: &mut StdRng, mutation_rate: f64) {
348 for weight in &mut self.weights {
349 for value in weight.data_mut().iter_mut() {
350 if rng.random::<f64>() < mutation_rate {
351 *value = rng.random_range(-1.0..1.0);
352 }
353 }
354 }
355
356 for bias in &mut self.biases {
357 for value in bias.data_mut().iter_mut() {
358 if rng.random::<f64>() < mutation_rate {
359 *value = rng.random_range(-1.0..1.0);
360 }
361 }
362 }
363 }
364}
365
366#[cfg(test)]
367pub mod nn_tests {
368 use rand::{SeedableRng, rngs::StdRng};
369 use serde_json;
370
371 fn sigmoid_scalar(value: f64) -> f64 {
372 1.0 / (1.0 + (-value).exp())
373 }
374
375 fn assert_close(actual: f64, expected: f64) {
376 assert!(
377 (actual - expected).abs() < 1e-12,
378 "expected {actual} to be within 1e-12 of {expected}"
379 );
380 }
381
382 #[test]
383 fn it_creates_a_neural_network() {
384 let m = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
385
386 assert_eq!(m.layer_sizes, vec![3, 5, 2]);
387 assert_eq!(m.input_size(), 3);
388 assert_eq!(m.output_size(), 2);
389
390 assert_eq!(m.weights.len(), 2);
391 assert_eq!(m.weights[0].rows(), 5);
392 assert_eq!(m.weights[0].cols(), 3);
393 assert_eq!(m.weights[1].rows(), 2);
394 assert_eq!(m.weights[1].cols(), 5);
395
396 assert_eq!(m.biases.len(), 2);
397 assert_eq!(m.biases[0].rows(), 5);
398 assert_eq!(m.biases[0].cols(), 1);
399 assert_eq!(m.biases[1].rows(), 2);
400 assert_eq!(m.biases[1].cols(), 1);
401 }
402
403 #[test]
404 fn it_creates_a_deep_neural_network() {
405 let m = super::NeuralNetwork::new(vec![3, 4, 4, 2], None).unwrap();
406
407 assert_eq!(m.weights.len(), 3);
408 assert_eq!(m.weights[0].rows(), 4);
409 assert_eq!(m.weights[0].cols(), 3);
410 assert_eq!(m.weights[1].rows(), 4);
411 assert_eq!(m.weights[1].cols(), 4);
412 assert_eq!(m.weights[2].rows(), 2);
413 assert_eq!(m.weights[2].cols(), 4);
414
415 assert_eq!(m.biases.len(), 3);
416 assert_eq!(m.biases[0].rows(), 4);
417 assert_eq!(m.biases[1].rows(), 4);
418 assert_eq!(m.biases[2].rows(), 2);
419 }
420
421 #[test]
422 fn it_creates_a_no_hidden_layer_network() {
423 let m = super::NeuralNetwork::new(vec![3, 2], None).unwrap();
424
425 assert_eq!(m.weights.len(), 1);
426 assert_eq!(m.biases.len(), 1);
427 assert_eq!(m.weights[0].rows(), 2);
428 assert_eq!(m.weights[0].cols(), 3);
429 assert_eq!(m.biases[0].rows(), 2);
430 assert_eq!(m.biases[0].cols(), 1);
431 }
432
433 #[test]
434 pub fn it_predicts() {
435 let m = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
436 let input = vec![0.5, 0.2, 0.1];
437 let output = m.predict(input).unwrap();
438 assert_eq!(output.len(), 2);
439 assert_ne!(output[0], output[1]);
440 }
441
442 #[test]
443 fn predict_handles_deep_and_no_hidden_architectures() {
444 let deep = super::NeuralNetwork::new(vec![3, 4, 4, 2], None).unwrap();
445 let no_hidden = super::NeuralNetwork::new(vec![3, 2], None).unwrap();
446
447 assert_eq!(deep.predict(vec![0.1, 0.2, 0.3]).unwrap().len(), 2);
448 assert_eq!(no_hidden.predict(vec![0.1, 0.2, 0.3]).unwrap().len(), 2);
449 }
450
451 #[test]
452 fn predict_linear_activation_matches_manual_multilayer_math() {
453 let mut nn = super::NeuralNetwork::new(vec![2, 2, 1], None).unwrap();
454 nn.set_activation_function(super::ActivationFunction::Linear);
455
456 nn.weights[0] = super::Matrix::from_vec(2, 2, vec![1.0, 2.0, 3.0, 4.0]);
457 nn.biases[0] = super::Matrix::from_col_vec(vec![0.5, -0.5]);
458 nn.weights[1] = super::Matrix::from_vec(1, 2, vec![2.0, -1.0]);
459 nn.biases[1] = super::Matrix::from_col_vec(vec![1.0]);
460
461 let output = nn.predict(vec![0.25, 0.75]).unwrap();
462 assert!((output[0] - 2.25).abs() < 1e-12);
463 }
464
465 #[test]
466 fn train_updates_all_layers_in_deep_network() {
467 let mut nn = super::NeuralNetwork::new(vec![2, 3, 2, 1], None).unwrap();
468 nn.set_activation_function(super::ActivationFunction::Linear);
469 nn.set_learning_rate(0.1);
470
471 nn.weights[0] = super::Matrix::from_vec(3, 2, vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6]);
472 nn.weights[1] = super::Matrix::from_vec(2, 3, vec![0.2, 0.1, 0.4, 0.3, 0.5, 0.7]);
473 nn.weights[2] = super::Matrix::from_vec(1, 2, vec![0.9, 0.8]);
474 nn.biases[0] = super::Matrix::from_col_vec(vec![0.0, 0.0, 0.0]);
475 nn.biases[1] = super::Matrix::from_col_vec(vec![0.0, 0.0]);
476 nn.biases[2] = super::Matrix::from_col_vec(vec![0.0]);
477
478 let weights_before: Vec<Vec<f64>> = nn.weights.iter().map(|w| w.data().to_vec()).collect();
479 let biases_before: Vec<Vec<f64>> = nn.biases.iter().map(|b| b.data().to_vec()).collect();
480
481 nn.train(vec![0.9, 0.1], vec![0.2]).unwrap();
482
483 for (idx, weight) in nn.weights.iter().enumerate() {
484 assert_ne!(
485 weight.data(),
486 &weights_before[idx],
487 "expected weights at layer {idx} to change"
488 );
489 }
490 for (idx, bias) in nn.biases.iter().enumerate() {
491 assert_ne!(
492 bias.data(),
493 &biases_before[idx],
494 "expected biases at layer {idx} to change"
495 );
496 }
497 }
498
499 #[test]
500 fn train_uses_downstream_delta_when_updating_hidden_layers() {
501 let mut nn = super::NeuralNetwork::new(vec![1, 1, 1], None).unwrap();
502 nn.set_learning_rate(0.5);
503
504 nn.weights[0] = super::Matrix::from_vec(1, 1, vec![0.5]);
505 nn.biases[0] = super::Matrix::from_col_vec(vec![0.0]);
506 nn.weights[1] = super::Matrix::from_vec(1, 1, vec![-0.4]);
507 nn.biases[1] = super::Matrix::from_col_vec(vec![0.1]);
508
509 let input = 1.0;
510 let target = 0.8;
511 let hidden = sigmoid_scalar(0.5 * input);
512 let output = sigmoid_scalar(-0.4 * hidden + 0.1);
513 let output_error = target - output;
514 let output_delta = output_error * output * (1.0 - output);
515 let hidden_error = -0.4 * output_delta;
516 let hidden_delta = hidden_error * hidden * (1.0 - hidden);
517
518 let expected_hidden_weight = 0.5 + 0.5 * hidden_delta * input;
519 let expected_hidden_bias = 0.0 + 0.5 * hidden_delta;
520 let expected_output_weight = -0.4 + 0.5 * output_delta * hidden;
521 let expected_output_bias = 0.1 + 0.5 * output_delta;
522
523 nn.train(vec![input], vec![target]).unwrap();
524
525 assert_close(nn.weights[0].get(0, 0), expected_hidden_weight);
526 assert_close(nn.biases[0].get(0, 0), expected_hidden_bias);
527 assert_close(nn.weights[1].get(0, 0), expected_output_weight);
528 assert_close(nn.biases[1].get(0, 0), expected_output_bias);
529 }
530
531 #[test]
532 fn mutate_honors_rate_extremes() {
533 let mut rng = StdRng::seed_from_u64(77);
534 let mut nn = super::NeuralNetwork::new(vec![3, 4, 2], Some(&mut rng)).unwrap();
535
536 let original_weights: Vec<Vec<f64>> =
537 nn.weights.iter().map(|w| w.data().to_vec()).collect();
538 let original_biases: Vec<Vec<f64>> = nn.biases.iter().map(|b| b.data().to_vec()).collect();
539
540 nn.mutate(&mut rng, 0.0);
541 for (idx, weight) in nn.weights.iter().enumerate() {
542 assert_eq!(weight.data(), &original_weights[idx]);
543 }
544 for (idx, bias) in nn.biases.iter().enumerate() {
545 assert_eq!(bias.data(), &original_biases[idx]);
546 }
547
548 nn.mutate(&mut rng, 1.0);
549 let mut any_changed = false;
550
551 for (idx, weight) in nn.weights.iter().enumerate() {
552 if weight.data() != original_weights[idx].as_slice() {
553 any_changed = true;
554 }
555 assert!(
556 weight
557 .data()
558 .iter()
559 .all(|value| *value >= -1.0 && *value < 1.0)
560 );
561 }
562 for (idx, bias) in nn.biases.iter().enumerate() {
563 if bias.data() != original_biases[idx].as_slice() {
564 any_changed = true;
565 }
566 assert!(
567 bias.data()
568 .iter()
569 .all(|value| *value >= -1.0 && *value < 1.0)
570 );
571 }
572
573 assert!(any_changed, "expected at least one parameter to change");
574 }
575
576 #[test]
577 fn it_learns_the_or_function() {
578 let mut rng = StdRng::seed_from_u64(42);
579 let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
580 nn.set_learning_rate(0.5);
581
582 let training_data = [
583 (vec![0.0, 0.0], vec![0.0]),
584 (vec![0.0, 1.0], vec![1.0]),
585 (vec![1.0, 0.0], vec![1.0]),
586 (vec![1.0, 1.0], vec![1.0]),
587 ];
588
589 for _ in 0..10_000 {
590 for (input, target) in &training_data {
591 nn.train(input.clone(), target.clone()).unwrap();
592 }
593 }
594
595 assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
596 assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
597 assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
598 assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] > 0.8);
599 }
600
601 #[test]
602 fn tanh_derivative_uses_activated_output() {
603 let mut x = crate::Matrix::from_col_vec(vec![0.5, -0.25]);
604 super::tanh_derivative(&mut x);
605
606 assert!((x.get(0, 0) - 0.75).abs() < 1e-12);
607 assert!((x.get(1, 0) - 0.9375).abs() < 1e-12);
608 }
609
610 #[test]
611 fn predict_returns_clear_error_for_wrong_input_size() {
612 let nn = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
613
614 assert_eq!(
615 nn.predict(vec![0.1, 0.2]),
616 Err(super::NeuralNetworkError::InputLengthMismatch {
617 expected: 3,
618 got: 2,
619 })
620 );
621 }
622
623 #[test]
624 fn train_returns_clear_error_for_wrong_target_size() {
625 let mut nn = super::NeuralNetwork::new(vec![3, 5, 2], None).unwrap();
626
627 assert_eq!(
628 nn.train(vec![0.1, 0.2, 0.3], vec![1.0]),
629 Err(super::NeuralNetworkError::TargetLengthMismatch {
630 expected: 2,
631 got: 1,
632 })
633 );
634 }
635
636 #[test]
637 fn new_rejects_invalid_layer_vectors() {
638 assert_eq!(
639 super::NeuralNetwork::new(vec![], None).unwrap_err(),
640 super::NeuralNetworkError::InvalidLayerCount { got: 0 }
641 );
642
643 assert_eq!(
644 super::NeuralNetwork::new(vec![3], None).unwrap_err(),
645 super::NeuralNetworkError::InvalidLayerCount { got: 1 }
646 );
647
648 assert_eq!(
649 super::NeuralNetwork::new(vec![0, 5, 2], None).unwrap_err(),
650 super::NeuralNetworkError::InvalidLayerSize {
651 layer_index: 0,
652 size: 0,
653 }
654 );
655
656 assert_eq!(
657 super::NeuralNetwork::new(vec![3, 0, 2], None).unwrap_err(),
658 super::NeuralNetworkError::InvalidLayerSize {
659 layer_index: 1,
660 size: 0,
661 }
662 );
663
664 assert_eq!(
665 super::NeuralNetwork::new(vec![3, 5, 0], None).unwrap_err(),
666 super::NeuralNetworkError::InvalidLayerSize {
667 layer_index: 2,
668 size: 0,
669 }
670 );
671 }
672
673 #[test]
674 fn new_uses_zero_biases() {
675 let nn = super::NeuralNetwork::new(vec![3, 5, 4, 2], None).unwrap();
676
677 assert!(
678 nn.biases
679 .iter()
680 .all(|bias| bias.data().iter().all(|value| *value == 0.0))
681 );
682 }
683
684 #[test]
685 fn new_uses_xavier_weight_ranges() {
686 let mut rng = StdRng::seed_from_u64(7);
687 let layer_sizes = vec![3, 5, 4, 2];
688 let nn = super::NeuralNetwork::new(layer_sizes.clone(), Some(&mut rng)).unwrap();
689
690 for (weight, pair) in nn.weights.iter().zip(layer_sizes.windows(2)) {
691 let fan_in = pair[0] as f64;
692 let fan_out = pair[1] as f64;
693 let limit = (6.0_f64 / (fan_in + fan_out)).sqrt();
694
695 assert!(
696 weight
697 .data()
698 .iter()
699 .all(|value| *value >= -limit && *value < limit)
700 );
701 }
702 }
703
704 #[test]
705 fn it_learns_the_xor_function() {
706 let mut rng = StdRng::seed_from_u64(99);
707 let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
708 nn.set_learning_rate(0.5);
709
710 let training_data = [
711 (vec![0.0, 0.0], vec![0.0]),
712 (vec![0.0, 1.0], vec![1.0]),
713 (vec![1.0, 0.0], vec![1.0]),
714 (vec![1.0, 1.0], vec![0.0]),
715 ];
716
717 for _ in 0..20_000 {
718 for (input, target) in &training_data {
719 nn.train(input.clone(), target.clone()).unwrap();
720 }
721 }
722
723 assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
724 assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
725 assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
726 assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] < 0.2);
727 }
728
729 #[test]
730 fn it_learns_the_xor_function_with_deeper_network() {
731 let mut rng = StdRng::seed_from_u64(314);
732 let mut nn = super::NeuralNetwork::new(vec![2, 4, 4, 1], Some(&mut rng)).unwrap();
733 nn.set_learning_rate(0.5);
734
735 let training_data = [
736 (vec![0.0, 0.0], vec![0.0]),
737 (vec![0.0, 1.0], vec![1.0]),
738 (vec![1.0, 0.0], vec![1.0]),
739 (vec![1.0, 1.0], vec![0.0]),
740 ];
741
742 for _ in 0..20_000 {
743 for (input, target) in &training_data {
744 nn.train(input.clone(), target.clone()).unwrap();
745 }
746 }
747
748 assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
749 assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] > 0.8);
750 assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
751 assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] < 0.2);
752 }
753
754 #[test]
755 fn perceptron_architecture_learns_linearly_separable_data() {
756 let mut rng = StdRng::seed_from_u64(202);
757 let mut nn = super::NeuralNetwork::new(vec![2, 1], Some(&mut rng)).unwrap();
758 nn.set_learning_rate(0.5);
759
760 let training_data = [
761 (vec![0.0, 0.0], vec![0.0]),
762 (vec![0.0, 1.0], vec![0.0]),
763 (vec![1.0, 0.0], vec![1.0]),
764 (vec![1.0, 1.0], vec![1.0]),
765 ];
766
767 for _ in 0..12_000 {
768 for (input, target) in &training_data {
769 nn.train(input.clone(), target.clone()).unwrap();
770 }
771 }
772
773 assert!(nn.predict(vec![0.0, 0.0]).unwrap()[0] < 0.2);
774 assert!(nn.predict(vec![0.0, 1.0]).unwrap()[0] < 0.2);
775 assert!(nn.predict(vec![1.0, 0.0]).unwrap()[0] > 0.8);
776 assert!(nn.predict(vec![1.0, 1.0]).unwrap()[0] > 0.8);
777 }
778
779 #[test]
780 fn serde_round_trip_preserves_predictions() {
781 let mut rng = StdRng::seed_from_u64(123);
782 let mut nn = super::NeuralNetwork::new(vec![2, 4, 1], Some(&mut rng)).unwrap();
783 nn.set_learning_rate(0.5);
784
785 let training_data = [
786 (vec![0.0, 0.0], vec![0.0]),
787 (vec![0.0, 1.0], vec![1.0]),
788 (vec![1.0, 0.0], vec![1.0]),
789 (vec![1.0, 1.0], vec![0.0]),
790 ];
791
792 for _ in 0..5_000 {
793 for (input, target) in &training_data {
794 nn.train(input.clone(), target.clone()).unwrap();
795 }
796 }
797
798 let probe_input = vec![0.25, 0.75];
799 let before = nn.predict(probe_input.clone()).unwrap();
800
801 let json = serde_json::to_string(&nn).unwrap();
802 let restored: super::NeuralNetwork = serde_json::from_str(&json).unwrap();
803 let after = restored.predict(probe_input).unwrap();
804
805 assert_eq!(before, after);
806 }
807
808 #[test]
809 fn serde_rejects_networks_with_invalid_matrix_shapes() {
810 let json = r#"{
811 "layer_sizes": [2, 1],
812 "weights": [
813 { "rows": 1, "cols": 2, "data": [0.25] }
814 ],
815 "biases": [
816 { "rows": 1, "cols": 1, "data": [0.0] }
817 ],
818 "learning_rate": 0.1,
819 "activation_function": "Sigmoid"
820 }"#;
821
822 let result = serde_json::from_str::<super::NeuralNetwork>(json);
823
824 assert!(
825 result.is_err(),
826 "deserialization should reject matrix data whose length does not match rows * cols"
827 );
828 }
829
830 #[test]
831 fn serde_rejects_networks_whose_shapes_do_not_match_layer_sizes() {
832 let json = r#"{
833 "layer_sizes": [2, 1],
834 "weights": [
835 { "rows": 1, "cols": 1, "data": [0.25] }
836 ],
837 "biases": [
838 { "rows": 1, "cols": 1, "data": [0.0] }
839 ],
840 "learning_rate": 0.1,
841 "activation_function": "Sigmoid"
842 }"#;
843
844 let result = serde_json::from_str::<super::NeuralNetwork>(json);
845
846 assert!(
847 result.is_err(),
848 "deserialization should reject weights whose shape does not match layer_sizes"
849 );
850 }
851}