1use crate::error::{
12 RillError, checked_finite_add, checked_increment, ensure_finite, validate_features,
13};
14use crate::loss::log_loss::sigmoid;
15#[cfg(feature = "serde")]
16use crate::persistence::ValidateState;
17use crate::traits::OnlineBinaryClassifier;
18
19#[derive(Debug, Clone)]
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22#[non_exhaustive]
23pub struct NaiveBayesConfig {
24 pub alpha: f64,
28}
29
30impl Default for NaiveBayesConfig {
31 fn default() -> Self {
32 Self { alpha: 1.0 }
33 }
34}
35
36fn validate_config(config: &NaiveBayesConfig) -> Result<(), RillError> {
37 ensure_finite("alpha", config.alpha)?;
38 if config.alpha <= 0.0 {
39 return Err(RillError::InvalidParameter {
40 name: "alpha",
41 value: config.alpha,
42 });
43 }
44 Ok(())
45}
46
47fn approx_equal_non_negative(a: f64, b: f64, atol: f64, rtol: f64) -> bool {
63 if !a.is_finite() || !b.is_finite() {
64 return false;
65 }
66 let abs_diff = (a - b).abs();
67 if abs_diff <= atol {
68 return true;
69 }
70 let larger = a.abs().max(b.abs());
71 abs_diff <= rtol * larger
72}
73
74const MULTINOMIAL_TOTAL_ATOL: f64 = 1e-9;
83const MULTINOMIAL_TOTAL_RTOL: f64 = 1e-6;
84
85fn ensure_log_domain(field: &'static str, value: f64) -> Result<(), RillError> {
93 if value.is_nan() || (value.is_infinite() && value > 0.0) {
94 return Err(RillError::NonFiniteValue { field, value });
95 }
96 Ok(())
97}
98
99fn checked_log_add(current: f64, delta: f64, field: &'static str) -> Result<f64, RillError> {
104 let value = current + delta;
105 ensure_log_domain(field, value)?;
106 Ok(value)
107}
108
109fn validate_non_negative(feature_count: usize, features: &[f64]) -> Result<(), RillError> {
111 validate_features(feature_count, features)?;
112 for &x in features {
113 if x < 0.0 {
114 return Err(RillError::InvalidParameter {
115 name: "feature",
116 value: x,
117 });
118 }
119 }
120 Ok(())
121}
122
123#[derive(Debug, Clone)]
129#[cfg_attr(feature = "serde", derive(serde::Serialize))]
130struct GaussianClassStats {
131 counts: Vec<u64>,
132 means: Vec<f64>,
133 m2s: Vec<f64>,
134 class_count: u64,
135}
136
137impl GaussianClassStats {
138 fn new(feature_count: usize) -> Self {
139 Self {
140 counts: vec![0; feature_count],
141 means: vec![0.0; feature_count],
142 m2s: vec![0.0; feature_count],
143 class_count: 0,
144 }
145 }
146
147 fn variance(&self, idx: usize) -> f64 {
148 if self.counts[idx] < 2 {
149 0.0
150 } else {
151 self.m2s[idx] / self.counts[idx] as f64
152 }
153 }
154
155 fn reset(&mut self) {
156 self.counts.fill(0);
157 self.means.fill(0.0);
158 self.m2s.fill(0.0);
159 self.class_count = 0;
160 }
161
162 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
167 fn validate_invariants(&self) -> Result<(), RillError> {
168 let n = self.counts.len();
169 if self.means.len() != n || self.m2s.len() != n {
170 return Err(RillError::InvalidState(
171 "gaussian class stats: counts/means/m2s length mismatch".to_owned(),
172 ));
173 }
174 for &c in &self.counts {
177 if c != self.class_count {
178 return Err(RillError::InvalidState(format!(
179 "gaussian class stats: feature count {c} != class_count {}",
180 self.class_count
181 )));
182 }
183 }
184 for &m in &self.means {
185 ensure_finite("gaussian mean", m)?;
186 }
187 for &m2 in &self.m2s {
188 ensure_finite("gaussian m2", m2)?;
189 if m2 < 0.0 {
190 return Err(RillError::InvalidState(format!(
191 "gaussian m2 must be non-negative, got {m2}"
192 )));
193 }
194 }
195 if self.class_count == 0 {
197 for &m in &self.means {
198 if m != 0.0 {
199 return Err(RillError::InvalidState(format!(
200 "gaussian class_count=0 but mean={m} (must be 0)"
201 )));
202 }
203 }
204 for &m2 in &self.m2s {
205 if m2 != 0.0 {
206 return Err(RillError::InvalidState(format!(
207 "gaussian class_count=0 but m2={m2} (must be 0)"
208 )));
209 }
210 }
211 }
212 if self.class_count == 1 {
214 for &m2 in &self.m2s {
215 if m2 != 0.0 {
216 return Err(RillError::InvalidState(format!(
217 "gaussian class_count=1 but m2={m2} (must be 0)"
218 )));
219 }
220 }
221 }
222 Ok(())
223 }
224}
225
226#[cfg(feature = "serde")]
227impl<'de> serde::Deserialize<'de> for GaussianClassStats {
228 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
229 where
230 D: serde::Deserializer<'de>,
231 {
232 #[derive(serde::Deserialize)]
233 struct State {
234 counts: Vec<u64>,
235 means: Vec<f64>,
236 m2s: Vec<f64>,
237 class_count: u64,
238 }
239 let s = State::deserialize(deserializer)?;
240 let stats = GaussianClassStats {
241 counts: s.counts,
242 means: s.means,
243 m2s: s.m2s,
244 class_count: s.class_count,
245 };
246 stats
247 .validate_invariants()
248 .map_err(serde::de::Error::custom)?;
249 Ok(stats)
250 }
251}
252
253#[derive(Debug, Clone)]
272#[cfg_attr(feature = "serde", derive(serde::Serialize))]
273pub struct GaussianNaiveBayes {
274 feature_count: usize,
275 config: NaiveBayesConfig,
276 class_false: GaussianClassStats,
277 class_true: GaussianClassStats,
278 samples_seen: u64,
279}
280
281impl GaussianNaiveBayes {
282 pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
286 validate_config(&config)?;
287 if feature_count == 0 {
288 return Err(RillError::EmptyFeatures);
289 }
290 Ok(Self {
291 feature_count,
292 config,
293 class_false: GaussianClassStats::new(feature_count),
294 class_true: GaussianClassStats::new(feature_count),
295 samples_seen: 0,
296 })
297 }
298
299 pub const fn alpha(&self) -> f64 {
301 self.config.alpha
302 }
303
304 fn gaussian_log_pdf(x: f64, mean: f64, variance: f64) -> f64 {
311 if variance <= 0.0 {
312 return 0.0;
313 }
314 let sigma = variance.sqrt();
315 -0.5 * ((x - mean) / sigma).powi(2) - sigma.ln() - 0.5 * (2.0 * std::f64::consts::PI).ln()
316 }
317
318 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
320 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
321 if self.feature_count == 0 {
322 return Err(RillError::EmptyFeatures);
323 }
324 validate_config(&self.config)?;
325 if self.class_false.counts.len() != self.feature_count
327 || self.class_false.means.len() != self.feature_count
328 || self.class_false.m2s.len() != self.feature_count
329 {
330 return Err(RillError::InvalidState(
331 "gaussian class_false vector length != feature_count".to_owned(),
332 ));
333 }
334 if self.class_true.counts.len() != self.feature_count
335 || self.class_true.means.len() != self.feature_count
336 || self.class_true.m2s.len() != self.feature_count
337 {
338 return Err(RillError::InvalidState(
339 "gaussian class_true vector length != feature_count".to_owned(),
340 ));
341 }
342 self.class_false.validate_invariants()?;
343 self.class_true.validate_invariants()?;
344 let total_class = self
345 .class_false
346 .class_count
347 .checked_add(self.class_true.class_count)
348 .ok_or_else(|| {
349 RillError::InvalidState(format!(
350 "gaussian class_false({}) + class_true({}) overflow",
351 self.class_false.class_count, self.class_true.class_count
352 ))
353 })?;
354 if total_class != self.samples_seen {
355 return Err(RillError::InvalidState(format!(
356 "gaussian samples_seen={} != class_false({}) + class_true({})",
357 self.samples_seen, self.class_false.class_count, self.class_true.class_count
358 )));
359 }
360 Ok(())
361 }
362}
363
364#[cfg(feature = "serde")]
365impl<'de> serde::Deserialize<'de> for GaussianNaiveBayes {
366 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
367 where
368 D: serde::Deserializer<'de>,
369 {
370 #[derive(serde::Deserialize)]
371 struct State {
372 feature_count: usize,
373 config: NaiveBayesConfig,
374 class_false: GaussianClassStats,
375 class_true: GaussianClassStats,
376 samples_seen: u64,
377 }
378 let s = State::deserialize(deserializer)?;
379 let model = GaussianNaiveBayes {
380 feature_count: s.feature_count,
381 config: s.config,
382 class_false: s.class_false,
383 class_true: s.class_true,
384 samples_seen: s.samples_seen,
385 };
386 model
387 .validate_invariants()
388 .map_err(serde::de::Error::custom)?;
389 Ok(model)
390 }
391}
392
393impl OnlineBinaryClassifier for GaussianNaiveBayes {
394 fn feature_count(&self) -> usize {
395 self.feature_count
396 }
397
398 fn samples_seen(&self) -> u64 {
399 self.samples_seen
400 }
401
402 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
403 validate_features(self.feature_count, features)?;
404
405 if self.samples_seen == 0 {
406 return Ok(0.5);
407 }
408
409 let count_true = self.class_true.class_count as f64;
410 let count_false = self.class_false.class_count as f64;
411 let total = count_true + count_false;
412
413 let log_prior_true = (count_true / total).ln();
414 ensure_log_domain("nb_gaussian_log_prior_true", log_prior_true)?;
415 let log_prior_false = (count_false / total).ln();
416 ensure_log_domain("nb_gaussian_log_prior_false", log_prior_false)?;
417
418 let mut log_likelihood_true = 0.0;
419 let mut log_likelihood_false = 0.0;
420
421 for (i, &x) in features.iter().enumerate() {
422 let ll_true =
423 Self::gaussian_log_pdf(x, self.class_true.means[i], self.class_true.variance(i));
424 ensure_log_domain("nb_gaussian_log_likelihood_true", ll_true)?;
425 log_likelihood_true = checked_log_add(
426 log_likelihood_true,
427 ll_true,
428 "nb_gaussian_log_likelihood_true",
429 )?;
430 let ll_false =
431 Self::gaussian_log_pdf(x, self.class_false.means[i], self.class_false.variance(i));
432 ensure_log_domain("nb_gaussian_log_likelihood_false", ll_false)?;
433 log_likelihood_false = checked_log_add(
434 log_likelihood_false,
435 ll_false,
436 "nb_gaussian_log_likelihood_false",
437 )?;
438 }
439
440 let log_p_true = checked_log_add(
441 log_prior_true,
442 log_likelihood_true,
443 "nb_gaussian_log_p_true",
444 )?;
445 let log_p_false = checked_log_add(
446 log_prior_false,
447 log_likelihood_false,
448 "nb_gaussian_log_p_false",
449 )?;
450
451 let log_odds = log_p_true - log_p_false;
452 if log_odds.is_nan() {
456 return Err(RillError::NonFiniteValue {
457 field: "nb_gaussian_log_odds",
458 value: log_odds,
459 });
460 }
461
462 let probability = sigmoid(log_odds);
463 ensure_finite("nb_gaussian_probability", probability)?;
464 if !(0.0..=1.0).contains(&probability) {
465 return Err(RillError::InvalidProbability(probability));
466 }
467 Ok(probability)
468 }
469
470 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
471 validate_features(self.feature_count, features)?;
472
473 let stats = if target {
477 &self.class_true
478 } else {
479 &self.class_false
480 };
481 let mut next_states: Vec<(usize, u64, f64, f64)> = Vec::with_capacity(features.len());
482 for (i, &x) in features.iter().enumerate() {
483 let n = checked_increment(stats.counts[i], "feature count")?;
484 let delta = x - stats.means[i];
485 ensure_finite("mean delta", delta)?;
486 let new_mean = checked_finite_add(stats.means[i], delta / n as f64, "mean")?;
487 let delta2 = x - new_mean;
488 ensure_finite("mean delta2", delta2)?;
489 let new_m2 = checked_finite_add(stats.m2s[i], delta * delta2, "m2")?;
490 next_states.push((i, n, new_mean, new_m2));
491 }
492 let new_class_count = checked_increment(stats.class_count, "class_count")?;
493 let new_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
494
495 let stats = if target {
497 &mut self.class_true
498 } else {
499 &mut self.class_false
500 };
501 for (i, n, new_mean, new_m2) in next_states {
502 stats.counts[i] = n;
503 stats.means[i] = new_mean;
504 stats.m2s[i] = new_m2;
505 }
506 stats.class_count = new_class_count;
507 self.samples_seen = new_samples_seen;
508 Ok(())
509 }
510
511 fn reset(&mut self) {
512 self.class_false.reset();
513 self.class_true.reset();
514 self.samples_seen = 0;
515 }
516}
517
518#[derive(Debug, Clone)]
540#[cfg_attr(feature = "serde", derive(serde::Serialize))]
541pub struct BernoulliNaiveBayes {
542 feature_count: usize,
543 config: NaiveBayesConfig,
544 feature_true_counts_false: Vec<u64>,
545 feature_true_counts_true: Vec<u64>,
546 class_false_count: u64,
547 class_true_count: u64,
548 samples_seen: u64,
549}
550
551impl BernoulliNaiveBayes {
552 pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
556 validate_config(&config)?;
557 if feature_count == 0 {
558 return Err(RillError::EmptyFeatures);
559 }
560 Ok(Self {
561 feature_count,
562 config,
563 feature_true_counts_false: vec![0; feature_count],
564 feature_true_counts_true: vec![0; feature_count],
565 class_false_count: 0,
566 class_true_count: 0,
567 samples_seen: 0,
568 })
569 }
570
571 fn log_bernoulli(x: f64, p: f64) -> f64 {
573 x * p.ln() + (1.0 - x) * (1.0 - p).ln()
574 }
575
576 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
583 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
584 if self.feature_count == 0 {
585 return Err(RillError::EmptyFeatures);
586 }
587 validate_config(&self.config)?;
588 if self.feature_true_counts_false.len() != self.feature_count
589 || self.feature_true_counts_true.len() != self.feature_count
590 {
591 return Err(RillError::InvalidState(
592 "bernoulli feature count vector length != feature_count".to_owned(),
593 ));
594 }
595 if let Some(total) = self.class_false_count.checked_add(self.class_true_count) {
596 if total != self.samples_seen {
597 return Err(RillError::InvalidState(format!(
598 "bernoulli samples_seen={} != class_false({}) + class_true({})",
599 self.samples_seen, self.class_false_count, self.class_true_count
600 )));
601 }
602 } else {
603 return Err(RillError::InvalidState(format!(
604 "bernoulli class_false({}) + class_true({}) overflow",
605 self.class_false_count, self.class_true_count
606 )));
607 }
608 for &c in &self.feature_true_counts_false {
609 if c > self.class_false_count {
610 return Err(RillError::InvalidState(format!(
611 "bernoulli feature_true_counts_false entry {c} > class_false_count {}",
612 self.class_false_count
613 )));
614 }
615 }
616 for &c in &self.feature_true_counts_true {
617 if c > self.class_true_count {
618 return Err(RillError::InvalidState(format!(
619 "bernoulli feature_true_counts_true entry {c} > class_true_count {}",
620 self.class_true_count
621 )));
622 }
623 }
624 Ok(())
625 }
626}
627
628#[cfg(feature = "serde")]
629impl<'de> serde::Deserialize<'de> for BernoulliNaiveBayes {
630 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
631 where
632 D: serde::Deserializer<'de>,
633 {
634 #[derive(serde::Deserialize)]
635 struct State {
636 feature_count: usize,
637 config: NaiveBayesConfig,
638 feature_true_counts_false: Vec<u64>,
639 feature_true_counts_true: Vec<u64>,
640 class_false_count: u64,
641 class_true_count: u64,
642 samples_seen: u64,
643 }
644 let s = State::deserialize(deserializer)?;
645 let model = BernoulliNaiveBayes {
646 feature_count: s.feature_count,
647 config: s.config,
648 feature_true_counts_false: s.feature_true_counts_false,
649 feature_true_counts_true: s.feature_true_counts_true,
650 class_false_count: s.class_false_count,
651 class_true_count: s.class_true_count,
652 samples_seen: s.samples_seen,
653 };
654 model
655 .validate_invariants()
656 .map_err(serde::de::Error::custom)?;
657 Ok(model)
658 }
659}
660
661impl OnlineBinaryClassifier for BernoulliNaiveBayes {
662 fn feature_count(&self) -> usize {
663 self.feature_count
664 }
665
666 fn samples_seen(&self) -> u64 {
667 self.samples_seen
668 }
669
670 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
671 validate_non_negative(self.feature_count, features)?;
672
673 if self.samples_seen == 0 {
674 return Ok(0.5);
675 }
676
677 let count_true = self.class_true_count as f64;
678 let count_false = self.class_false_count as f64;
679 let total = count_true + count_false;
680
681 let log_prior_true = (count_true / total).ln();
682 ensure_log_domain("nb_bernoulli_log_prior_true", log_prior_true)?;
683 let log_prior_false = (count_false / total).ln();
684 ensure_log_domain("nb_bernoulli_log_prior_false", log_prior_false)?;
685
686 let mut log_likelihood_true = 0.0;
687 let mut log_likelihood_false = 0.0;
688
689 for (i, &x) in features.iter().enumerate() {
690 let p_true = (self.feature_true_counts_true[i] as f64 + self.config.alpha)
691 / (count_true + 2.0 * self.config.alpha);
692 let p_false = (self.feature_true_counts_false[i] as f64 + self.config.alpha)
693 / (count_false + 2.0 * self.config.alpha);
694 let ll_true = Self::log_bernoulli(x, p_true);
695 ensure_log_domain("nb_bernoulli_log_likelihood_true", ll_true)?;
696 log_likelihood_true = checked_log_add(
697 log_likelihood_true,
698 ll_true,
699 "nb_bernoulli_log_likelihood_true",
700 )?;
701 let ll_false = Self::log_bernoulli(x, p_false);
702 ensure_log_domain("nb_bernoulli_log_likelihood_false", ll_false)?;
703 log_likelihood_false = checked_log_add(
704 log_likelihood_false,
705 ll_false,
706 "nb_bernoulli_log_likelihood_false",
707 )?;
708 }
709
710 let log_p_true = checked_log_add(
711 log_prior_true,
712 log_likelihood_true,
713 "nb_bernoulli_log_p_true",
714 )?;
715 let log_p_false = checked_log_add(
716 log_prior_false,
717 log_likelihood_false,
718 "nb_bernoulli_log_p_false",
719 )?;
720
721 let log_odds = log_p_true - log_p_false;
722 if log_odds.is_nan() {
725 return Err(RillError::NonFiniteValue {
726 field: "nb_bernoulli_log_odds",
727 value: log_odds,
728 });
729 }
730
731 let probability = sigmoid(log_odds);
732 ensure_finite("nb_bernoulli_probability", probability)?;
733 if !(0.0..=1.0).contains(&probability) {
734 return Err(RillError::InvalidProbability(probability));
735 }
736 Ok(probability)
737 }
738
739 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
740 validate_non_negative(self.feature_count, features)?;
741
742 let mut next_feature_counts = if target {
746 self.feature_true_counts_true.clone()
747 } else {
748 self.feature_true_counts_false.clone()
749 };
750 for (i, &x) in features.iter().enumerate() {
751 if x > 0.5 {
752 next_feature_counts[i] =
753 checked_increment(next_feature_counts[i], "feature_true_count")?;
754 }
755 }
756 let next_class_count = if target {
757 checked_increment(self.class_true_count, "class_true_count")?
758 } else {
759 checked_increment(self.class_false_count, "class_false_count")?
760 };
761 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
762
763 if target {
765 self.feature_true_counts_true = next_feature_counts;
766 self.class_true_count = next_class_count;
767 } else {
768 self.feature_true_counts_false = next_feature_counts;
769 self.class_false_count = next_class_count;
770 }
771 self.samples_seen = next_samples_seen;
772 Ok(())
773 }
774
775 fn reset(&mut self) {
776 self.feature_true_counts_false.fill(0);
777 self.feature_true_counts_true.fill(0);
778 self.class_false_count = 0;
779 self.class_true_count = 0;
780 self.samples_seen = 0;
781 }
782}
783
784#[derive(Debug, Clone)]
806#[cfg_attr(feature = "serde", derive(serde::Serialize))]
807pub struct MultinomialNaiveBayes {
808 feature_count: usize,
809 config: NaiveBayesConfig,
810 feature_sums_false: Vec<f64>,
811 feature_sums_true: Vec<f64>,
812 total_false: f64,
813 total_true: f64,
814 class_false_count: u64,
815 class_true_count: u64,
816 samples_seen: u64,
817}
818
819impl MultinomialNaiveBayes {
820 pub fn new(feature_count: usize, config: NaiveBayesConfig) -> Result<Self, RillError> {
824 validate_config(&config)?;
825 if feature_count == 0 {
826 return Err(RillError::EmptyFeatures);
827 }
828 Ok(Self {
829 feature_count,
830 config,
831 feature_sums_false: vec![0.0; feature_count],
832 feature_sums_true: vec![0.0; feature_count],
833 total_false: 0.0,
834 total_true: 0.0,
835 class_false_count: 0,
836 class_true_count: 0,
837 samples_seen: 0,
838 })
839 }
840
841 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
851 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
852 if self.feature_count == 0 {
853 return Err(RillError::EmptyFeatures);
854 }
855 validate_config(&self.config)?;
856 if self.feature_sums_false.len() != self.feature_count
857 || self.feature_sums_true.len() != self.feature_count
858 {
859 return Err(RillError::InvalidState(
860 "multinomial feature sum vector length != feature_count".to_owned(),
861 ));
862 }
863 for &s in &self.feature_sums_false {
864 ensure_finite("multinomial feature_sums_false", s)?;
865 if s < 0.0 {
866 return Err(RillError::InvalidState(format!(
867 "multinomial feature_sums_false entry {s} is negative"
868 )));
869 }
870 }
871 for &s in &self.feature_sums_true {
872 ensure_finite("multinomial feature_sums_true", s)?;
873 if s < 0.0 {
874 return Err(RillError::InvalidState(format!(
875 "multinomial feature_sums_true entry {s} is negative"
876 )));
877 }
878 }
879 ensure_finite("multinomial total_false", self.total_false)?;
880 if self.total_false < 0.0 {
881 return Err(RillError::InvalidState(format!(
882 "multinomial total_false is negative: {}",
883 self.total_false
884 )));
885 }
886 ensure_finite("multinomial total_true", self.total_true)?;
887 if self.total_true < 0.0 {
888 return Err(RillError::InvalidState(format!(
889 "multinomial total_true is negative: {}",
890 self.total_true
891 )));
892 }
893 if let Some(total) = self.class_false_count.checked_add(self.class_true_count) {
894 if total != self.samples_seen {
895 return Err(RillError::InvalidState(format!(
896 "multinomial samples_seen={} != class_false({}) + class_true({})",
897 self.samples_seen, self.class_false_count, self.class_true_count
898 )));
899 }
900 } else {
901 return Err(RillError::InvalidState(format!(
902 "multinomial class_false({}) + class_true({}) overflow",
903 self.class_false_count, self.class_true_count
904 )));
905 }
906 let summed_false: f64 = self.feature_sums_false.iter().sum();
910 let summed_true: f64 = self.feature_sums_true.iter().sum();
911 if !approx_equal_non_negative(
912 summed_false,
913 self.total_false,
914 MULTINOMIAL_TOTAL_ATOL,
915 MULTINOMIAL_TOTAL_RTOL,
916 ) {
917 return Err(RillError::InvalidState(format!(
918 "multinomial sum(feature_sums_false)={summed_false} mismatches total_false={} beyond tolerance",
919 self.total_false
920 )));
921 }
922 if !approx_equal_non_negative(
923 summed_true,
924 self.total_true,
925 MULTINOMIAL_TOTAL_ATOL,
926 MULTINOMIAL_TOTAL_RTOL,
927 ) {
928 return Err(RillError::InvalidState(format!(
929 "multinomial sum(feature_sums_true)={summed_true} mismatches total_true={} beyond tolerance",
930 self.total_true
931 )));
932 }
933 Ok(())
934 }
935}
936
937#[cfg(feature = "serde")]
938impl<'de> serde::Deserialize<'de> for MultinomialNaiveBayes {
939 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
940 where
941 D: serde::Deserializer<'de>,
942 {
943 #[derive(serde::Deserialize)]
944 struct State {
945 feature_count: usize,
946 config: NaiveBayesConfig,
947 feature_sums_false: Vec<f64>,
948 feature_sums_true: Vec<f64>,
949 total_false: f64,
950 total_true: f64,
951 class_false_count: u64,
952 class_true_count: u64,
953 samples_seen: u64,
954 }
955 let s = State::deserialize(deserializer)?;
956 let model = MultinomialNaiveBayes {
957 feature_count: s.feature_count,
958 config: s.config,
959 feature_sums_false: s.feature_sums_false,
960 feature_sums_true: s.feature_sums_true,
961 total_false: s.total_false,
962 total_true: s.total_true,
963 class_false_count: s.class_false_count,
964 class_true_count: s.class_true_count,
965 samples_seen: s.samples_seen,
966 };
967 model
968 .validate_invariants()
969 .map_err(serde::de::Error::custom)?;
970 Ok(model)
971 }
972}
973
974#[cfg(feature = "serde")]
975impl ValidateState for GaussianNaiveBayes {
976 fn validate_state(&self) -> Result<(), RillError> {
977 GaussianNaiveBayes::validate_invariants(self)
978 }
979}
980
981#[cfg(feature = "serde")]
982impl ValidateState for BernoulliNaiveBayes {
983 fn validate_state(&self) -> Result<(), RillError> {
984 BernoulliNaiveBayes::validate_invariants(self)
985 }
986}
987
988#[cfg(feature = "serde")]
989impl ValidateState for MultinomialNaiveBayes {
990 fn validate_state(&self) -> Result<(), RillError> {
991 MultinomialNaiveBayes::validate_invariants(self)
992 }
993}
994
995impl OnlineBinaryClassifier for MultinomialNaiveBayes {
996 fn feature_count(&self) -> usize {
997 self.feature_count
998 }
999
1000 fn samples_seen(&self) -> u64 {
1001 self.samples_seen
1002 }
1003
1004 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
1005 validate_non_negative(self.feature_count, features)?;
1006
1007 if self.samples_seen == 0 {
1008 return Ok(0.5);
1009 }
1010
1011 let count_true = self.class_true_count as f64;
1012 let count_false = self.class_false_count as f64;
1013 let total = count_true + count_false;
1014
1015 let log_prior_true = (count_true / total).ln();
1016 ensure_log_domain("nb_multinomial_log_prior_true", log_prior_true)?;
1017 let log_prior_false = (count_false / total).ln();
1018 ensure_log_domain("nb_multinomial_log_prior_false", log_prior_false)?;
1019
1020 let denom_true = self.total_true + self.config.alpha * self.feature_count as f64;
1021 let denom_false = self.total_false + self.config.alpha * self.feature_count as f64;
1022
1023 let mut log_likelihood_true = 0.0;
1024 let mut log_likelihood_false = 0.0;
1025
1026 for (i, &x) in features.iter().enumerate() {
1027 let p_true = (self.feature_sums_true[i] + self.config.alpha) / denom_true;
1028 let p_false = (self.feature_sums_false[i] + self.config.alpha) / denom_false;
1029 let ll_true = x * p_true.ln();
1030 ensure_log_domain("nb_multinomial_log_likelihood_true", ll_true)?;
1031 log_likelihood_true = checked_log_add(
1032 log_likelihood_true,
1033 ll_true,
1034 "nb_multinomial_log_likelihood_true",
1035 )?;
1036 let ll_false = x * p_false.ln();
1037 ensure_log_domain("nb_multinomial_log_likelihood_false", ll_false)?;
1038 log_likelihood_false = checked_log_add(
1039 log_likelihood_false,
1040 ll_false,
1041 "nb_multinomial_log_likelihood_false",
1042 )?;
1043 }
1044
1045 let log_p_true = checked_log_add(
1046 log_prior_true,
1047 log_likelihood_true,
1048 "nb_multinomial_log_p_true",
1049 )?;
1050 let log_p_false = checked_log_add(
1051 log_prior_false,
1052 log_likelihood_false,
1053 "nb_multinomial_log_p_false",
1054 )?;
1055
1056 let log_odds = log_p_true - log_p_false;
1057 if log_odds.is_nan() {
1060 return Err(RillError::NonFiniteValue {
1061 field: "nb_multinomial_log_odds",
1062 value: log_odds,
1063 });
1064 }
1065
1066 let probability = sigmoid(log_odds);
1067 ensure_finite("nb_multinomial_probability", probability)?;
1068 if !(0.0..=1.0).contains(&probability) {
1069 return Err(RillError::InvalidProbability(probability));
1070 }
1071 Ok(probability)
1072 }
1073
1074 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
1075 validate_non_negative(self.feature_count, features)?;
1076
1077 let mut next_feature_sums = if target {
1081 self.feature_sums_true.clone()
1082 } else {
1083 self.feature_sums_false.clone()
1084 };
1085 let mut next_total = if target {
1086 self.total_true
1087 } else {
1088 self.total_false
1089 };
1090 for (i, &x) in features.iter().enumerate() {
1091 next_feature_sums[i] = checked_finite_add(next_feature_sums[i], x, "feature_sum")?;
1092 next_total = checked_finite_add(next_total, x, "total")?;
1093 }
1094 let next_class_count = if target {
1095 checked_increment(self.class_true_count, "class_true_count")?
1096 } else {
1097 checked_increment(self.class_false_count, "class_false_count")?
1098 };
1099 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
1100
1101 if target {
1103 self.feature_sums_true = next_feature_sums;
1104 self.total_true = next_total;
1105 self.class_true_count = next_class_count;
1106 } else {
1107 self.feature_sums_false = next_feature_sums;
1108 self.total_false = next_total;
1109 self.class_false_count = next_class_count;
1110 }
1111 self.samples_seen = next_samples_seen;
1112 Ok(())
1113 }
1114
1115 fn reset(&mut self) {
1116 self.feature_sums_false.fill(0.0);
1117 self.feature_sums_true.fill(0.0);
1118 self.total_false = 0.0;
1119 self.total_true = 0.0;
1120 self.class_false_count = 0;
1121 self.class_true_count = 0;
1122 self.samples_seen = 0;
1123 }
1124}
1125
1126#[cfg(test)]
1127mod tests {
1128 use super::*;
1129 use rand::SeedableRng;
1130
1131 #[test]
1136 fn gaussian_cold_start_returns_0_5() {
1137 let model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1138 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1139 assert!((p - 0.5).abs() < 1e-12);
1140 }
1141
1142 #[test]
1143 fn gaussian_learn_separable_data() {
1144 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1145 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1146 for _ in 0..200 {
1147 let x1 = 2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1148 let x2 = 2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1149 model.learn(&[x1, x2], true).unwrap();
1150 let x1 = -2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1151 let x2 = -2.0 + rand::Rng::gen_range(&mut rng, -1.0..1.0);
1152 model.learn(&[x1, x2], false).unwrap();
1153 }
1154 let p_pos = model.predict_proba(&[2.0, 2.0]).unwrap();
1155 let p_neg = model.predict_proba(&[-2.0, -2.0]).unwrap();
1156 assert!(p_pos > 0.7, "p_pos = {p_pos}");
1157 assert!(p_neg < 0.3, "p_neg = {p_neg}");
1158 }
1159
1160 #[test]
1161 fn gaussian_dimension_mismatch_rejected() {
1162 let mut model = GaussianNaiveBayes::new(3, Default::default()).unwrap();
1163 assert!(model.predict_proba(&[1.0, 2.0]).is_err());
1164 assert!(model.learn(&[1.0, 2.0], true).is_err());
1165 }
1166
1167 #[test]
1168 fn gaussian_non_finite_rejected() {
1169 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1170 assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1171 assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1172 }
1173
1174 #[test]
1175 fn gaussian_reset_clears_state() {
1176 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1177 model.learn(&[1.0, 2.0], true).unwrap();
1178 model.learn(&[-1.0, -2.0], false).unwrap();
1179 model.reset();
1180 assert_eq!(model.samples_seen(), 0);
1181 assert!((model.predict_proba(&[1.0, 2.0]).unwrap() - 0.5).abs() < 1e-12);
1182 }
1183
1184 #[test]
1185 fn gaussian_invalid_alpha_rejected() {
1186 assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: 0.0 }).is_err());
1187 assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: -1.0 }).is_err());
1188 assert!(GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1189 }
1190
1191 #[test]
1192 fn gaussian_predict_does_not_update_state() {
1193 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1194 model.learn(&[1.0, 2.0], true).unwrap();
1195 let before = model.samples_seen();
1196 let _ = model.predict_proba(&[0.5, 0.5]).unwrap();
1197 assert_eq!(model.samples_seen(), before);
1198 }
1199
1200 #[cfg(feature = "serde")]
1201 #[test]
1202 fn gaussian_serde_roundtrip() {
1203 let mut model = GaussianNaiveBayes::new(2, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1204 model.learn(&[1.0, 2.0], true).unwrap();
1205 model.learn(&[1.5, 2.5], true).unwrap();
1206 model.learn(&[-1.0, -2.0], false).unwrap();
1207 model.learn(&[-1.5, -2.5], false).unwrap();
1208 let json = serde_json::to_string(&model).unwrap();
1209 let restored: GaussianNaiveBayes = serde_json::from_str(&json).unwrap();
1210 assert_eq!(restored.samples_seen(), model.samples_seen());
1211 assert_eq!(restored.feature_count(), model.feature_count());
1212 let p1 = model.predict_proba(&[0.5, 0.5]).unwrap();
1213 let p2 = restored.predict_proba(&[0.5, 0.5]).unwrap();
1214 assert!((p1 - p2).abs() < 1e-12);
1215 }
1216
1217 #[test]
1218 fn gaussian_predict_proba_in_range() {
1219 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1220 model.learn(&[1.0, 2.0], true).unwrap();
1221 model.learn(&[3.0, 4.0], true).unwrap();
1222 model.learn(&[-1.0, -2.0], false).unwrap();
1223 model.learn(&[-3.0, -4.0], false).unwrap();
1224 let p = model.predict_proba(&[0.5, 1.0]).unwrap();
1225 assert!(p > 0.0 && p < 1.0, "p = {p}");
1226 }
1227
1228 #[test]
1229 fn gaussian_zero_features_rejected() {
1230 assert!(GaussianNaiveBayes::new(0, Default::default()).is_err());
1231 }
1232
1233 #[test]
1234 fn gaussian_learns_gaussian_distribution() {
1235 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1236 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1237 for _ in 0..500 {
1238 let x1 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1239 let x2 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1240 model.learn(&[x1, x2], true).unwrap();
1241 let x1 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1242 let x2 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1243 model.learn(&[x1, x2], false).unwrap();
1244 }
1245 let mut correct = 0;
1246 let total = 100;
1247 for _ in 0..total {
1248 let x1 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1249 let x2 = 3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1250 if model.predict(&[x1, x2]).unwrap() {
1251 correct += 1;
1252 }
1253 let x1 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1254 let x2 = -3.0 + 0.5 * rand::Rng::gen_range(&mut rng, -3.0..3.0);
1255 if !model.predict(&[x1, x2]).unwrap() {
1256 correct += 1;
1257 }
1258 }
1259 let accuracy = correct as f64 / (total * 2) as f64;
1260 assert!(accuracy > 0.95, "accuracy = {accuracy}");
1261 }
1262
1263 #[test]
1264 fn gaussian_single_class_predicts_that_class() {
1265 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1266 model.learn(&[1.0, 2.0], true).unwrap();
1267 model.learn(&[1.5, 2.5], true).unwrap();
1268 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1269 assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1270 }
1271
1272 #[test]
1277 fn bernoulli_cold_start_returns_0_5() {
1278 let model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1279 let p = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1280 assert!((p - 0.5).abs() < 1e-12);
1281 }
1282
1283 #[test]
1284 fn bernoulli_learn_separable_data() {
1285 let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1286 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1287 for _ in 0..200 {
1288 let f0 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.9 {
1289 1.0
1290 } else {
1291 0.0
1292 };
1293 let f1 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.1 {
1294 1.0
1295 } else {
1296 0.0
1297 };
1298 let f2 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.5 {
1299 1.0
1300 } else {
1301 0.0
1302 };
1303 model.learn(&[f0, f1, f2], true).unwrap();
1304 let f0 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.1 {
1305 1.0
1306 } else {
1307 0.0
1308 };
1309 let f1 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.9 {
1310 1.0
1311 } else {
1312 0.0
1313 };
1314 let f2 = if rand::Rng::gen_range(&mut rng, 0.0..1.0) < 0.5 {
1315 1.0
1316 } else {
1317 0.0
1318 };
1319 model.learn(&[f0, f1, f2], false).unwrap();
1320 }
1321 let p_pos = model.predict_proba(&[1.0, 0.0, 0.0]).unwrap();
1322 let p_neg = model.predict_proba(&[0.0, 1.0, 0.0]).unwrap();
1323 assert!(p_pos > 0.7, "p_pos = {p_pos}");
1324 assert!(p_neg < 0.3, "p_neg = {p_neg}");
1325 }
1326
1327 #[test]
1328 fn bernoulli_dimension_mismatch_rejected() {
1329 let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1330 assert!(model.predict_proba(&[1.0, 0.0]).is_err());
1331 assert!(model.learn(&[1.0, 0.0], true).is_err());
1332 }
1333
1334 #[test]
1335 fn bernoulli_non_finite_rejected() {
1336 let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1337 assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1338 assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1339 }
1340
1341 #[test]
1342 fn bernoulli_reset_clears_state() {
1343 let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1344 model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1345 model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1346 model.reset();
1347 assert_eq!(model.samples_seen(), 0);
1348 assert!((model.predict_proba(&[1.0, 0.0, 1.0]).unwrap() - 0.5).abs() < 1e-12);
1349 }
1350
1351 #[test]
1352 fn bernoulli_invalid_alpha_rejected() {
1353 assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.0 }).is_err());
1354 assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: -1.0 }).is_err());
1355 assert!(BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1356 }
1357
1358 #[test]
1359 fn bernoulli_predict_does_not_update_state() {
1360 let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1361 model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1362 let before = model.samples_seen();
1363 let _ = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1364 assert_eq!(model.samples_seen(), before);
1365 }
1366
1367 #[cfg(feature = "serde")]
1368 #[test]
1369 fn bernoulli_serde_roundtrip() {
1370 let mut model = BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1371 model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1372 model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1373 model.learn(&[1.0, 1.0, 0.0], true).unwrap();
1374 let json = serde_json::to_string(&model).unwrap();
1375 let restored: BernoulliNaiveBayes = serde_json::from_str(&json).unwrap();
1376 assert_eq!(restored.samples_seen(), model.samples_seen());
1377 assert_eq!(restored.feature_count(), model.feature_count());
1378 let p1 = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1379 let p2 = restored.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1380 assert!((p1 - p2).abs() < 1e-12);
1381 }
1382
1383 #[test]
1384 fn bernoulli_predict_proba_in_range() {
1385 let mut model = BernoulliNaiveBayes::new(3, Default::default()).unwrap();
1386 model.learn(&[1.0, 0.0, 1.0], true).unwrap();
1387 model.learn(&[0.0, 1.0, 0.0], false).unwrap();
1388 let p = model.predict_proba(&[1.0, 0.0, 1.0]).unwrap();
1389 assert!(p > 0.0 && p < 1.0, "p = {p}");
1390 }
1391
1392 #[test]
1393 fn bernoulli_zero_features_rejected() {
1394 assert!(BernoulliNaiveBayes::new(0, Default::default()).is_err());
1395 }
1396
1397 #[test]
1398 fn bernoulli_rejects_negative_values() {
1399 let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1400 assert!(model.learn(&[-1.0, 0.0], true).is_err());
1401 assert!(model.predict_proba(&[-0.5, 0.0]).is_err());
1402 }
1403
1404 #[test]
1409 fn multinomial_cold_start_returns_0_5() {
1410 let model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1411 let p = model.predict_proba(&[1.0, 2.0, 3.0]).unwrap();
1412 assert!((p - 0.5).abs() < 1e-12);
1413 }
1414
1415 #[test]
1416 fn multinomial_learn_separable_data() {
1417 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1418 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1419 for _ in 0..200 {
1420 let f0 = rand::Rng::gen_range(&mut rng, 3.0..6.0);
1421 let f1 = rand::Rng::gen_range(&mut rng, 2.0..5.0);
1422 let f2 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1423 model.learn(&[f0, f1, f2], true).unwrap();
1424 let f0 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1425 let f1 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
1426 let f2 = rand::Rng::gen_range(&mut rng, 3.0..6.0);
1427 model.learn(&[f0, f1, f2], false).unwrap();
1428 }
1429 let p_pos = model.predict_proba(&[4.0, 3.0, 0.0]).unwrap();
1430 let p_neg = model.predict_proba(&[0.0, 0.0, 4.0]).unwrap();
1431 assert!(p_pos > 0.7, "p_pos = {p_pos}");
1432 assert!(p_neg < 0.3, "p_neg = {p_neg}");
1433 }
1434
1435 #[test]
1436 fn multinomial_dimension_mismatch_rejected() {
1437 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1438 assert!(model.predict_proba(&[1.0, 2.0]).is_err());
1439 assert!(model.learn(&[1.0, 2.0], true).is_err());
1440 }
1441
1442 #[test]
1443 fn multinomial_non_finite_rejected() {
1444 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1445 assert!(model.learn(&[f64::NAN, 1.0], true).is_err());
1446 assert!(model.learn(&[1.0, f64::INFINITY], true).is_err());
1447 }
1448
1449 #[test]
1450 fn multinomial_reset_clears_state() {
1451 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1452 model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1453 model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1454 model.reset();
1455 assert_eq!(model.samples_seen(), 0);
1456 assert!((model.predict_proba(&[1.0, 1.0, 1.0]).unwrap() - 0.5).abs() < 1e-12);
1457 }
1458
1459 #[test]
1460 fn multinomial_invalid_alpha_rejected() {
1461 assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.0 }).is_err());
1462 assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: -1.0 }).is_err());
1463 assert!(MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: f64::NAN }).is_err());
1464 }
1465
1466 #[test]
1467 fn multinomial_predict_does_not_update_state() {
1468 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1469 model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1470 let before = model.samples_seen();
1471 let _ = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1472 assert_eq!(model.samples_seen(), before);
1473 }
1474
1475 #[cfg(feature = "serde")]
1476 #[test]
1477 fn multinomial_serde_roundtrip() {
1478 let mut model = MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
1479 model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1480 model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1481 model.learn(&[1.0, 2.0, 1.0], true).unwrap();
1482 let json = serde_json::to_string(&model).unwrap();
1483 let restored: MultinomialNaiveBayes = serde_json::from_str(&json).unwrap();
1484 assert_eq!(restored.samples_seen(), model.samples_seen());
1485 assert_eq!(restored.feature_count(), model.feature_count());
1486 let p1 = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1487 let p2 = restored.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1488 assert!((p1 - p2).abs() < 1e-12);
1489 }
1490
1491 #[test]
1492 fn multinomial_predict_proba_in_range() {
1493 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1494 model.learn(&[2.0, 1.0, 0.0], true).unwrap();
1495 model.learn(&[0.0, 1.0, 3.0], false).unwrap();
1496 let p = model.predict_proba(&[1.0, 1.0, 0.0]).unwrap();
1497 assert!(p > 0.0 && p < 1.0, "p = {p}");
1498 }
1499
1500 #[test]
1501 fn multinomial_zero_features_rejected() {
1502 assert!(MultinomialNaiveBayes::new(0, Default::default()).is_err());
1503 }
1504
1505 #[test]
1506 fn multinomial_rejects_negative_values() {
1507 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1508 assert!(model.learn(&[-1.0, 0.0], true).is_err());
1509 assert!(model.predict_proba(&[-0.5, 0.0]).is_err());
1510 }
1511
1512 #[test]
1513 fn multinomial_handles_all_zero_features() {
1514 let mut model = MultinomialNaiveBayes::new(3, Default::default()).unwrap();
1515 model.learn(&[0.0, 0.0, 0.0], true).unwrap();
1516 model.learn(&[0.0, 0.0, 0.0], false).unwrap();
1517 let p = model.predict_proba(&[0.0, 0.0, 0.0]).unwrap();
1518 assert!((p - 0.5).abs() < 1e-12, "p = {p}");
1519 }
1520
1521 #[test]
1526 fn gaussian_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1527 let mut model = GaussianNaiveBayes::new(1, Default::default()).unwrap();
1532 model.learn(&[1.0], true).unwrap();
1533 model.learn(&[2.0], true).unwrap();
1534 model.learn(&[1.0], false).unwrap();
1535 model.learn(&[2.0], false).unwrap();
1536 let result = model.predict_proba(&[1e200]);
1537 assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1538 }
1539
1540 #[test]
1541 fn bernoulli_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1542 let mut model = BernoulliNaiveBayes::new(1, Default::default()).unwrap();
1548 for _ in 0..10 {
1549 model.learn(&[1.0], true).unwrap();
1550 model.learn(&[1.0], false).unwrap();
1551 }
1552 let result = model.predict_proba(&[1e308]);
1553 assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1554 }
1555
1556 #[test]
1557 fn multinomial_predict_proba_rejects_extreme_features_causing_nan_log_odds() {
1558 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1563 model.learn(&[1e10, 0.0], true).unwrap();
1564 model.learn(&[0.0, 1e10], false).unwrap();
1565 let result = model.predict_proba(&[1e308, 1e308]);
1566 assert!(result.is_err(), "expected Err for NaN log_odds, got Ok");
1567 }
1568
1569 #[test]
1570 fn gaussian_single_class_predict_proba_returns_0_or_1() {
1571 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1572 model.learn(&[1.0, 2.0], true).unwrap();
1573 model.learn(&[1.5, 2.5], true).unwrap();
1574 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
1575 assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1576 }
1577
1578 #[test]
1579 fn bernoulli_single_class_predict_proba_returns_0_or_1() {
1580 let mut model = BernoulliNaiveBayes::new(2, Default::default()).unwrap();
1581 model.learn(&[1.0, 0.0], true).unwrap();
1582 model.learn(&[0.0, 1.0], true).unwrap();
1583 let p = model.predict_proba(&[1.0, 0.0]).unwrap();
1584 assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1585 }
1586
1587 #[test]
1588 fn multinomial_single_class_predict_proba_returns_0_or_1() {
1589 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1590 model.learn(&[1.0, 0.0], true).unwrap();
1591 model.learn(&[0.0, 1.0], true).unwrap();
1592 let p = model.predict_proba(&[1.0, 0.0]).unwrap();
1593 assert!((p - 1.0).abs() < 1e-12, "p = {p}");
1594 }
1595
1596 #[test]
1597 fn gaussian_predict_proba_does_not_modify_state_on_error() {
1598 let mut model = GaussianNaiveBayes::new(1, Default::default()).unwrap();
1599 model.learn(&[1.0], true).unwrap();
1600 model.learn(&[2.0], true).unwrap();
1601 model.learn(&[1.0], false).unwrap();
1602 model.learn(&[2.0], false).unwrap();
1603 let before = model.samples_seen();
1604 let _ = model.predict_proba(&[1e200]);
1605 assert_eq!(model.samples_seen(), before);
1606 }
1607
1608 #[test]
1613 fn gaussian_learn_failure_leaves_state_unchanged() {
1614 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1620 model.learn(&[1.0, 1.0], true).unwrap();
1621 let before_samples = model.samples_seen();
1622 let before_class_count = model.class_true.class_count;
1623
1624 let result = model.learn(&[2.0, 1e200], true);
1625 assert!(result.is_err(), "expected overflow error");
1626
1627 assert_eq!(model.samples_seen(), before_samples);
1629 assert_eq!(model.class_true.class_count, before_class_count);
1630 assert_eq!(model.class_true.counts, vec![1, 1]);
1631 assert_eq!(model.class_true.means, vec![1.0, 1.0]);
1632 assert_eq!(model.class_true.m2s, vec![0.0, 0.0]);
1633 }
1634
1635 #[test]
1636 fn gaussian_learn_succeeds_after_failed_attempt() {
1637 let mut model = GaussianNaiveBayes::new(2, Default::default()).unwrap();
1638 model.learn(&[1.0, 1.0], true).unwrap();
1639 let _ = model.learn(&[2.0, 1e200], true);
1641 model.learn(&[2.0, 3.0], true).unwrap();
1643 assert_eq!(model.samples_seen(), 2);
1644 }
1645
1646 #[test]
1647 fn multinomial_learn_failure_leaves_state_unchanged() {
1648 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1653 let before_samples = model.samples_seen();
1654
1655 let result = model.learn(&[1e308, 1e308], true);
1656 assert!(result.is_err(), "expected overflow error");
1657
1658 assert_eq!(model.samples_seen(), before_samples);
1660 assert_eq!(model.class_true_count, 0);
1661 assert_eq!(model.feature_sums_true, vec![0.0, 0.0]);
1662 assert_eq!(model.total_true, 0.0);
1663 }
1664
1665 #[test]
1666 fn multinomial_learn_succeeds_after_failed_attempt() {
1667 let mut model = MultinomialNaiveBayes::new(2, Default::default()).unwrap();
1668 let _ = model.learn(&[1e308, 1e308], true);
1669 model.learn(&[1.0, 2.0], true).unwrap();
1671 assert_eq!(model.samples_seen(), 1);
1672 assert_eq!(model.feature_sums_true, vec![1.0, 2.0]);
1673 }
1674
1675 #[cfg(feature = "serde")]
1676 #[test]
1677 fn bernoulli_learn_failure_leaves_state_unchanged() {
1678 let json = r#"{
1686 "feature_count": 2,
1687 "config": {"alpha": 1.0},
1688 "feature_true_counts_false": [0, 0],
1689 "feature_true_counts_true": [0, 18446744073709551615],
1690 "class_false_count": 0,
1691 "class_true_count": 18446744073709551615,
1692 "samples_seen": 18446744073709551615
1693 }"#;
1694 let mut model: BernoulliNaiveBayes = serde_json::from_str(json).unwrap();
1695 let before_samples = model.samples_seen();
1696
1697 let result = model.learn(&[1.0, 1.0], true);
1698 assert!(result.is_err(), "expected counter overflow");
1699
1700 assert_eq!(model.samples_seen(), before_samples);
1702 assert_eq!(model.feature_true_counts_true, vec![0u64, u64::MAX]);
1703 assert_eq!(model.class_true_count, u64::MAX);
1704 }
1705
1706 #[cfg(feature = "serde")]
1707 #[test]
1708 fn gaussian_learn_failure_on_samples_seen_overflow_leaves_state_unchanged() {
1709 let max_minus_one = u64::MAX - 1;
1721 let json = format!(
1722 r#"{{
1723 "feature_count": 1,
1724 "config": {{"alpha": 1.0}},
1725 "class_false": {{
1726 "counts": [1],
1727 "means": [3.0],
1728 "m2s": [0.0],
1729 "class_count": 1
1730 }},
1731 "class_true": {{
1732 "counts": [{max_minus_one}],
1733 "means": [5.0],
1734 "m2s": [0.0],
1735 "class_count": {max_minus_one}
1736 }},
1737 "samples_seen": 18446744073709551615
1738 }}"#
1739 );
1740 let mut model: GaussianNaiveBayes = serde_json::from_str(&json).unwrap();
1741 let before_samples = model.samples_seen();
1742
1743 let result = model.learn(&[6.0], true);
1744 assert!(result.is_err(), "expected samples_seen overflow");
1745
1746 assert_eq!(model.samples_seen(), before_samples);
1750 assert_eq!(model.class_true.counts, vec![max_minus_one]);
1751 assert_eq!(model.class_true.means, vec![5.0]);
1752 assert_eq!(model.class_true.m2s, vec![0.0]);
1753 }
1754
1755 #[cfg(feature = "serde")]
1756 #[test]
1757 fn multinomial_learn_failure_on_class_count_overflow_leaves_state_unchanged() {
1758 let json = r#"{
1764 "feature_count": 1,
1765 "config": {"alpha": 1.0},
1766 "feature_sums_false": [0.0],
1767 "feature_sums_true": [5.0],
1768 "total_false": 0.0,
1769 "total_true": 5.0,
1770 "class_false_count": 0,
1771 "class_true_count": 18446744073709551615,
1772 "samples_seen": 18446744073709551615
1773 }"#;
1774 let mut model: MultinomialNaiveBayes = serde_json::from_str(json).unwrap();
1775 let before_samples = model.samples_seen();
1776
1777 let result = model.learn(&[3.0], true);
1778 assert!(result.is_err(), "expected class_count overflow");
1779
1780 assert_eq!(model.samples_seen(), before_samples);
1781 assert_eq!(model.feature_sums_true, vec![5.0]);
1782 assert_eq!(model.total_true, 5.0);
1783 assert_eq!(model.class_true_count, u64::MAX);
1784 }
1785
1786 #[test]
1791 fn approx_equal_non_negative_tolerances() {
1792 assert!(approx_equal_non_negative(0.0, 0.0, 1e-9, 1e-6));
1794 assert!(approx_equal_non_negative(5.0, 5.0, 1e-9, 1e-6));
1795 assert!(approx_equal_non_negative(0.0, 1e-10, 1e-9, 1e-6));
1797 assert!(!approx_equal_non_negative(0.0, 1e-8, 1e-9, 1e-6));
1798 assert!(approx_equal_non_negative(
1800 1_000_000.0,
1801 1_000_000.5,
1802 1e-9,
1803 1e-6
1804 ));
1805 assert!(!approx_equal_non_negative(
1806 1_000_000.0,
1807 1_000_002.0,
1808 1e-9,
1809 1e-6
1810 ));
1811 assert!(!approx_equal_non_negative(f64::NAN, 0.0, 1e-9, 1e-6));
1813 assert!(!approx_equal_non_negative(
1814 f64::INFINITY,
1815 f64::INFINITY,
1816 1e-9,
1817 1e-6
1818 ));
1819 assert!(!approx_equal_non_negative(1.0, 1e6, 1e-9, 1e-6));
1821 }
1822
1823 #[cfg(feature = "serde")]
1824 #[test]
1825 fn gaussian_serde_rejects_feature_vector_length_mismatch() {
1826 let json = r#"{
1828 "feature_count": 2,
1829 "config": {"alpha": 1.0},
1830 "class_false": {
1831 "counts": [0, 0],
1832 "means": [0.0, 0.0],
1833 "m2s": [0.0, 0.0],
1834 "class_count": 0
1835 },
1836 "class_true": {
1837 "counts": [0],
1838 "means": [0.0],
1839 "m2s": [0.0],
1840 "class_count": 0
1841 },
1842 "samples_seen": 0
1843 }"#;
1844 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1845 assert!(
1846 result.is_err(),
1847 "feature vector length mismatch must be rejected"
1848 );
1849 }
1850
1851 #[cfg(feature = "serde")]
1852 #[test]
1853 fn gaussian_serde_rejects_negative_m2() {
1854 let json = r#"{
1855 "feature_count": 1,
1856 "config": {"alpha": 1.0},
1857 "class_false": {
1858 "counts": [0],
1859 "means": [0.0],
1860 "m2s": [0.0],
1861 "class_count": 0
1862 },
1863 "class_true": {
1864 "counts": [2],
1865 "means": [5.0],
1866 "m2s": [-1.0],
1867 "class_count": 2
1868 },
1869 "samples_seen": 2
1870 }"#;
1871 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1872 assert!(result.is_err(), "negative m2 must be rejected");
1873 }
1874
1875 #[cfg(feature = "serde")]
1876 #[test]
1877 fn gaussian_serde_rejects_count_zero_with_nonzero_state() {
1878 let json = r#"{
1881 "feature_count": 1,
1882 "config": {"alpha": 1.0},
1883 "class_false": {
1884 "counts": [0],
1885 "means": [0.0],
1886 "m2s": [0.0],
1887 "class_count": 0
1888 },
1889 "class_true": {
1890 "counts": [0],
1891 "means": [7.0],
1892 "m2s": [0.0],
1893 "class_count": 0
1894 },
1895 "samples_seen": 0
1896 }"#;
1897 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1898 assert!(
1899 result.is_err(),
1900 "class_count=0 with non-zero mean must be rejected"
1901 );
1902 }
1903
1904 #[cfg(feature = "serde")]
1905 #[test]
1906 fn gaussian_serde_rejects_count_one_with_nonzero_m2() {
1907 let json = r#"{
1910 "feature_count": 1,
1911 "config": {"alpha": 1.0},
1912 "class_false": {
1913 "counts": [0],
1914 "means": [0.0],
1915 "m2s": [0.0],
1916 "class_count": 0
1917 },
1918 "class_true": {
1919 "counts": [1],
1920 "means": [5.0],
1921 "m2s": [0.25],
1922 "class_count": 1
1923 },
1924 "samples_seen": 1
1925 }"#;
1926 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1927 assert!(
1928 result.is_err(),
1929 "class_count=1 with non-zero m2 must be rejected"
1930 );
1931 }
1932
1933 #[cfg(feature = "serde")]
1934 #[test]
1935 fn gaussian_serde_rejects_samples_seen_mismatch() {
1936 let json = r#"{
1939 "feature_count": 1,
1940 "config": {"alpha": 1.0},
1941 "class_false": {
1942 "counts": [1],
1943 "means": [3.0],
1944 "m2s": [0.0],
1945 "class_count": 1
1946 },
1947 "class_true": {
1948 "counts": [1],
1949 "means": [5.0],
1950 "m2s": [0.0],
1951 "class_count": 1
1952 },
1953 "samples_seen": 5
1954 }"#;
1955 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1956 assert!(result.is_err(), "samples_seen mismatch must be rejected");
1957 }
1958
1959 #[cfg(feature = "serde")]
1960 #[test]
1961 fn gaussian_serde_rejects_feature_count_not_equal_class_count() {
1962 let json = r#"{
1966 "feature_count": 1,
1967 "config": {"alpha": 1.0},
1968 "class_false": {
1969 "counts": [0],
1970 "means": [0.0],
1971 "m2s": [0.0],
1972 "class_count": 0
1973 },
1974 "class_true": {
1975 "counts": [3],
1976 "means": [5.0],
1977 "m2s": [0.5],
1978 "class_count": 1
1979 },
1980 "samples_seen": 1
1981 }"#;
1982 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
1983 assert!(
1984 result.is_err(),
1985 "per-feature count != class_count must be rejected"
1986 );
1987 }
1988
1989 #[cfg(feature = "serde")]
1990 #[test]
1991 fn gaussian_serde_rejects_invalid_alpha() {
1992 let json = r#"{
1993 "feature_count": 1,
1994 "config": {"alpha": 0.0},
1995 "class_false": {
1996 "counts": [0],
1997 "means": [0.0],
1998 "m2s": [0.0],
1999 "class_count": 0
2000 },
2001 "class_true": {
2002 "counts": [0],
2003 "means": [0.0],
2004 "m2s": [0.0],
2005 "class_count": 0
2006 },
2007 "samples_seen": 0
2008 }"#;
2009 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2010 assert!(result.is_err(), "invalid alpha must be rejected");
2011 }
2012
2013 #[cfg(feature = "serde")]
2014 #[test]
2015 fn gaussian_serde_rejects_non_finite_state() {
2016 let json = r#"{
2017 "feature_count": 1,
2018 "config": {"alpha": 1.0},
2019 "class_false": {
2020 "counts": [0],
2021 "means": [0.0],
2022 "m2s": [0.0],
2023 "class_count": 0
2024 },
2025 "class_true": {
2026 "counts": [1],
2027 "means": [NaN],
2028 "m2s": [0.0],
2029 "class_count": 1
2030 },
2031 "samples_seen": 1
2032 }"#;
2033 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2034 assert!(result.is_err(), "non-finite mean must be rejected");
2035 }
2036
2037 #[cfg(feature = "serde")]
2038 #[test]
2039 fn gaussian_serde_rejects_malicious_state_without_panic() {
2040 let json = r#"{
2043 "feature_count": 3,
2044 "config": {"alpha": 1.0},
2045 "class_false": {
2046 "counts": [0],
2047 "means": [0.0],
2048 "m2s": [0.0],
2049 "class_count": 0
2050 },
2051 "class_true": {
2052 "counts": [0],
2053 "means": [0.0],
2054 "m2s": [0.0],
2055 "class_count": 0
2056 },
2057 "samples_seen": 0
2058 }"#;
2059 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2060 assert!(
2061 result.is_err(),
2062 "malicious length-mismatched state must be rejected, not panicked on"
2063 );
2064 }
2065
2066 #[cfg(feature = "serde")]
2067 #[test]
2068 fn bernoulli_serde_rejects_feature_count_above_class_count() {
2069 let json = r#"{
2073 "feature_count": 1,
2074 "config": {"alpha": 1.0},
2075 "feature_true_counts_false": [0],
2076 "feature_true_counts_true": [5],
2077 "class_false_count": 0,
2078 "class_true_count": 1,
2079 "samples_seen": 1
2080 }"#;
2081 let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2082 assert!(
2083 result.is_err(),
2084 "feature count > class_count must be rejected"
2085 );
2086 }
2087
2088 #[cfg(feature = "serde")]
2089 #[test]
2090 fn bernoulli_serde_rejects_samples_seen_mismatch() {
2091 let json = r#"{
2092 "feature_count": 1,
2093 "config": {"alpha": 1.0},
2094 "feature_true_counts_false": [0],
2095 "feature_true_counts_true": [0],
2096 "class_false_count": 1,
2097 "class_true_count": 1,
2098 "samples_seen": 5
2099 }"#;
2100 let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2101 assert!(result.is_err(), "samples_seen mismatch must be rejected");
2102 }
2103
2104 #[cfg(feature = "serde")]
2105 #[test]
2106 fn bernoulli_serde_rejects_feature_vector_length_mismatch() {
2107 let json = r#"{
2108 "feature_count": 2,
2109 "config": {"alpha": 1.0},
2110 "feature_true_counts_false": [0],
2111 "feature_true_counts_true": [0, 0],
2112 "class_false_count": 0,
2113 "class_true_count": 0,
2114 "samples_seen": 0
2115 }"#;
2116 let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2117 assert!(
2118 result.is_err(),
2119 "feature vector length mismatch must be rejected"
2120 );
2121 }
2122
2123 #[cfg(feature = "serde")]
2124 #[test]
2125 fn bernoulli_serde_rejects_invalid_alpha() {
2126 let json = r#"{
2127 "feature_count": 1,
2128 "config": {"alpha": -1.0},
2129 "feature_true_counts_false": [0],
2130 "feature_true_counts_true": [0],
2131 "class_false_count": 0,
2132 "class_true_count": 0,
2133 "samples_seen": 0
2134 }"#;
2135 let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2136 assert!(result.is_err(), "invalid alpha must be rejected");
2137 }
2138
2139 #[cfg(feature = "serde")]
2140 #[test]
2141 fn multinomial_serde_rejects_negative_feature_sum() {
2142 let json = r#"{
2143 "feature_count": 1,
2144 "config": {"alpha": 1.0},
2145 "feature_sums_false": [0.0],
2146 "feature_sums_true": [-3.0],
2147 "total_false": 0.0,
2148 "total_true": -3.0,
2149 "class_false_count": 0,
2150 "class_true_count": 1,
2151 "samples_seen": 1
2152 }"#;
2153 let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2154 assert!(result.is_err(), "negative feature sum must be rejected");
2155 }
2156
2157 #[cfg(feature = "serde")]
2158 #[test]
2159 fn multinomial_serde_rejects_total_mismatch() {
2160 let json = r#"{
2163 "feature_count": 2,
2164 "config": {"alpha": 1.0},
2165 "feature_sums_false": [0.0, 0.0],
2166 "feature_sums_true": [1.0, 2.0],
2167 "total_false": 0.0,
2168 "total_true": 100.0,
2169 "class_false_count": 0,
2170 "class_true_count": 1,
2171 "samples_seen": 1
2172 }"#;
2173 let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2174 assert!(
2175 result.is_err(),
2176 "total mismatch beyond tolerance must be rejected"
2177 );
2178 }
2179
2180 #[cfg(feature = "serde")]
2181 #[test]
2182 fn multinomial_serde_rejects_samples_seen_mismatch() {
2183 let json = r#"{
2184 "feature_count": 1,
2185 "config": {"alpha": 1.0},
2186 "feature_sums_false": [0.0],
2187 "feature_sums_true": [0.0],
2188 "total_false": 0.0,
2189 "total_true": 0.0,
2190 "class_false_count": 1,
2191 "class_true_count": 1,
2192 "samples_seen": 5
2193 }"#;
2194 let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2195 assert!(result.is_err(), "samples_seen mismatch must be rejected");
2196 }
2197
2198 #[cfg(feature = "serde")]
2199 #[test]
2200 fn multinomial_serde_accepts_tiny_total_roundoff() {
2201 let json = r#"{
2205 "feature_count": 2,
2206 "config": {"alpha": 1.0},
2207 "feature_sums_false": [0.0, 0.0],
2208 "feature_sums_true": [1.0, 2.0],
2209 "total_false": 0.0,
2210 "total_true": 3.000000000001,
2211 "class_false_count": 0,
2212 "class_true_count": 1,
2213 "samples_seen": 1
2214 }"#;
2215 let model: MultinomialNaiveBayes = serde_json::from_str(json).unwrap();
2216 assert_eq!(model.samples_seen(), 1);
2217 }
2218
2219 #[cfg(feature = "serde")]
2220 #[test]
2221 fn multinomial_serde_rejects_feature_vector_length_mismatch() {
2222 let json = r#"{
2223 "feature_count": 2,
2224 "config": {"alpha": 1.0},
2225 "feature_sums_false": [0.0, 0.0],
2226 "feature_sums_true": [0.0],
2227 "total_false": 0.0,
2228 "total_true": 0.0,
2229 "class_false_count": 0,
2230 "class_true_count": 0,
2231 "samples_seen": 0
2232 }"#;
2233 let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2234 assert!(
2235 result.is_err(),
2236 "feature vector length mismatch must be rejected"
2237 );
2238 }
2239
2240 #[cfg(feature = "serde")]
2241 #[test]
2242 fn gaussian_serde_rejects_class_count_overflow() {
2243 let json = r#"{
2246 "feature_count": 1,
2247 "config": {"alpha": 1.0},
2248 "class_false": {
2249 "counts": [1],
2250 "means": [0.0],
2251 "m2s": [0.0],
2252 "class_count": 18446744073709551615
2253 },
2254 "class_true": {
2255 "counts": [1],
2256 "means": [0.0],
2257 "m2s": [0.0],
2258 "class_count": 18446744073709551615
2259 },
2260 "samples_seen": 18446744073709551614
2261 }"#;
2262 let result: Result<GaussianNaiveBayes, _> = serde_json::from_str(json);
2263 assert!(
2264 result.is_err(),
2265 "u64 overflow on class_count sum must be rejected, not panic"
2266 );
2267 }
2268
2269 #[cfg(feature = "serde")]
2270 #[test]
2271 fn bernoulli_serde_rejects_class_count_overflow() {
2272 let json = r#"{
2273 "feature_count": 1,
2274 "config": {"alpha": 1.0},
2275 "feature_true_counts_false": [1],
2276 "feature_true_counts_true": [1],
2277 "class_false_count": 18446744073709551615,
2278 "class_true_count": 18446744073709551615,
2279 "samples_seen": 18446744073709551614
2280 }"#;
2281 let result: Result<BernoulliNaiveBayes, _> = serde_json::from_str(json);
2282 assert!(
2283 result.is_err(),
2284 "u64 overflow on class_count sum must be rejected, not panic"
2285 );
2286 }
2287
2288 #[cfg(feature = "serde")]
2289 #[test]
2290 fn multinomial_serde_rejects_class_count_overflow() {
2291 let json = r#"{
2292 "feature_count": 1,
2293 "config": {"alpha": 1.0},
2294 "feature_sums_false": [1.0],
2295 "feature_sums_true": [1.0],
2296 "total_false": 1.0,
2297 "total_true": 1.0,
2298 "class_false_count": 18446744073709551615,
2299 "class_true_count": 18446744073709551615,
2300 "samples_seen": 18446744073709551614
2301 }"#;
2302 let result: Result<MultinomialNaiveBayes, _> = serde_json::from_str(json);
2303 assert!(
2304 result.is_err(),
2305 "u64 overflow on class_count sum must be rejected, not panic"
2306 );
2307 }
2308
2309 #[cfg(feature = "serde")]
2310 #[test]
2311 fn naive_bayes_valid_roundtrip_preserves_prediction() {
2312 let probe = [0.5, 1.0, 0.0];
2317
2318 let mut g = GaussianNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2319 g.learn(&[1.0, 2.0, 0.0], true).unwrap();
2320 g.learn(&[1.5, 2.5, 1.0], true).unwrap();
2321 g.learn(&[-1.0, -2.0, 0.0], false).unwrap();
2322 g.learn(&[-1.5, -2.5, 1.0], false).unwrap();
2323 let g_json = serde_json::to_string(&g).unwrap();
2324 let g_restored: GaussianNaiveBayes = serde_json::from_str(&g_json).unwrap();
2325 assert_eq!(g_restored.samples_seen(), g.samples_seen());
2326 assert_eq!(g_restored.feature_count(), g.feature_count());
2327 let g_p1 = g.predict_proba(&probe).unwrap();
2328 let g_p2 = g_restored.predict_proba(&probe).unwrap();
2329 assert!((g_p1 - g_p2).abs() < 1e-12);
2330
2331 let mut b = BernoulliNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2332 b.learn(&[1.0, 0.0, 1.0], true).unwrap();
2333 b.learn(&[0.0, 1.0, 0.0], false).unwrap();
2334 b.learn(&[1.0, 1.0, 0.0], true).unwrap();
2335 let b_json = serde_json::to_string(&b).unwrap();
2336 let b_restored: BernoulliNaiveBayes = serde_json::from_str(&b_json).unwrap();
2337 assert_eq!(b_restored.samples_seen(), b.samples_seen());
2338 assert_eq!(b_restored.feature_count(), b.feature_count());
2339 let b_p1 = b.predict_proba(&probe).unwrap();
2340 let b_p2 = b_restored.predict_proba(&probe).unwrap();
2341 assert!((b_p1 - b_p2).abs() < 1e-12);
2342
2343 let mut m = MultinomialNaiveBayes::new(3, NaiveBayesConfig { alpha: 0.5 }).unwrap();
2344 m.learn(&[2.0, 1.0, 0.0], true).unwrap();
2345 m.learn(&[0.0, 1.0, 3.0], false).unwrap();
2346 m.learn(&[1.0, 2.0, 1.0], true).unwrap();
2347 let m_json = serde_json::to_string(&m).unwrap();
2348 let m_restored: MultinomialNaiveBayes = serde_json::from_str(&m_json).unwrap();
2349 assert_eq!(m_restored.samples_seen(), m.samples_seen());
2350 assert_eq!(m_restored.feature_count(), m.feature_count());
2351 let m_p1 = m.predict_proba(&probe).unwrap();
2352 let m_p2 = m_restored.predict_proba(&probe).unwrap();
2353 assert!((m_p1 - m_p2).abs() < 1e-12);
2354 }
2355}