1use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
45use crate::loss::log_loss::sigmoid;
46use crate::sparse::{FeatureId, SparseFeatures};
47use crate::traits::{SparseClassifier, SparseRegressor};
48use std::collections::BTreeMap;
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
56#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
57pub enum NewFeaturePolicy {
58 #[default]
62 Reject,
63 Ignore,
67}
68
69#[derive(Debug, Clone)]
75#[cfg_attr(feature = "serde", derive(serde::Serialize))]
76pub struct FtrlConfig {
77 pub alpha: f64,
79 pub beta: f64,
81 pub l1: f64,
83 pub l2: f64,
85 pub max_features: Option<usize>,
92 pub new_feature_policy: NewFeaturePolicy,
95}
96
97impl Default for FtrlConfig {
98 fn default() -> Self {
99 Self {
100 alpha: 0.1,
101 beta: 1.0,
102 l1: 1.0,
103 l2: 1.0,
104 max_features: None,
105 new_feature_policy: NewFeaturePolicy::default(),
106 }
107 }
108}
109
110impl FtrlConfig {
111 fn validate(&self) -> Result<(), RillError> {
113 ensure_finite("alpha", self.alpha)?;
114 ensure_finite("beta", self.beta)?;
115 ensure_finite("l1", self.l1)?;
116 ensure_finite("l2", self.l2)?;
117 if self.alpha <= 0.0 {
118 return Err(RillError::InvalidParameter {
119 name: "alpha",
120 value: self.alpha,
121 });
122 }
123 if self.beta < 0.0 {
124 return Err(RillError::InvalidParameter {
125 name: "beta",
126 value: self.beta,
127 });
128 }
129 if self.l1 < 0.0 {
130 return Err(RillError::InvalidParameter {
131 name: "l1",
132 value: self.l1,
133 });
134 }
135 if self.l2 < 0.0 {
136 return Err(RillError::InvalidParameter {
137 name: "l2",
138 value: self.l2,
139 });
140 }
141 if let Some(max_features) = self.max_features
142 && max_features == 0
143 {
144 return Err(RillError::InvalidParameter {
145 name: "max_features",
146 value: 0.0,
147 });
148 }
149 Ok(())
150 }
151}
152
153#[cfg(feature = "serde")]
154impl<'de> serde::Deserialize<'de> for FtrlConfig {
155 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
156 where
157 D: serde::Deserializer<'de>,
158 {
159 #[derive(serde::Deserialize)]
160 struct FtrlConfigState {
161 alpha: f64,
162 beta: f64,
163 l1: f64,
164 l2: f64,
165 #[serde(default)]
166 max_features: Option<usize>,
167 #[serde(default)]
168 new_feature_policy: NewFeaturePolicy,
169 }
170
171 let state = FtrlConfigState::deserialize(deserializer)?;
172 let config = FtrlConfig {
173 alpha: state.alpha,
174 beta: state.beta,
175 l1: state.l1,
176 l2: state.l2,
177 max_features: state.max_features,
178 new_feature_policy: state.new_feature_policy,
179 };
180 config.validate().map_err(serde::de::Error::custom)?;
181 Ok(config)
182 }
183}
184
185#[derive(Debug, Clone, Default)]
191#[cfg_attr(feature = "serde", derive(serde::Serialize))]
192pub struct FtrlParam {
193 z: f64,
195 n: f64,
197}
198
199impl FtrlParam {
200 fn weight(&self, config: &FtrlConfig) -> f64 {
204 if self.z.abs() <= config.l1 {
205 0.0
206 } else {
207 let sign = self.z.signum();
208 let numerator = -(self.z - sign * config.l1);
209 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
210 numerator / denominator
211 }
212 }
213
214 fn intercept_weight(&self, config: &FtrlConfig) -> f64 {
219 if self.n == 0.0 {
220 0.0
221 } else {
222 let numerator = -self.z;
223 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
224 numerator / denominator
225 }
226 }
227
228 fn next_updated(
239 &self,
240 gradient: f64,
241 weight: f64,
242 config: &FtrlConfig,
243 ) -> Result<(f64, f64), RillError> {
244 let gradient_sq = gradient * gradient;
245 ensure_finite("ftrl_gradient_squared", gradient_sq)?;
246 let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
247 let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
248 ensure_finite("ftrl_sigma", sigma)?;
249 let sigma_w = sigma * weight;
250 ensure_finite("ftrl_sigma_weight", sigma_w)?;
251 let z_delta = gradient - sigma_w;
252 ensure_finite("ftrl_z_delta", z_delta)?;
253 let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
254 Ok((z_new, n_new))
255 }
256
257 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
259 fn validate(&self) -> Result<(), RillError> {
260 ensure_finite("ftrl_z", self.z)?;
261 ensure_finite("ftrl_n", self.n)?;
262 if self.n < 0.0 {
263 return Err(RillError::InvalidState(format!(
264 "ftrl n must be non-negative, got {0}",
265 self.n
266 )));
267 }
268 Ok(())
269 }
270}
271
272#[cfg(feature = "serde")]
273impl<'de> serde::Deserialize<'de> for FtrlParam {
274 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
275 where
276 D: serde::Deserializer<'de>,
277 {
278 #[derive(serde::Deserialize)]
279 struct FtrlParamState {
280 z: f64,
281 n: f64,
282 }
283
284 let state = FtrlParamState::deserialize(deserializer)?;
285 let param = FtrlParam {
286 z: state.z,
287 n: state.n,
288 };
289 param.validate().map_err(serde::de::Error::custom)?;
290 Ok(param)
291 }
292}
293
294fn compute_dot(
301 params: &BTreeMap<FeatureId, FtrlParam>,
302 config: &FtrlConfig,
303 features: &SparseFeatures,
304) -> Result<f64, RillError> {
305 if features.is_empty() {
306 return Err(RillError::EmptyFeatures);
307 }
308 let mut dot = 0.0;
309 for &(id, value) in features.values() {
310 ensure_finite("sparse_value", value)?;
311 if let Some(param) = params.get(&id) {
312 let w = param.weight(config);
313 ensure_finite("ftrl_weight", w)?;
314 let contribution = w * value;
315 ensure_finite("ftrl_dot_contribution", contribution)?;
316 dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
317 }
318 }
319 Ok(dot)
320}
321
322#[derive(Debug, Clone)]
341#[cfg_attr(feature = "serde", derive(serde::Serialize))]
342pub struct FtrlRegressor {
343 config: FtrlConfig,
344 params: BTreeMap<FeatureId, FtrlParam>,
345 intercept: FtrlParam,
346 samples_seen: u64,
347}
348
349impl FtrlRegressor {
350 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
354 config.validate()?;
355 Ok(Self {
356 config,
357 params: BTreeMap::new(),
358 intercept: FtrlParam::default(),
359 samples_seen: 0,
360 })
361 }
362
363 pub const fn config(&self) -> &FtrlConfig {
365 &self.config
366 }
367
368 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
373 self.params
374 .iter()
375 .map(|(&id, param)| (id, param.weight(&self.config)))
376 .filter(|&(_, w)| w != 0.0)
377 .collect()
378 }
379
380 pub fn intercept(&self) -> f64 {
382 self.intercept.intercept_weight(&self.config)
383 }
384
385 pub fn feature_count(&self) -> usize {
387 self.params.len()
388 }
389
390 fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
392 let dot = compute_dot(&self.params, &self.config, features)?;
393 let intercept = self.intercept.intercept_weight(&self.config);
394 ensure_finite("ftrl_intercept", intercept)?;
395 Ok(dot + intercept)
396 }
397}
398
399#[cfg(feature = "serde")]
400impl<'de> serde::Deserialize<'de> for FtrlRegressor {
401 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
402 where
403 D: serde::Deserializer<'de>,
404 {
405 #[derive(serde::Deserialize)]
406 struct FtrlRegressorState {
407 config: FtrlConfig,
408 params: BTreeMap<FeatureId, FtrlParam>,
409 intercept: FtrlParam,
410 samples_seen: u64,
411 }
412
413 let state = FtrlRegressorState::deserialize(deserializer)?;
414 let model = FtrlRegressor {
415 config: state.config,
416 params: state.params,
417 intercept: state.intercept,
418 samples_seen: state.samples_seen,
419 };
420 model
423 .validate_invariants()
424 .map_err(serde::de::Error::custom)?;
425 Ok(model)
426 }
427}
428
429impl FtrlRegressor {
430 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
431 fn validate_invariants(&self) -> Result<(), RillError> {
432 self.config.validate()?;
435 self.intercept.validate()?;
436 for param in self.params.values() {
437 param.validate()?;
438 }
439 Ok(())
440 }
441}
442
443impl SparseRegressor for FtrlRegressor {
444 fn samples_seen(&self) -> u64 {
445 self.samples_seen
446 }
447
448 fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
449 self.predict_inner(features)
450 }
451
452 fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
453 if features.is_empty() {
454 return Err(RillError::EmptyFeatures);
455 }
456 ensure_finite("target", target)?;
457
458 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
460
461 let prediction = self.predict_inner(features)?;
462 ensure_finite("ftrl_prediction", prediction)?;
463 let grad = prediction - target;
464 ensure_finite("ftrl_gradient", grad)?;
465
466 let new_ids_count = features
470 .values()
471 .iter()
472 .filter(|(id, _)| !self.params.contains_key(id))
473 .count();
474 let mut skip_new_features = false;
475 if let Some(max_features) = self.config.max_features {
476 let projected = self.params.len().saturating_add(new_ids_count);
477 if projected > max_features {
478 match self.config.new_feature_policy {
479 NewFeaturePolicy::Reject => {
480 return Err(RillError::InvalidState(format!(
481 "FTRL feature count {projected} exceeds max_features {max_features}"
482 )));
483 }
484 NewFeaturePolicy::Ignore => {
485 skip_new_features = true;
486 }
487 }
488 }
489 }
490
491 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
493 for &(id, value) in features.values() {
494 let g = grad * value;
495 ensure_finite("ftrl_feature_gradient", g)?;
496
497 let is_new = !self.params.contains_key(&id);
498 if is_new && skip_new_features {
499 continue;
500 }
501
502 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
503 let w = param.weight(&self.config);
504 param.next_updated(g, w, &self.config)?
505 } else {
506 let param = FtrlParam::default();
507 let w = param.weight(&self.config);
508 param.next_updated(g, w, &self.config)?
509 };
510 updates.push((id, new_z, new_n));
511 }
512
513 let w_b = self.intercept.intercept_weight(&self.config);
515 let (new_intercept_z, new_intercept_n) =
516 self.intercept.next_updated(grad, w_b, &self.config)?;
517
518 for (id, new_z, new_n) in updates {
520 let param = self.params.entry(id).or_default();
521 param.z = new_z;
522 param.n = new_n;
523 }
524 self.intercept.z = new_intercept_z;
525 self.intercept.n = new_intercept_n;
526 self.samples_seen = next_samples_seen;
527
528 Ok(())
529 }
530
531 fn reset(&mut self) {
532 self.params.clear();
533 self.intercept = FtrlParam::default();
534 self.samples_seen = 0;
535 }
536}
537
538#[derive(Debug, Clone)]
561#[cfg_attr(feature = "serde", derive(serde::Serialize))]
562pub struct FtrlClassifier {
563 config: FtrlConfig,
564 params: BTreeMap<FeatureId, FtrlParam>,
565 intercept: FtrlParam,
566 samples_seen: u64,
567}
568
569impl FtrlClassifier {
570 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
574 config.validate()?;
575 Ok(Self {
576 config,
577 params: BTreeMap::new(),
578 intercept: FtrlParam::default(),
579 samples_seen: 0,
580 })
581 }
582
583 pub const fn config(&self) -> &FtrlConfig {
585 &self.config
586 }
587
588 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
593 self.params
594 .iter()
595 .map(|(&id, param)| (id, param.weight(&self.config)))
596 .filter(|&(_, w)| w != 0.0)
597 .collect()
598 }
599
600 pub fn intercept(&self) -> f64 {
602 self.intercept.intercept_weight(&self.config)
603 }
604
605 pub fn feature_count(&self) -> usize {
607 self.params.len()
608 }
609
610 fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
612 let dot = compute_dot(&self.params, &self.config, features)?;
613 let intercept = self.intercept.intercept_weight(&self.config);
614 ensure_finite("ftrl_intercept", intercept)?;
615 let logit = dot + intercept;
616 ensure_finite("ftrl_logit", logit)?;
617 Ok(sigmoid(logit))
618 }
619}
620
621#[cfg(feature = "serde")]
622impl<'de> serde::Deserialize<'de> for FtrlClassifier {
623 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
624 where
625 D: serde::Deserializer<'de>,
626 {
627 #[derive(serde::Deserialize)]
628 struct FtrlClassifierState {
629 config: FtrlConfig,
630 params: BTreeMap<FeatureId, FtrlParam>,
631 intercept: FtrlParam,
632 samples_seen: u64,
633 }
634
635 let state = FtrlClassifierState::deserialize(deserializer)?;
636 let model = FtrlClassifier {
637 config: state.config,
638 params: state.params,
639 intercept: state.intercept,
640 samples_seen: state.samples_seen,
641 };
642 model
643 .validate_invariants()
644 .map_err(serde::de::Error::custom)?;
645 Ok(model)
646 }
647}
648
649impl FtrlClassifier {
650 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
651 fn validate_invariants(&self) -> Result<(), RillError> {
652 self.config.validate()?;
653 self.intercept.validate()?;
654 for param in self.params.values() {
655 param.validate()?;
656 }
657 Ok(())
658 }
659}
660
661impl SparseClassifier for FtrlClassifier {
662 fn samples_seen(&self) -> u64 {
663 self.samples_seen
664 }
665
666 fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
667 self.predict_proba_inner(features)
668 }
669
670 fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
671 if features.is_empty() {
672 return Err(RillError::EmptyFeatures);
673 }
674
675 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
676
677 let probability = self.predict_proba_inner(features)?;
678 ensure_finite("ftrl_probability", probability)?;
679 let y = if target { 1.0 } else { 0.0 };
680 let grad = probability - y;
681 ensure_finite("ftrl_gradient", grad)?;
682
683 let new_ids_count = features
684 .values()
685 .iter()
686 .filter(|(id, _)| !self.params.contains_key(id))
687 .count();
688 let mut skip_new_features = false;
689 if let Some(max_features) = self.config.max_features {
690 let projected = self.params.len().saturating_add(new_ids_count);
691 if projected > max_features {
692 match self.config.new_feature_policy {
693 NewFeaturePolicy::Reject => {
694 return Err(RillError::InvalidState(format!(
695 "FTRL feature count {projected} exceeds max_features {max_features}"
696 )));
697 }
698 NewFeaturePolicy::Ignore => {
699 skip_new_features = true;
700 }
701 }
702 }
703 }
704
705 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
706 for &(id, value) in features.values() {
707 let g = grad * value;
708 ensure_finite("ftrl_feature_gradient", g)?;
709
710 let is_new = !self.params.contains_key(&id);
711 if is_new && skip_new_features {
712 continue;
713 }
714
715 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
716 let w = param.weight(&self.config);
717 param.next_updated(g, w, &self.config)?
718 } else {
719 let param = FtrlParam::default();
720 let w = param.weight(&self.config);
721 param.next_updated(g, w, &self.config)?
722 };
723 updates.push((id, new_z, new_n));
724 }
725
726 let w_b = self.intercept.intercept_weight(&self.config);
727 let (new_intercept_z, new_intercept_n) =
728 self.intercept.next_updated(grad, w_b, &self.config)?;
729
730 for (id, new_z, new_n) in updates {
731 let param = self.params.entry(id).or_default();
732 param.z = new_z;
733 param.n = new_n;
734 }
735 self.intercept.z = new_intercept_z;
736 self.intercept.n = new_intercept_n;
737 self.samples_seen = next_samples_seen;
738
739 Ok(())
740 }
741
742 fn reset(&mut self) {
743 self.params.clear();
744 self.intercept = FtrlParam::default();
745 self.samples_seen = 0;
746 }
747}
748
749#[cfg(test)]
750mod tests {
751 use super::*;
752 use rand::SeedableRng;
753
754 #[test]
759 fn cold_start_returns_zero() {
760 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
761 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
762 let pred = model.predict(&sf).unwrap();
763 assert!(pred.abs() < 1e-12);
764 }
765
766 #[test]
767 fn learn_linear_data_converges() {
768 let mut model = FtrlRegressor::new(FtrlConfig {
770 alpha: 0.5,
771 beta: 1.0,
772 l1: 0.0,
773 l2: 0.0,
774 max_features: None,
775 new_feature_policy: NewFeaturePolicy::default(),
776 })
777 .unwrap();
778 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
779 let mut first_err = 0.0;
780 let mut last_err = 0.0;
781 for i in 0..500 {
782 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
783 let y = 2.0 * x;
784 let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
785 let pred = model.predict(&sf).unwrap();
786 let err = (pred - y).abs();
787 if i < 10 {
788 first_err += err;
789 }
790 if i >= 490 {
791 last_err += err;
792 }
793 model.learn(&sf, y).unwrap();
794 }
795 assert!(last_err < first_err, "error should decrease");
796 let weights = model.weights();
797 assert_eq!(weights.len(), 1);
798 assert!(
799 (weights[0].1 - 2.0).abs() < 0.5,
800 "weight should approach 2.0"
801 );
802 }
803
804 #[test]
805 fn l1_produces_sparse_weights() {
806 let mut model = FtrlRegressor::new(FtrlConfig {
808 alpha: 0.1,
809 beta: 1.0,
810 l1: 100.0,
811 l2: 0.0,
812 max_features: None,
813 new_feature_policy: NewFeaturePolicy::default(),
814 })
815 .unwrap();
816 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
817 for _ in 0..200 {
818 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
819 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
820 let y = 0.5 * x1;
821 let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
822 model.learn(&sf, y).unwrap();
823 }
824 let weights = model.weights();
825 assert!(
827 weights.is_empty(),
828 "weights should all be zero, got {weights:?}"
829 );
830 }
831
832 #[test]
833 fn dynamic_features() {
834 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
835 assert_eq!(model.feature_count(), 0);
836 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
837 model.learn(&sf1, 1.0).unwrap();
838 assert_eq!(model.feature_count(), 1);
839 let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
841 model.learn(&sf2, 2.0).unwrap();
842 assert_eq!(model.feature_count(), 2);
843 assert!(model.params.contains_key(&0));
845 assert!(model.params.contains_key(&5));
846 }
847
848 #[test]
849 fn predict_does_not_update_state() {
850 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
851 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
852 let _ = model.predict(&sf).unwrap();
853 assert_eq!(model.samples_seen(), 0);
854 assert_eq!(model.feature_count(), 0);
855 model.learn(&sf, 1.0).unwrap();
857 let count_after_learn = model.feature_count();
858 let _ = model.predict(&sf).unwrap();
859 assert_eq!(model.feature_count(), count_after_learn);
860 assert_eq!(model.samples_seen(), 1);
861 }
862
863 #[test]
864 fn non_finite_value_rejected() {
865 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
866 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
868 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
869 assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
870 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
871 assert!(model.predict(&sf).is_ok());
872 }
873
874 #[test]
875 fn non_finite_target_rejected() {
876 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
877 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
878 assert!(model.learn(&sf, f64::NAN).is_err());
879 assert!(model.learn(&sf, f64::INFINITY).is_err());
880 assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
881 assert_eq!(model.samples_seen(), 0);
883 }
884
885 #[test]
886 fn empty_features_rejected() {
887 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
888 let sf = SparseFeatures::new();
889 assert!(model.predict(&sf).is_err());
890 assert!(model.learn(&sf, 1.0).is_err());
891 }
892
893 #[test]
894 fn reset_clears_state() {
895 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
896 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
897 model.learn(&sf, 3.0).unwrap();
898 model.learn(&sf, 3.0).unwrap();
899 assert_eq!(model.samples_seen(), 2);
900 assert_eq!(model.feature_count(), 2);
901 model.reset();
902 assert_eq!(model.samples_seen(), 0);
903 assert_eq!(model.feature_count(), 0);
904 assert!(model.predict(&sf).unwrap().abs() < 1e-12);
905 }
906
907 #[test]
908 fn invalid_config_rejected() {
909 assert!(
910 FtrlRegressor::new(FtrlConfig {
911 alpha: 0.0,
912 ..FtrlConfig::default()
913 })
914 .is_err()
915 );
916 assert!(
917 FtrlRegressor::new(FtrlConfig {
918 alpha: -1.0,
919 ..FtrlConfig::default()
920 })
921 .is_err()
922 );
923 assert!(
924 FtrlRegressor::new(FtrlConfig {
925 beta: -1.0,
926 ..FtrlConfig::default()
927 })
928 .is_err()
929 );
930 assert!(
931 FtrlRegressor::new(FtrlConfig {
932 l1: -1.0,
933 ..FtrlConfig::default()
934 })
935 .is_err()
936 );
937 assert!(
938 FtrlRegressor::new(FtrlConfig {
939 l2: -1.0,
940 ..FtrlConfig::default()
941 })
942 .is_err()
943 );
944 assert!(
945 FtrlRegressor::new(FtrlConfig {
946 alpha: f64::NAN,
947 ..FtrlConfig::default()
948 })
949 .is_err()
950 );
951 assert!(
952 FtrlRegressor::new(FtrlConfig {
953 max_features: Some(0),
954 ..FtrlConfig::default()
955 })
956 .is_err()
957 );
958 }
959
960 #[test]
961 #[cfg(feature = "serde")]
962 fn serde_roundtrip() {
963 let mut model = FtrlRegressor::new(FtrlConfig {
964 alpha: 0.2,
965 beta: 0.5,
966 l1: 0.5,
967 l2: 0.5,
968 max_features: Some(100),
969 new_feature_policy: NewFeaturePolicy::Reject,
970 })
971 .unwrap();
972 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
973 model.learn(&sf, 5.0).unwrap();
974 let json = serde_json::to_string(&model).unwrap();
975 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
976 assert_eq!(restored.samples_seen(), model.samples_seen());
977 assert_eq!(restored.feature_count(), model.feature_count());
978 let pred_orig = model.predict(&sf).unwrap();
979 let pred_restored = restored.predict(&sf).unwrap();
980 assert!((pred_orig - pred_restored).abs() < 1e-12);
981 }
982
983 #[test]
984 fn weights_returns_nonzero_only() {
985 let mut model = FtrlRegressor::new(FtrlConfig {
986 alpha: 0.5,
987 beta: 1.0,
988 l1: 0.0,
989 l2: 0.0,
990 max_features: None,
991 new_feature_policy: NewFeaturePolicy::default(),
992 })
993 .unwrap();
994 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
996 for _ in 0..50 {
997 model.learn(&sf, 1.0).unwrap();
998 }
999 let weights = model.weights();
1000 for &(_, w) in &weights {
1002 assert!(w != 0.0);
1003 }
1004 assert!(weights.iter().any(|&(id, _)| id == 0));
1006 }
1007
1008 #[test]
1009 fn multiple_features() {
1010 let mut model = FtrlRegressor::new(FtrlConfig {
1012 alpha: 0.5,
1013 beta: 1.0,
1014 l1: 0.0,
1015 l2: 0.0,
1016 max_features: None,
1017 new_feature_policy: NewFeaturePolicy::default(),
1018 })
1019 .unwrap();
1020 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1021 for _ in 0..500 {
1022 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1023 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1024 let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1025 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1026 model.learn(&sf, y).unwrap();
1027 }
1028 let weights = model.weights();
1029 assert_eq!(weights.len(), 2);
1030 let w0 = weights
1031 .iter()
1032 .find(|&&(id, _)| id == 0)
1033 .map(|&(_, w)| w)
1034 .unwrap();
1035 let w1 = weights
1036 .iter()
1037 .find(|&&(id, _)| id == 1)
1038 .map(|&(_, w)| w)
1039 .unwrap();
1040 assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1041 assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1042 assert!(
1043 (model.intercept() - 0.5).abs() < 0.5,
1044 "intercept should approach 0.5"
1045 );
1046 }
1047
1048 #[test]
1049 fn intercept_learned() {
1050 let mut model = FtrlRegressor::new(FtrlConfig {
1053 alpha: 0.5,
1054 beta: 1.0,
1055 l1: 0.0,
1056 l2: 0.0,
1057 max_features: None,
1058 new_feature_policy: NewFeaturePolicy::default(),
1059 })
1060 .unwrap();
1061 let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1062 for _ in 0..300 {
1063 model.learn(&sf, 3.0).unwrap();
1064 }
1065 let pred = model.predict(&sf).unwrap();
1066 assert!(
1067 (pred - 3.0).abs() < 0.5,
1068 "prediction should approach 3.0, got {pred}"
1069 );
1070 assert!(
1071 (model.intercept() - 3.0).abs() < 0.5,
1072 "intercept should approach 3.0"
1073 );
1074 assert!(model.weights().is_empty());
1076 }
1077
1078 #[test]
1079 fn high_dim_sparse() {
1080 let mut model = FtrlRegressor::new(FtrlConfig {
1083 alpha: 0.3,
1084 beta: 1.0,
1085 l1: 0.0,
1086 l2: 0.0,
1087 max_features: None,
1088 new_feature_policy: NewFeaturePolicy::default(),
1089 })
1090 .unwrap();
1091 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1092 let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1094 let mut first_err = 0.0;
1095 let mut last_err = 0.0;
1096 for i in 0..2000 {
1097 let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1098 for (j, &w) in true_w.iter().enumerate() {
1099 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1100 active.push((j as u64, x * w));
1101 }
1102 for k in 5..10 {
1104 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1105 active.push((k as u64 + 100, x));
1106 }
1107 active.sort_by_key(|(id, _)| *id);
1108 let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1109 let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1110 let pred = model.predict(&sf).unwrap();
1111 let err = (pred - y).abs();
1112 if i < 20 {
1113 first_err += err;
1114 }
1115 if i >= 1980 {
1116 last_err += err;
1117 }
1118 model.learn(&sf, y).unwrap();
1119 }
1120 assert!(
1121 last_err < first_err,
1122 "error should decrease in high-dim sparse"
1123 );
1124 }
1125
1126 #[test]
1131 fn regressor_overflow_does_not_mutate_state() {
1132 let mut model = FtrlRegressor::new(FtrlConfig {
1136 alpha: 0.1,
1137 beta: 1.0,
1138 l1: 0.0,
1139 l2: 0.0,
1140 max_features: None,
1141 new_feature_policy: NewFeaturePolicy::default(),
1142 })
1143 .unwrap();
1144 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1145 let result = model.learn(&sf, 1e100);
1146 assert!(result.is_err(), "expected overflow error, got {result:?}");
1147 assert_eq!(model.samples_seen(), 0);
1148 assert_eq!(model.feature_count(), 0);
1149 assert!(model.params.is_empty());
1150 assert_eq!(model.intercept.z, 0.0);
1151 assert_eq!(model.intercept.n, 0.0);
1152 }
1153
1154 #[test]
1155 fn regressor_partial_update_is_atomic() {
1156 let mut model = FtrlRegressor::new(FtrlConfig {
1159 alpha: 0.1,
1160 beta: 1.0,
1161 l1: 0.0,
1162 l2: 0.0,
1163 max_features: None,
1164 new_feature_policy: NewFeaturePolicy::default(),
1165 })
1166 .unwrap();
1167 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1168 assert!(model.learn(&sf, 1e100).is_err());
1169 assert!(!model.params.contains_key(&0));
1171 assert!(!model.params.contains_key(&1));
1172 assert_eq!(model.samples_seen(), 0);
1173 }
1174
1175 #[test]
1176 #[cfg(feature = "serde")]
1177 fn regressor_samples_seen_overflow_is_atomic() {
1178 let json = format!(
1179 "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1180 u64::MAX
1181 );
1182 let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1183 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1184 let result = model.learn(&sf, 1.0);
1185 assert!(result.is_err(), "expected counter overflow");
1186 assert_eq!(model.samples_seen(), u64::MAX);
1187 assert_eq!(model.feature_count(), 0);
1188 assert_eq!(model.intercept.z, 0.0);
1189 assert_eq!(model.intercept.n, 0.0);
1190 }
1191
1192 #[test]
1197 fn regressor_max_features_reject_at_limit() {
1198 let mut model = FtrlRegressor::new(FtrlConfig {
1199 alpha: 0.5,
1200 beta: 1.0,
1201 l1: 0.0,
1202 l2: 0.0,
1203 max_features: Some(2),
1204 new_feature_policy: NewFeaturePolicy::Reject,
1205 })
1206 .unwrap();
1207 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1209 model.learn(&sf, 1.0).unwrap();
1210 assert_eq!(model.feature_count(), 2);
1211 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1213 assert!(model.learn(&sf_new, 1.0).is_err());
1214 assert_eq!(model.feature_count(), 2);
1215 assert_eq!(model.samples_seen(), 1);
1216 let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1218 model.learn(&sf_existing, 1.0).unwrap();
1219 assert_eq!(model.feature_count(), 2);
1220 assert_eq!(model.samples_seen(), 2);
1221 }
1222
1223 #[test]
1224 fn regressor_max_features_ignore_skips_new() {
1225 let mut model = FtrlRegressor::new(FtrlConfig {
1226 alpha: 0.5,
1227 beta: 1.0,
1228 l1: 0.0,
1229 l2: 0.0,
1230 max_features: Some(2),
1231 new_feature_policy: NewFeaturePolicy::Ignore,
1232 })
1233 .unwrap();
1234 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1235 model.learn(&sf, 1.0).unwrap();
1236 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1240 model.learn(&sf_mixed, 1.0).unwrap();
1241 assert_eq!(model.feature_count(), 2);
1242 assert!(!model.params.contains_key(&2));
1243 assert_eq!(model.samples_seen(), 2);
1244 }
1245
1246 #[test]
1247 fn regressor_max_features_multi_new_prejudge() {
1248 let mut model = FtrlRegressor::new(FtrlConfig {
1249 alpha: 0.5,
1250 beta: 1.0,
1251 l1: 0.0,
1252 l2: 0.0,
1253 max_features: Some(2),
1254 new_feature_policy: NewFeaturePolicy::Reject,
1255 })
1256 .unwrap();
1257 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1260 assert!(model.learn(&sf, 1.0).is_err());
1261 assert_eq!(model.feature_count(), 0);
1262 assert_eq!(model.samples_seen(), 0);
1263 }
1264
1265 #[test]
1270 #[cfg(feature = "serde")]
1271 fn regressor_serde_rejects_negative_n() {
1272 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1273 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1274 assert!(result.is_err(), "negative n must be rejected");
1275 }
1276
1277 #[test]
1278 #[cfg(feature = "serde")]
1279 fn regressor_serde_rejects_invalid_config() {
1280 let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1281 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1282 assert!(result.is_err(), "invalid alpha must be rejected");
1283 }
1284
1285 #[test]
1286 #[cfg(feature = "serde")]
1287 fn regressor_serde_accepts_missing_optional_fields() {
1288 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1290 let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1291 assert!(model.config().max_features.is_none());
1292 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1293 }
1294
1295 #[test]
1300 fn cold_start_returns_0_5() {
1301 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1302 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1303 let p = model.predict_proba(&sf).unwrap();
1304 assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1305 }
1306
1307 #[test]
1308 fn learn_separable_data() {
1309 let mut model = FtrlClassifier::new(FtrlConfig {
1311 alpha: 0.5,
1312 beta: 1.0,
1313 l1: 0.0,
1314 l2: 0.0,
1315 max_features: None,
1316 new_feature_policy: NewFeaturePolicy::default(),
1317 })
1318 .unwrap();
1319 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1320 for _ in 0..1000 {
1321 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1322 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1323 let y = x0 > 0.0;
1324 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1325 model.learn(&sf, y).unwrap();
1326 }
1327 let p_pos = model
1328 .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1329 .unwrap();
1330 let p_neg = model
1331 .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1332 .unwrap();
1333 assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1334 assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1335 }
1336
1337 #[test]
1338 fn classifier_l1_produces_sparse_weights() {
1339 let mut model = FtrlClassifier::new(FtrlConfig {
1340 alpha: 0.1,
1341 beta: 1.0,
1342 l1: 100.0,
1343 l2: 0.0,
1344 max_features: None,
1345 new_feature_policy: NewFeaturePolicy::default(),
1346 })
1347 .unwrap();
1348 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1349 for _ in 0..200 {
1350 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1351 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1352 let y = x0 > 0.0;
1353 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1354 model.learn(&sf, y).unwrap();
1355 }
1356 let weights = model.weights();
1357 assert!(
1358 weights.is_empty(),
1359 "weights should all be zero with high L1, got {weights:?}"
1360 );
1361 }
1362
1363 #[test]
1364 fn classifier_dynamic_features() {
1365 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1366 assert_eq!(model.feature_count(), 0);
1367 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1368 model.learn(&sf1, true).unwrap();
1369 assert_eq!(model.feature_count(), 1);
1370 let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1371 model.learn(&sf2, false).unwrap();
1372 assert_eq!(model.feature_count(), 2);
1373 }
1374
1375 #[test]
1376 fn classifier_predict_does_not_update_state() {
1377 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1378 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1379 let _ = model.predict_proba(&sf).unwrap();
1380 assert_eq!(model.samples_seen(), 0);
1381 assert_eq!(model.feature_count(), 0);
1382 model.learn(&sf, true).unwrap();
1383 let count = model.feature_count();
1384 let _ = model.predict_proba(&sf).unwrap();
1385 assert_eq!(model.feature_count(), count);
1386 assert_eq!(model.samples_seen(), 1);
1387 }
1388
1389 #[test]
1390 fn classifier_non_finite_value_rejected() {
1391 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1392 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1393 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1394 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1395 assert!(model.predict_proba(&sf).is_ok());
1396 }
1397
1398 #[test]
1399 fn classifier_empty_features_rejected() {
1400 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1401 let sf = SparseFeatures::new();
1402 assert!(model.predict_proba(&sf).is_err());
1403 assert!(model.learn(&sf, true).is_err());
1404 }
1405
1406 #[test]
1407 fn classifier_reset_clears_state() {
1408 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1409 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1410 model.learn(&sf, true).unwrap();
1411 model.learn(&sf, false).unwrap();
1412 assert_eq!(model.samples_seen(), 2);
1413 assert!(model.feature_count() > 0);
1414 model.reset();
1415 assert_eq!(model.samples_seen(), 0);
1416 assert_eq!(model.feature_count(), 0);
1417 let p = model.predict_proba(&sf).unwrap();
1418 assert!((p - 0.5).abs() < 1e-12);
1419 }
1420
1421 #[test]
1422 fn classifier_invalid_config_rejected() {
1423 assert!(
1424 FtrlClassifier::new(FtrlConfig {
1425 alpha: 0.0,
1426 ..FtrlConfig::default()
1427 })
1428 .is_err()
1429 );
1430 assert!(
1431 FtrlClassifier::new(FtrlConfig {
1432 beta: -0.1,
1433 ..FtrlConfig::default()
1434 })
1435 .is_err()
1436 );
1437 assert!(
1438 FtrlClassifier::new(FtrlConfig {
1439 l1: -1.0,
1440 ..FtrlConfig::default()
1441 })
1442 .is_err()
1443 );
1444 assert!(
1445 FtrlClassifier::new(FtrlConfig {
1446 l2: -1.0,
1447 ..FtrlConfig::default()
1448 })
1449 .is_err()
1450 );
1451 assert!(
1452 FtrlClassifier::new(FtrlConfig {
1453 alpha: f64::INFINITY,
1454 ..FtrlConfig::default()
1455 })
1456 .is_err()
1457 );
1458 }
1459
1460 #[test]
1461 #[cfg(feature = "serde")]
1462 fn classifier_serde_roundtrip() {
1463 let mut model = FtrlClassifier::new(FtrlConfig {
1464 alpha: 0.3,
1465 beta: 0.5,
1466 l1: 0.1,
1467 l2: 0.2,
1468 max_features: Some(100),
1469 new_feature_policy: NewFeaturePolicy::Reject,
1470 })
1471 .unwrap();
1472 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1473 model.learn(&sf, true).unwrap();
1474 model.learn(&sf, false).unwrap();
1475 let json = serde_json::to_string(&model).unwrap();
1476 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1477 assert_eq!(restored.samples_seen(), model.samples_seen());
1478 assert_eq!(restored.feature_count(), model.feature_count());
1479 let p1 = model.predict_proba(&sf).unwrap();
1480 let p2 = restored.predict_proba(&sf).unwrap();
1481 assert!((p1 - p2).abs() < 1e-12);
1482 }
1483
1484 #[test]
1485 fn predict_proba_in_range() {
1486 let mut model = FtrlClassifier::new(FtrlConfig {
1487 alpha: 0.5,
1488 beta: 1.0,
1489 l1: 0.0,
1490 l2: 0.0,
1491 max_features: None,
1492 new_feature_policy: NewFeaturePolicy::default(),
1493 })
1494 .unwrap();
1495 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1496 for _ in 0..200 {
1497 let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1498 let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1499 let y = x0 > 0.0;
1500 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1501 model.learn(&sf, y).unwrap();
1502 let p = model.predict_proba(&sf).unwrap();
1503 assert!(
1504 (0.0..=1.0).contains(&p),
1505 "probability must be in [0,1], got {p}"
1506 );
1507 }
1508 }
1509
1510 #[test]
1511 fn learn_improves_accuracy() {
1512 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1513 let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1515 .map(|_| {
1516 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1517 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1518 let y = x0 + x1 > 0.0;
1519 (
1520 SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1521 y,
1522 )
1523 })
1524 .collect();
1525
1526 let mut model = FtrlClassifier::new(FtrlConfig {
1527 alpha: 0.5,
1528 beta: 1.0,
1529 l1: 0.0,
1530 l2: 0.0,
1531 max_features: None,
1532 new_feature_policy: NewFeaturePolicy::default(),
1533 })
1534 .unwrap();
1535
1536 let acc_before: f64 = test_set
1538 .iter()
1539 .map(|(sf, y)| {
1540 let pred = model.predict(sf).unwrap();
1541 if pred == *y { 1.0 } else { 0.0 }
1542 })
1543 .sum::<f64>()
1544 / test_set.len() as f64;
1545
1546 for _ in 0..1000 {
1548 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1549 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1550 let y = x0 + x1 > 0.0;
1551 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1552 model.learn(&sf, y).unwrap();
1553 }
1554
1555 let acc_after: f64 = test_set
1556 .iter()
1557 .map(|(sf, y)| {
1558 let pred = model.predict(sf).unwrap();
1559 if pred == *y { 1.0 } else { 0.0 }
1560 })
1561 .sum::<f64>()
1562 / test_set.len() as f64;
1563
1564 assert!(
1565 acc_after > acc_before,
1566 "accuracy should improve: {acc_before} -> {acc_after}"
1567 );
1568 }
1569
1570 #[test]
1571 fn classifier_weights_returns_nonzero_only() {
1572 let mut model = FtrlClassifier::new(FtrlConfig {
1573 alpha: 0.5,
1574 beta: 1.0,
1575 l1: 0.0,
1576 l2: 0.0,
1577 max_features: None,
1578 new_feature_policy: NewFeaturePolicy::default(),
1579 })
1580 .unwrap();
1581 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1582 for _ in 0..50 {
1583 model.learn(&sf, true).unwrap();
1584 }
1585 let weights = model.weights();
1586 for &(_, w) in &weights {
1587 assert!(w != 0.0);
1588 }
1589 }
1590
1591 #[test]
1592 fn classifier_multiple_features() {
1593 let mut model = FtrlClassifier::new(FtrlConfig {
1594 alpha: 0.5,
1595 beta: 1.0,
1596 l1: 0.0,
1597 l2: 0.0,
1598 max_features: None,
1599 new_feature_policy: NewFeaturePolicy::default(),
1600 })
1601 .unwrap();
1602 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1603 for _ in 0..1000 {
1604 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1605 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1606 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1607 let y = x0 + x1 > 0.0;
1609 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1610 model.learn(&sf, y).unwrap();
1611 }
1612 let weights = model.weights();
1613 assert!(weights.iter().any(|&(id, _)| id == 0));
1615 assert!(weights.iter().any(|&(id, _)| id == 1));
1616 let p_pos = model
1618 .predict_proba(
1619 &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1620 )
1621 .unwrap();
1622 let p_neg = model
1623 .predict_proba(
1624 &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1625 )
1626 .unwrap();
1627 assert!(p_pos > 0.8);
1628 assert!(p_neg < 0.2);
1629 }
1630
1631 #[test]
1632 fn log_loss_converges() {
1633 let mut model = FtrlClassifier::new(FtrlConfig {
1635 alpha: 0.5,
1636 beta: 1.0,
1637 l1: 0.0,
1638 l2: 0.0,
1639 max_features: None,
1640 new_feature_policy: NewFeaturePolicy::default(),
1641 })
1642 .unwrap();
1643 let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1644 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1645 let mut first_loss = 0.0;
1646 let mut last_loss = 0.0;
1647 for i in 0..1000 {
1648 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1649 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1650 let y = x0 > 0.0;
1651 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1652 let p = model.predict_proba(&sf).unwrap();
1653 let loss = loss_fn.loss(p, y);
1654 if i < 20 {
1655 first_loss += loss;
1656 }
1657 if i >= 980 {
1658 last_loss += loss;
1659 }
1660 model.learn(&sf, y).unwrap();
1661 }
1662 assert!(last_loss < first_loss, "log loss should decrease");
1663 }
1664
1665 #[test]
1670 fn classifier_overflow_does_not_mutate_state() {
1671 let mut model = FtrlClassifier::new(FtrlConfig {
1672 alpha: 0.1,
1673 beta: 1.0,
1674 l1: 0.0,
1675 l2: 0.0,
1676 max_features: None,
1677 new_feature_policy: NewFeaturePolicy::default(),
1678 })
1679 .unwrap();
1680 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1681 let result = model.learn(&sf, false);
1684 assert!(result.is_err(), "expected overflow error, got {result:?}");
1685 assert_eq!(model.samples_seen(), 0);
1686 assert_eq!(model.feature_count(), 0);
1687 assert!(model.params.is_empty());
1688 assert_eq!(model.intercept.z, 0.0);
1689 assert_eq!(model.intercept.n, 0.0);
1690 }
1691
1692 #[test]
1693 fn classifier_partial_update_is_atomic() {
1694 let mut model = FtrlClassifier::new(FtrlConfig {
1695 alpha: 0.1,
1696 beta: 1.0,
1697 l1: 0.0,
1698 l2: 0.0,
1699 max_features: None,
1700 new_feature_policy: NewFeaturePolicy::default(),
1701 })
1702 .unwrap();
1703 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1704 assert!(model.learn(&sf, false).is_err());
1705 assert!(!model.params.contains_key(&0));
1706 assert!(!model.params.contains_key(&1));
1707 assert_eq!(model.samples_seen(), 0);
1708 }
1709
1710 #[test]
1711 #[cfg(feature = "serde")]
1712 fn classifier_samples_seen_overflow_is_atomic() {
1713 let json = format!(
1714 "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1715 u64::MAX
1716 );
1717 let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1718 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1719 let result = model.learn(&sf, true);
1720 assert!(result.is_err(), "expected counter overflow");
1721 assert_eq!(model.samples_seen(), u64::MAX);
1722 assert_eq!(model.feature_count(), 0);
1723 assert_eq!(model.intercept.z, 0.0);
1724 assert_eq!(model.intercept.n, 0.0);
1725 }
1726
1727 #[test]
1732 fn classifier_max_features_reject_at_limit() {
1733 let mut model = FtrlClassifier::new(FtrlConfig {
1734 alpha: 0.5,
1735 beta: 1.0,
1736 l1: 0.0,
1737 l2: 0.0,
1738 max_features: Some(2),
1739 new_feature_policy: NewFeaturePolicy::Reject,
1740 })
1741 .unwrap();
1742 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1743 model.learn(&sf, true).unwrap();
1744 assert_eq!(model.feature_count(), 2);
1745 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1746 assert!(model.learn(&sf_new, true).is_err());
1747 assert_eq!(model.feature_count(), 2);
1748 assert_eq!(model.samples_seen(), 1);
1749 }
1750
1751 #[test]
1752 fn classifier_max_features_ignore_skips_new() {
1753 let mut model = FtrlClassifier::new(FtrlConfig {
1754 alpha: 0.5,
1755 beta: 1.0,
1756 l1: 0.0,
1757 l2: 0.0,
1758 max_features: Some(2),
1759 new_feature_policy: NewFeaturePolicy::Ignore,
1760 })
1761 .unwrap();
1762 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1763 model.learn(&sf, true).unwrap();
1764 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1765 model.learn(&sf_mixed, false).unwrap();
1766 assert_eq!(model.feature_count(), 2);
1767 assert!(!model.params.contains_key(&2));
1768 assert_eq!(model.samples_seen(), 2);
1769 }
1770
1771 #[test]
1772 fn classifier_max_features_multi_new_prejudge() {
1773 let mut model = FtrlClassifier::new(FtrlConfig {
1774 alpha: 0.5,
1775 beta: 1.0,
1776 l1: 0.0,
1777 l2: 0.0,
1778 max_features: Some(2),
1779 new_feature_policy: NewFeaturePolicy::Reject,
1780 })
1781 .unwrap();
1782 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1783 assert!(model.learn(&sf, true).is_err());
1784 assert_eq!(model.feature_count(), 0);
1785 assert_eq!(model.samples_seen(), 0);
1786 }
1787
1788 #[test]
1793 #[cfg(feature = "serde")]
1794 fn classifier_serde_rejects_negative_n() {
1795 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1796 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1797 assert!(result.is_err(), "negative n must be rejected");
1798 }
1799
1800 #[test]
1801 #[cfg(feature = "serde")]
1802 fn classifier_serde_rejects_invalid_config() {
1803 let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1804 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1805 assert!(result.is_err(), "invalid alpha must be rejected");
1806 }
1807
1808 #[test]
1809 #[cfg(feature = "serde")]
1810 fn classifier_serde_accepts_missing_optional_fields() {
1811 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1812 let model: FtrlClassifier = serde_json::from_str(json).unwrap();
1813 assert!(model.config().max_features.is_none());
1814 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1815 }
1816}