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 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
247 fn weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
248 ensure_finite("ftrl_z", self.z)?;
249 ensure_finite("ftrl_n", self.n)?;
250 if self.n < 0.0 {
251 return Err(RillError::InvalidState(format!(
252 "ftrl n must be non-negative, got {}",
253 self.n
254 )));
255 }
256 if self.z.abs() <= config.l1 {
258 return Ok(0.0);
259 }
260 let sign = self.z.signum();
261 let numerator = -(self.z - sign * config.l1);
262 ensure_finite("ftrl_weight_numerator", numerator)?;
263 let sqrt_n = self.n.sqrt();
264 ensure_finite("ftrl_weight_sqrt_n", sqrt_n)?;
265 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
266 ensure_finite("ftrl_weight_denominator", denominator)?;
267 if denominator == 0.0 {
268 return Err(RillError::InvalidState(format!(
269 "ftrl weight denominator is zero (z={}, n={}, alpha={}, beta={}, l1={}, l2={})",
270 self.z, self.n, config.alpha, config.beta, config.l1, config.l2
271 )));
272 }
273 let weight = numerator / denominator;
274 ensure_finite("ftrl_weight", weight)?;
275 Ok(weight)
276 }
277
278 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
287 fn intercept_weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
288 ensure_finite("ftrl_z", self.z)?;
289 ensure_finite("ftrl_n", self.n)?;
290 if self.n < 0.0 {
291 return Err(RillError::InvalidState(format!(
292 "ftrl n must be non-negative, got {}",
293 self.n
294 )));
295 }
296 if self.n == 0.0 {
300 if self.z != 0.0 {
301 return Err(RillError::InvalidState(format!(
302 "ftrl intercept has n=0 but z={} (non-zero); cannot produce a finite weight",
303 self.z
304 )));
305 }
306 return Ok(0.0);
307 }
308 let numerator = -self.z;
309 ensure_finite("ftrl_intercept_numerator", numerator)?;
310 let sqrt_n = self.n.sqrt();
311 ensure_finite("ftrl_intercept_sqrt_n", sqrt_n)?;
312 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
313 ensure_finite("ftrl_intercept_denominator", denominator)?;
314 if denominator == 0.0 {
315 return Err(RillError::InvalidState(format!(
316 "ftrl intercept denominator is zero (z={}, n={}, alpha={}, beta={}, l2={})",
317 self.z, self.n, config.alpha, config.beta, config.l2
318 )));
319 }
320 let weight = numerator / denominator;
321 ensure_finite("ftrl_intercept_weight", weight)?;
322 Ok(weight)
323 }
324
325 fn next_updated(
336 &self,
337 gradient: f64,
338 weight: f64,
339 config: &FtrlConfig,
340 ) -> Result<(f64, f64), RillError> {
341 let gradient_sq = gradient * gradient;
342 ensure_finite("ftrl_gradient_squared", gradient_sq)?;
343 if gradient != 0.0 && gradient_sq == 0.0 {
350 return Err(RillError::NonFiniteValue {
351 field: "ftrl_gradient_squared",
352 value: gradient_sq,
353 });
354 }
355 let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
356 let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
357 ensure_finite("ftrl_sigma", sigma)?;
358 let sigma_w = sigma * weight;
359 ensure_finite("ftrl_sigma_weight", sigma_w)?;
360 let z_delta = gradient - sigma_w;
361 ensure_finite("ftrl_z_delta", z_delta)?;
362 let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
363 Ok((z_new, n_new))
364 }
365
366 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
378 fn validate(&self) -> Result<(), RillError> {
379 ensure_finite("ftrl_z", self.z)?;
380 ensure_finite("ftrl_n", self.n)?;
381 if self.n < 0.0 {
382 return Err(RillError::InvalidState(format!(
383 "ftrl n must be non-negative, got {0}",
384 self.n
385 )));
386 }
387 if self.n == 0.0 && self.z != 0.0 {
388 return Err(RillError::InvalidState(format!(
389 "ftrl param has n=0 but z={0} (non-zero); this state cannot \
390 produce a finite weight",
391 self.z
392 )));
393 }
394 Ok(())
395 }
396}
397
398#[cfg(feature = "serde")]
399impl<'de> serde::Deserialize<'de> for FtrlParam {
400 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
401 where
402 D: serde::Deserializer<'de>,
403 {
404 #[derive(serde::Deserialize)]
405 struct FtrlParamState {
406 z: f64,
407 n: f64,
408 }
409
410 let state = FtrlParamState::deserialize(deserializer)?;
411 let param = FtrlParam {
412 z: state.z,
413 n: state.n,
414 };
415 param.validate().map_err(serde::de::Error::custom)?;
416 Ok(param)
417 }
418}
419
420fn compute_dot(
427 params: &BTreeMap<FeatureId, FtrlParam>,
428 config: &FtrlConfig,
429 features: &SparseFeatures,
430) -> Result<f64, RillError> {
431 if features.is_empty() {
432 return Err(RillError::EmptyFeatures);
433 }
434 let mut dot = 0.0;
435 for &(id, value) in features.values() {
436 ensure_finite("sparse_value", value)?;
437 if let Some(param) = params.get(&id) {
438 let w = param.weight(config);
439 ensure_finite("ftrl_weight", w)?;
440 let contribution = w * value;
441 ensure_finite("ftrl_dot_contribution", contribution)?;
442 dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
443 }
444 }
445 Ok(dot)
446}
447
448#[derive(Debug, Clone)]
467#[cfg_attr(feature = "serde", derive(serde::Serialize))]
468pub struct FtrlRegressor {
469 config: FtrlConfig,
470 params: BTreeMap<FeatureId, FtrlParam>,
471 intercept: FtrlParam,
472 samples_seen: u64,
473}
474
475impl FtrlRegressor {
476 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
480 config.validate()?;
481 Ok(Self {
482 config,
483 params: BTreeMap::new(),
484 intercept: FtrlParam::default(),
485 samples_seen: 0,
486 })
487 }
488
489 pub const fn config(&self) -> &FtrlConfig {
491 &self.config
492 }
493
494 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
499 self.params
500 .iter()
501 .map(|(&id, param)| (id, param.weight(&self.config)))
502 .filter(|&(_, w)| w != 0.0)
503 .collect()
504 }
505
506 pub fn intercept(&self) -> f64 {
508 self.intercept.intercept_weight(&self.config)
509 }
510
511 pub fn feature_count(&self) -> usize {
513 self.params.len()
514 }
515
516 fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
518 let dot = compute_dot(&self.params, &self.config, features)?;
519 let intercept = self.intercept.intercept_weight(&self.config);
520 ensure_finite("ftrl_intercept", intercept)?;
521 checked_finite_add(dot, intercept, "ftrl_prediction")
522 }
523}
524
525#[cfg(feature = "serde")]
526impl<'de> serde::Deserialize<'de> for FtrlRegressor {
527 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
528 where
529 D: serde::Deserializer<'de>,
530 {
531 #[derive(serde::Deserialize)]
532 struct FtrlRegressorState {
533 config: FtrlConfig,
534 params: BTreeMap<FeatureId, FtrlParam>,
535 intercept: FtrlParam,
536 samples_seen: u64,
537 }
538
539 let state = FtrlRegressorState::deserialize(deserializer)?;
540 let model = FtrlRegressor {
541 config: state.config,
542 params: state.params,
543 intercept: state.intercept,
544 samples_seen: state.samples_seen,
545 };
546 model
549 .validate_invariants()
550 .map_err(serde::de::Error::custom)?;
551 Ok(model)
552 }
553}
554
555impl FtrlRegressor {
556 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
557 fn validate_invariants(&self) -> Result<(), RillError> {
558 self.config.validate()?;
565 if let Some(max_features) = self.config.max_features
570 && self.params.len() > max_features
571 {
572 return Err(RillError::InvalidState(format!(
573 "FTRL stored feature count {} exceeds max_features {}",
574 self.params.len(),
575 max_features
576 )));
577 }
578 for (id, param) in &self.params {
579 param.validate()?;
580 param
581 .weight_checked(&self.config)
582 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
583 }
584 self.intercept.validate()?;
585 self.intercept
586 .intercept_weight_checked(&self.config)
587 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
588 Ok(())
589 }
590}
591
592impl SparseRegressor for FtrlRegressor {
593 fn samples_seen(&self) -> u64 {
594 self.samples_seen
595 }
596
597 fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
598 self.predict_inner(features)
599 }
600
601 fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
602 if features.is_empty() {
603 return Err(RillError::EmptyFeatures);
604 }
605 ensure_finite("target", target)?;
606
607 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
609
610 let prediction = self.predict_inner(features)?;
611 ensure_finite("ftrl_prediction", prediction)?;
612 let grad = prediction - target;
613 ensure_finite("ftrl_gradient", grad)?;
614
615 let new_ids_count = features
619 .values()
620 .iter()
621 .filter(|(id, _)| !self.params.contains_key(id))
622 .count();
623 let mut skip_new_features = false;
624 if let Some(max_features) = self.config.max_features {
625 let projected = self.params.len().saturating_add(new_ids_count);
626 if projected > max_features {
627 match self.config.new_feature_policy {
628 NewFeaturePolicy::Reject => {
629 return Err(RillError::InvalidState(format!(
630 "FTRL feature count {projected} exceeds max_features {max_features}"
631 )));
632 }
633 NewFeaturePolicy::Ignore => {
634 skip_new_features = true;
635 }
636 }
637 }
638 }
639
640 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
642 for &(id, value) in features.values() {
643 let is_new = !self.params.contains_key(&id);
648 if is_new && skip_new_features {
649 continue;
650 }
651
652 let g = grad * value;
653 ensure_finite("ftrl_feature_gradient", g)?;
654
655 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
656 let w = param.weight(&self.config);
657 param.next_updated(g, w, &self.config)?
658 } else {
659 let param = FtrlParam::default();
660 let w = param.weight(&self.config);
661 param.next_updated(g, w, &self.config)?
662 };
663 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
668 ensure_finite("ftrl_next_weight", next_w)?;
669 updates.push((id, new_z, new_n));
670 }
671
672 let w_b = self.intercept.intercept_weight(&self.config);
674 let (new_intercept_z, new_intercept_n) =
675 self.intercept.next_updated(grad, w_b, &self.config)?;
676 let next_intercept_w = FtrlParam {
678 z: new_intercept_z,
679 n: new_intercept_n,
680 }
681 .intercept_weight(&self.config);
682 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
683
684 for (id, new_z, new_n) in updates {
686 let param = self.params.entry(id).or_default();
687 param.z = new_z;
688 param.n = new_n;
689 }
690 self.intercept.z = new_intercept_z;
691 self.intercept.n = new_intercept_n;
692 self.samples_seen = next_samples_seen;
693
694 Ok(())
695 }
696
697 fn reset(&mut self) {
698 self.params.clear();
699 self.intercept = FtrlParam::default();
700 self.samples_seen = 0;
701 }
702}
703
704#[derive(Debug, Clone)]
727#[cfg_attr(feature = "serde", derive(serde::Serialize))]
728pub struct FtrlClassifier {
729 config: FtrlConfig,
730 params: BTreeMap<FeatureId, FtrlParam>,
731 intercept: FtrlParam,
732 samples_seen: u64,
733}
734
735impl FtrlClassifier {
736 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
740 config.validate()?;
741 Ok(Self {
742 config,
743 params: BTreeMap::new(),
744 intercept: FtrlParam::default(),
745 samples_seen: 0,
746 })
747 }
748
749 pub const fn config(&self) -> &FtrlConfig {
751 &self.config
752 }
753
754 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
759 self.params
760 .iter()
761 .map(|(&id, param)| (id, param.weight(&self.config)))
762 .filter(|&(_, w)| w != 0.0)
763 .collect()
764 }
765
766 pub fn intercept(&self) -> f64 {
768 self.intercept.intercept_weight(&self.config)
769 }
770
771 pub fn feature_count(&self) -> usize {
773 self.params.len()
774 }
775
776 fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
778 let dot = compute_dot(&self.params, &self.config, features)?;
779 let intercept = self.intercept.intercept_weight(&self.config);
780 ensure_finite("ftrl_intercept", intercept)?;
781 let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
782 Ok(sigmoid(logit))
783 }
784}
785
786#[cfg(feature = "serde")]
787impl<'de> serde::Deserialize<'de> for FtrlClassifier {
788 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
789 where
790 D: serde::Deserializer<'de>,
791 {
792 #[derive(serde::Deserialize)]
793 struct FtrlClassifierState {
794 config: FtrlConfig,
795 params: BTreeMap<FeatureId, FtrlParam>,
796 intercept: FtrlParam,
797 samples_seen: u64,
798 }
799
800 let state = FtrlClassifierState::deserialize(deserializer)?;
801 let model = FtrlClassifier {
802 config: state.config,
803 params: state.params,
804 intercept: state.intercept,
805 samples_seen: state.samples_seen,
806 };
807 model
808 .validate_invariants()
809 .map_err(serde::de::Error::custom)?;
810 Ok(model)
811 }
812}
813
814impl FtrlClassifier {
815 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
816 fn validate_invariants(&self) -> Result<(), RillError> {
817 self.config.validate()?;
819 if let Some(max_features) = self.config.max_features
822 && self.params.len() > max_features
823 {
824 return Err(RillError::InvalidState(format!(
825 "FTRL stored feature count {} exceeds max_features {}",
826 self.params.len(),
827 max_features
828 )));
829 }
830 for (id, param) in &self.params {
831 param.validate()?;
832 param
833 .weight_checked(&self.config)
834 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
835 }
836 self.intercept.validate()?;
837 self.intercept
838 .intercept_weight_checked(&self.config)
839 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
840 Ok(())
841 }
842}
843
844impl SparseClassifier for FtrlClassifier {
845 fn samples_seen(&self) -> u64 {
846 self.samples_seen
847 }
848
849 fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
850 self.predict_proba_inner(features)
851 }
852
853 fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
854 if features.is_empty() {
855 return Err(RillError::EmptyFeatures);
856 }
857
858 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
859
860 let probability = self.predict_proba_inner(features)?;
861 ensure_finite("ftrl_probability", probability)?;
862 let y = if target { 1.0 } else { 0.0 };
863 let grad = probability - y;
864 ensure_finite("ftrl_gradient", grad)?;
865
866 let new_ids_count = features
867 .values()
868 .iter()
869 .filter(|(id, _)| !self.params.contains_key(id))
870 .count();
871 let mut skip_new_features = false;
872 if let Some(max_features) = self.config.max_features {
873 let projected = self.params.len().saturating_add(new_ids_count);
874 if projected > max_features {
875 match self.config.new_feature_policy {
876 NewFeaturePolicy::Reject => {
877 return Err(RillError::InvalidState(format!(
878 "FTRL feature count {projected} exceeds max_features {max_features}"
879 )));
880 }
881 NewFeaturePolicy::Ignore => {
882 skip_new_features = true;
883 }
884 }
885 }
886 }
887
888 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
889 for &(id, value) in features.values() {
890 let is_new = !self.params.contains_key(&id);
895 if is_new && skip_new_features {
896 continue;
897 }
898
899 let g = grad * value;
900 ensure_finite("ftrl_feature_gradient", g)?;
901
902 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
903 let w = param.weight(&self.config);
904 param.next_updated(g, w, &self.config)?
905 } else {
906 let param = FtrlParam::default();
907 let w = param.weight(&self.config);
908 param.next_updated(g, w, &self.config)?
909 };
910 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
912 ensure_finite("ftrl_next_weight", next_w)?;
913 updates.push((id, new_z, new_n));
914 }
915
916 let w_b = self.intercept.intercept_weight(&self.config);
917 let (new_intercept_z, new_intercept_n) =
918 self.intercept.next_updated(grad, w_b, &self.config)?;
919 let next_intercept_w = FtrlParam {
921 z: new_intercept_z,
922 n: new_intercept_n,
923 }
924 .intercept_weight(&self.config);
925 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
926
927 for (id, new_z, new_n) in updates {
928 let param = self.params.entry(id).or_default();
929 param.z = new_z;
930 param.n = new_n;
931 }
932 self.intercept.z = new_intercept_z;
933 self.intercept.n = new_intercept_n;
934 self.samples_seen = next_samples_seen;
935
936 Ok(())
937 }
938
939 fn reset(&mut self) {
940 self.params.clear();
941 self.intercept = FtrlParam::default();
942 self.samples_seen = 0;
943 }
944}
945
946#[cfg(test)]
947mod tests {
948 use super::*;
949 use rand::SeedableRng;
950
951 #[test]
956 fn cold_start_returns_zero() {
957 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
958 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
959 let pred = model.predict(&sf).unwrap();
960 assert!(pred.abs() < 1e-12);
961 }
962
963 #[test]
964 fn learn_linear_data_converges() {
965 let mut model = FtrlRegressor::new(FtrlConfig {
967 alpha: 0.5,
968 beta: 1.0,
969 l1: 0.0,
970 l2: 0.0,
971 max_features: None,
972 new_feature_policy: NewFeaturePolicy::default(),
973 })
974 .unwrap();
975 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
976 let mut first_err = 0.0;
977 let mut last_err = 0.0;
978 for i in 0..500 {
979 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
980 let y = 2.0 * x;
981 let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
982 let pred = model.predict(&sf).unwrap();
983 let err = (pred - y).abs();
984 if i < 10 {
985 first_err += err;
986 }
987 if i >= 490 {
988 last_err += err;
989 }
990 model.learn(&sf, y).unwrap();
991 }
992 assert!(last_err < first_err, "error should decrease");
993 let weights = model.weights();
994 assert_eq!(weights.len(), 1);
995 assert!(
996 (weights[0].1 - 2.0).abs() < 0.5,
997 "weight should approach 2.0"
998 );
999 }
1000
1001 #[test]
1002 fn l1_produces_sparse_weights() {
1003 let mut model = FtrlRegressor::new(FtrlConfig {
1005 alpha: 0.1,
1006 beta: 1.0,
1007 l1: 100.0,
1008 l2: 0.0,
1009 max_features: None,
1010 new_feature_policy: NewFeaturePolicy::default(),
1011 })
1012 .unwrap();
1013 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
1014 for _ in 0..200 {
1015 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1016 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1017 let y = 0.5 * x1;
1018 let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
1019 model.learn(&sf, y).unwrap();
1020 }
1021 let weights = model.weights();
1022 assert!(
1024 weights.is_empty(),
1025 "weights should all be zero, got {weights:?}"
1026 );
1027 }
1028
1029 #[test]
1030 fn dynamic_features() {
1031 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1032 assert_eq!(model.feature_count(), 0);
1033 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1034 model.learn(&sf1, 1.0).unwrap();
1035 assert_eq!(model.feature_count(), 1);
1036 let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1038 model.learn(&sf2, 2.0).unwrap();
1039 assert_eq!(model.feature_count(), 2);
1040 assert!(model.params.contains_key(&0));
1042 assert!(model.params.contains_key(&5));
1043 }
1044
1045 #[test]
1046 fn predict_does_not_update_state() {
1047 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1048 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1049 let _ = model.predict(&sf).unwrap();
1050 assert_eq!(model.samples_seen(), 0);
1051 assert_eq!(model.feature_count(), 0);
1052 model.learn(&sf, 1.0).unwrap();
1054 let count_after_learn = model.feature_count();
1055 let _ = model.predict(&sf).unwrap();
1056 assert_eq!(model.feature_count(), count_after_learn);
1057 assert_eq!(model.samples_seen(), 1);
1058 }
1059
1060 #[test]
1061 fn non_finite_value_rejected() {
1062 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1063 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1065 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1066 assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1067 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1068 assert!(model.predict(&sf).is_ok());
1069 }
1070
1071 #[test]
1072 fn non_finite_target_rejected() {
1073 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1074 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1075 assert!(model.learn(&sf, f64::NAN).is_err());
1076 assert!(model.learn(&sf, f64::INFINITY).is_err());
1077 assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1078 assert_eq!(model.samples_seen(), 0);
1080 }
1081
1082 #[test]
1083 fn empty_features_rejected() {
1084 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1085 let sf = SparseFeatures::new();
1086 assert!(model.predict(&sf).is_err());
1087 assert!(model.learn(&sf, 1.0).is_err());
1088 }
1089
1090 #[test]
1091 fn reset_clears_state() {
1092 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1093 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1094 model.learn(&sf, 3.0).unwrap();
1095 model.learn(&sf, 3.0).unwrap();
1096 assert_eq!(model.samples_seen(), 2);
1097 assert_eq!(model.feature_count(), 2);
1098 model.reset();
1099 assert_eq!(model.samples_seen(), 0);
1100 assert_eq!(model.feature_count(), 0);
1101 assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1102 }
1103
1104 #[test]
1105 fn invalid_config_rejected() {
1106 assert!(
1107 FtrlRegressor::new(FtrlConfig {
1108 alpha: 0.0,
1109 ..FtrlConfig::default()
1110 })
1111 .is_err()
1112 );
1113 assert!(
1114 FtrlRegressor::new(FtrlConfig {
1115 alpha: -1.0,
1116 ..FtrlConfig::default()
1117 })
1118 .is_err()
1119 );
1120 assert!(
1121 FtrlRegressor::new(FtrlConfig {
1122 beta: -1.0,
1123 ..FtrlConfig::default()
1124 })
1125 .is_err()
1126 );
1127 assert!(
1128 FtrlRegressor::new(FtrlConfig {
1129 l1: -1.0,
1130 ..FtrlConfig::default()
1131 })
1132 .is_err()
1133 );
1134 assert!(
1135 FtrlRegressor::new(FtrlConfig {
1136 l2: -1.0,
1137 ..FtrlConfig::default()
1138 })
1139 .is_err()
1140 );
1141 assert!(
1142 FtrlRegressor::new(FtrlConfig {
1143 alpha: f64::NAN,
1144 ..FtrlConfig::default()
1145 })
1146 .is_err()
1147 );
1148 assert!(
1149 FtrlRegressor::new(FtrlConfig {
1150 max_features: Some(0),
1151 ..FtrlConfig::default()
1152 })
1153 .is_err()
1154 );
1155 }
1156
1157 #[test]
1158 #[cfg(feature = "serde")]
1159 fn serde_roundtrip() {
1160 let mut model = FtrlRegressor::new(FtrlConfig {
1161 alpha: 0.2,
1162 beta: 0.5,
1163 l1: 0.5,
1164 l2: 0.5,
1165 max_features: Some(100),
1166 new_feature_policy: NewFeaturePolicy::Reject,
1167 })
1168 .unwrap();
1169 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1170 model.learn(&sf, 5.0).unwrap();
1171 let json = serde_json::to_string(&model).unwrap();
1172 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1173 assert_eq!(restored.samples_seen(), model.samples_seen());
1174 assert_eq!(restored.feature_count(), model.feature_count());
1175 let pred_orig = model.predict(&sf).unwrap();
1176 let pred_restored = restored.predict(&sf).unwrap();
1177 assert!((pred_orig - pred_restored).abs() < 1e-12);
1178 }
1179
1180 #[test]
1181 fn weights_returns_nonzero_only() {
1182 let mut model = FtrlRegressor::new(FtrlConfig {
1183 alpha: 0.5,
1184 beta: 1.0,
1185 l1: 0.0,
1186 l2: 0.0,
1187 max_features: None,
1188 new_feature_policy: NewFeaturePolicy::default(),
1189 })
1190 .unwrap();
1191 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1193 for _ in 0..50 {
1194 model.learn(&sf, 1.0).unwrap();
1195 }
1196 let weights = model.weights();
1197 for &(_, w) in &weights {
1199 assert!(w != 0.0);
1200 }
1201 assert!(weights.iter().any(|&(id, _)| id == 0));
1203 }
1204
1205 #[test]
1206 fn multiple_features() {
1207 let mut model = FtrlRegressor::new(FtrlConfig {
1209 alpha: 0.5,
1210 beta: 1.0,
1211 l1: 0.0,
1212 l2: 0.0,
1213 max_features: None,
1214 new_feature_policy: NewFeaturePolicy::default(),
1215 })
1216 .unwrap();
1217 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1218 for _ in 0..500 {
1219 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1220 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1221 let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1222 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1223 model.learn(&sf, y).unwrap();
1224 }
1225 let weights = model.weights();
1226 assert_eq!(weights.len(), 2);
1227 let w0 = weights
1228 .iter()
1229 .find(|&&(id, _)| id == 0)
1230 .map(|&(_, w)| w)
1231 .unwrap();
1232 let w1 = weights
1233 .iter()
1234 .find(|&&(id, _)| id == 1)
1235 .map(|&(_, w)| w)
1236 .unwrap();
1237 assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1238 assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1239 assert!(
1240 (model.intercept() - 0.5).abs() < 0.5,
1241 "intercept should approach 0.5"
1242 );
1243 }
1244
1245 #[test]
1246 fn intercept_learned() {
1247 let mut model = FtrlRegressor::new(FtrlConfig {
1250 alpha: 0.5,
1251 beta: 1.0,
1252 l1: 0.0,
1253 l2: 0.0,
1254 max_features: None,
1255 new_feature_policy: NewFeaturePolicy::default(),
1256 })
1257 .unwrap();
1258 let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1259 for _ in 0..300 {
1260 model.learn(&sf, 3.0).unwrap();
1261 }
1262 let pred = model.predict(&sf).unwrap();
1263 assert!(
1264 (pred - 3.0).abs() < 0.5,
1265 "prediction should approach 3.0, got {pred}"
1266 );
1267 assert!(
1268 (model.intercept() - 3.0).abs() < 0.5,
1269 "intercept should approach 3.0"
1270 );
1271 assert!(model.weights().is_empty());
1273 }
1274
1275 #[test]
1276 fn high_dim_sparse() {
1277 let mut model = FtrlRegressor::new(FtrlConfig {
1280 alpha: 0.3,
1281 beta: 1.0,
1282 l1: 0.0,
1283 l2: 0.0,
1284 max_features: None,
1285 new_feature_policy: NewFeaturePolicy::default(),
1286 })
1287 .unwrap();
1288 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1289 let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1291 let mut first_err = 0.0;
1292 let mut last_err = 0.0;
1293 for i in 0..2000 {
1294 let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1295 for (j, &w) in true_w.iter().enumerate() {
1296 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1297 active.push((j as u64, x * w));
1298 }
1299 for k in 5..10 {
1301 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1302 active.push((k as u64 + 100, x));
1303 }
1304 active.sort_by_key(|(id, _)| *id);
1305 let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1306 let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1307 let pred = model.predict(&sf).unwrap();
1308 let err = (pred - y).abs();
1309 if i < 20 {
1310 first_err += err;
1311 }
1312 if i >= 1980 {
1313 last_err += err;
1314 }
1315 model.learn(&sf, y).unwrap();
1316 }
1317 assert!(
1318 last_err < first_err,
1319 "error should decrease in high-dim sparse"
1320 );
1321 }
1322
1323 #[test]
1328 fn regressor_overflow_does_not_mutate_state() {
1329 let mut model = FtrlRegressor::new(FtrlConfig {
1333 alpha: 0.1,
1334 beta: 1.0,
1335 l1: 0.0,
1336 l2: 0.0,
1337 max_features: None,
1338 new_feature_policy: NewFeaturePolicy::default(),
1339 })
1340 .unwrap();
1341 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1342 let result = model.learn(&sf, 1e100);
1343 assert!(result.is_err(), "expected overflow error, got {result:?}");
1344 assert_eq!(model.samples_seen(), 0);
1345 assert_eq!(model.feature_count(), 0);
1346 assert!(model.params.is_empty());
1347 assert_eq!(model.intercept.z, 0.0);
1348 assert_eq!(model.intercept.n, 0.0);
1349 }
1350
1351 #[test]
1352 fn regressor_partial_update_is_atomic() {
1353 let mut model = FtrlRegressor::new(FtrlConfig {
1356 alpha: 0.1,
1357 beta: 1.0,
1358 l1: 0.0,
1359 l2: 0.0,
1360 max_features: None,
1361 new_feature_policy: NewFeaturePolicy::default(),
1362 })
1363 .unwrap();
1364 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1365 assert!(model.learn(&sf, 1e100).is_err());
1366 assert!(!model.params.contains_key(&0));
1368 assert!(!model.params.contains_key(&1));
1369 assert_eq!(model.samples_seen(), 0);
1370 }
1371
1372 #[test]
1373 #[cfg(feature = "serde")]
1374 fn regressor_samples_seen_overflow_is_atomic() {
1375 let json = format!(
1376 "{{\"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\":{}}}",
1377 u64::MAX
1378 );
1379 let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1380 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1381 let result = model.learn(&sf, 1.0);
1382 assert!(result.is_err(), "expected counter overflow");
1383 assert_eq!(model.samples_seen(), u64::MAX);
1384 assert_eq!(model.feature_count(), 0);
1385 assert_eq!(model.intercept.z, 0.0);
1386 assert_eq!(model.intercept.n, 0.0);
1387 }
1388
1389 #[test]
1394 fn regressor_max_features_reject_at_limit() {
1395 let mut model = FtrlRegressor::new(FtrlConfig {
1396 alpha: 0.5,
1397 beta: 1.0,
1398 l1: 0.0,
1399 l2: 0.0,
1400 max_features: Some(2),
1401 new_feature_policy: NewFeaturePolicy::Reject,
1402 })
1403 .unwrap();
1404 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1406 model.learn(&sf, 1.0).unwrap();
1407 assert_eq!(model.feature_count(), 2);
1408 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1410 assert!(model.learn(&sf_new, 1.0).is_err());
1411 assert_eq!(model.feature_count(), 2);
1412 assert_eq!(model.samples_seen(), 1);
1413 let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1415 model.learn(&sf_existing, 1.0).unwrap();
1416 assert_eq!(model.feature_count(), 2);
1417 assert_eq!(model.samples_seen(), 2);
1418 }
1419
1420 #[test]
1421 fn regressor_max_features_ignore_skips_new() {
1422 let mut model = FtrlRegressor::new(FtrlConfig {
1423 alpha: 0.5,
1424 beta: 1.0,
1425 l1: 0.0,
1426 l2: 0.0,
1427 max_features: Some(2),
1428 new_feature_policy: NewFeaturePolicy::Ignore,
1429 })
1430 .unwrap();
1431 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1432 model.learn(&sf, 1.0).unwrap();
1433 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1437 model.learn(&sf_mixed, 1.0).unwrap();
1438 assert_eq!(model.feature_count(), 2);
1439 assert!(!model.params.contains_key(&2));
1440 assert_eq!(model.samples_seen(), 2);
1441 }
1442
1443 #[test]
1444 fn regressor_max_features_multi_new_prejudge() {
1445 let mut model = FtrlRegressor::new(FtrlConfig {
1446 alpha: 0.5,
1447 beta: 1.0,
1448 l1: 0.0,
1449 l2: 0.0,
1450 max_features: Some(2),
1451 new_feature_policy: NewFeaturePolicy::Reject,
1452 })
1453 .unwrap();
1454 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1457 assert!(model.learn(&sf, 1.0).is_err());
1458 assert_eq!(model.feature_count(), 0);
1459 assert_eq!(model.samples_seen(), 0);
1460 }
1461
1462 #[test]
1467 #[cfg(feature = "serde")]
1468 fn regressor_serde_rejects_negative_n() {
1469 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}";
1470 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1471 assert!(result.is_err(), "negative n must be rejected");
1472 }
1473
1474 #[test]
1475 #[cfg(feature = "serde")]
1476 fn regressor_serde_rejects_invalid_config() {
1477 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}";
1478 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1479 assert!(result.is_err(), "invalid alpha must be rejected");
1480 }
1481
1482 #[test]
1483 #[cfg(feature = "serde")]
1484 fn regressor_serde_accepts_missing_optional_fields() {
1485 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}";
1487 let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1488 assert!(model.config().max_features.is_none());
1489 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1490 }
1491
1492 #[test]
1497 fn cold_start_returns_0_5() {
1498 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1499 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1500 let p = model.predict_proba(&sf).unwrap();
1501 assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1502 }
1503
1504 #[test]
1505 fn learn_separable_data() {
1506 let mut model = FtrlClassifier::new(FtrlConfig {
1508 alpha: 0.5,
1509 beta: 1.0,
1510 l1: 0.0,
1511 l2: 0.0,
1512 max_features: None,
1513 new_feature_policy: NewFeaturePolicy::default(),
1514 })
1515 .unwrap();
1516 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1517 for _ in 0..1000 {
1518 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1519 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1520 let y = x0 > 0.0;
1521 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1522 model.learn(&sf, y).unwrap();
1523 }
1524 let p_pos = model
1525 .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1526 .unwrap();
1527 let p_neg = model
1528 .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1529 .unwrap();
1530 assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1531 assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1532 }
1533
1534 #[test]
1535 fn classifier_l1_produces_sparse_weights() {
1536 let mut model = FtrlClassifier::new(FtrlConfig {
1537 alpha: 0.1,
1538 beta: 1.0,
1539 l1: 100.0,
1540 l2: 0.0,
1541 max_features: None,
1542 new_feature_policy: NewFeaturePolicy::default(),
1543 })
1544 .unwrap();
1545 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1546 for _ in 0..200 {
1547 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1548 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1549 let y = x0 > 0.0;
1550 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1551 model.learn(&sf, y).unwrap();
1552 }
1553 let weights = model.weights();
1554 assert!(
1555 weights.is_empty(),
1556 "weights should all be zero with high L1, got {weights:?}"
1557 );
1558 }
1559
1560 #[test]
1561 fn classifier_dynamic_features() {
1562 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1563 assert_eq!(model.feature_count(), 0);
1564 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1565 model.learn(&sf1, true).unwrap();
1566 assert_eq!(model.feature_count(), 1);
1567 let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1568 model.learn(&sf2, false).unwrap();
1569 assert_eq!(model.feature_count(), 2);
1570 }
1571
1572 #[test]
1573 fn classifier_predict_does_not_update_state() {
1574 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1575 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1576 let _ = model.predict_proba(&sf).unwrap();
1577 assert_eq!(model.samples_seen(), 0);
1578 assert_eq!(model.feature_count(), 0);
1579 model.learn(&sf, true).unwrap();
1580 let count = model.feature_count();
1581 let _ = model.predict_proba(&sf).unwrap();
1582 assert_eq!(model.feature_count(), count);
1583 assert_eq!(model.samples_seen(), 1);
1584 }
1585
1586 #[test]
1587 fn classifier_non_finite_value_rejected() {
1588 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1589 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1590 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1591 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1592 assert!(model.predict_proba(&sf).is_ok());
1593 }
1594
1595 #[test]
1596 fn classifier_empty_features_rejected() {
1597 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1598 let sf = SparseFeatures::new();
1599 assert!(model.predict_proba(&sf).is_err());
1600 assert!(model.learn(&sf, true).is_err());
1601 }
1602
1603 #[test]
1604 fn classifier_reset_clears_state() {
1605 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1606 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1607 model.learn(&sf, true).unwrap();
1608 model.learn(&sf, false).unwrap();
1609 assert_eq!(model.samples_seen(), 2);
1610 assert!(model.feature_count() > 0);
1611 model.reset();
1612 assert_eq!(model.samples_seen(), 0);
1613 assert_eq!(model.feature_count(), 0);
1614 let p = model.predict_proba(&sf).unwrap();
1615 assert!((p - 0.5).abs() < 1e-12);
1616 }
1617
1618 #[test]
1619 fn classifier_invalid_config_rejected() {
1620 assert!(
1621 FtrlClassifier::new(FtrlConfig {
1622 alpha: 0.0,
1623 ..FtrlConfig::default()
1624 })
1625 .is_err()
1626 );
1627 assert!(
1628 FtrlClassifier::new(FtrlConfig {
1629 beta: -0.1,
1630 ..FtrlConfig::default()
1631 })
1632 .is_err()
1633 );
1634 assert!(
1635 FtrlClassifier::new(FtrlConfig {
1636 l1: -1.0,
1637 ..FtrlConfig::default()
1638 })
1639 .is_err()
1640 );
1641 assert!(
1642 FtrlClassifier::new(FtrlConfig {
1643 l2: -1.0,
1644 ..FtrlConfig::default()
1645 })
1646 .is_err()
1647 );
1648 assert!(
1649 FtrlClassifier::new(FtrlConfig {
1650 alpha: f64::INFINITY,
1651 ..FtrlConfig::default()
1652 })
1653 .is_err()
1654 );
1655 }
1656
1657 #[test]
1658 #[cfg(feature = "serde")]
1659 fn classifier_serde_roundtrip() {
1660 let mut model = FtrlClassifier::new(FtrlConfig {
1661 alpha: 0.3,
1662 beta: 0.5,
1663 l1: 0.1,
1664 l2: 0.2,
1665 max_features: Some(100),
1666 new_feature_policy: NewFeaturePolicy::Reject,
1667 })
1668 .unwrap();
1669 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1670 model.learn(&sf, true).unwrap();
1671 model.learn(&sf, false).unwrap();
1672 let json = serde_json::to_string(&model).unwrap();
1673 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1674 assert_eq!(restored.samples_seen(), model.samples_seen());
1675 assert_eq!(restored.feature_count(), model.feature_count());
1676 let p1 = model.predict_proba(&sf).unwrap();
1677 let p2 = restored.predict_proba(&sf).unwrap();
1678 assert!((p1 - p2).abs() < 1e-12);
1679 }
1680
1681 #[test]
1682 fn predict_proba_in_range() {
1683 let mut model = FtrlClassifier::new(FtrlConfig {
1684 alpha: 0.5,
1685 beta: 1.0,
1686 l1: 0.0,
1687 l2: 0.0,
1688 max_features: None,
1689 new_feature_policy: NewFeaturePolicy::default(),
1690 })
1691 .unwrap();
1692 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1693 for _ in 0..200 {
1694 let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1695 let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1696 let y = x0 > 0.0;
1697 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1698 model.learn(&sf, y).unwrap();
1699 let p = model.predict_proba(&sf).unwrap();
1700 assert!(
1701 (0.0..=1.0).contains(&p),
1702 "probability must be in [0,1], got {p}"
1703 );
1704 }
1705 }
1706
1707 #[test]
1708 fn learn_improves_accuracy() {
1709 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1710 let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1712 .map(|_| {
1713 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1714 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1715 let y = x0 + x1 > 0.0;
1716 (
1717 SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1718 y,
1719 )
1720 })
1721 .collect();
1722
1723 let mut model = FtrlClassifier::new(FtrlConfig {
1724 alpha: 0.5,
1725 beta: 1.0,
1726 l1: 0.0,
1727 l2: 0.0,
1728 max_features: None,
1729 new_feature_policy: NewFeaturePolicy::default(),
1730 })
1731 .unwrap();
1732
1733 let acc_before: f64 = test_set
1735 .iter()
1736 .map(|(sf, y)| {
1737 let pred = model.predict(sf).unwrap();
1738 if pred == *y { 1.0 } else { 0.0 }
1739 })
1740 .sum::<f64>()
1741 / test_set.len() as f64;
1742
1743 for _ in 0..1000 {
1745 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1746 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1747 let y = x0 + x1 > 0.0;
1748 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1749 model.learn(&sf, y).unwrap();
1750 }
1751
1752 let acc_after: f64 = test_set
1753 .iter()
1754 .map(|(sf, y)| {
1755 let pred = model.predict(sf).unwrap();
1756 if pred == *y { 1.0 } else { 0.0 }
1757 })
1758 .sum::<f64>()
1759 / test_set.len() as f64;
1760
1761 assert!(
1762 acc_after > acc_before,
1763 "accuracy should improve: {acc_before} -> {acc_after}"
1764 );
1765 }
1766
1767 #[test]
1768 fn classifier_weights_returns_nonzero_only() {
1769 let mut model = FtrlClassifier::new(FtrlConfig {
1770 alpha: 0.5,
1771 beta: 1.0,
1772 l1: 0.0,
1773 l2: 0.0,
1774 max_features: None,
1775 new_feature_policy: NewFeaturePolicy::default(),
1776 })
1777 .unwrap();
1778 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1779 for _ in 0..50 {
1780 model.learn(&sf, true).unwrap();
1781 }
1782 let weights = model.weights();
1783 for &(_, w) in &weights {
1784 assert!(w != 0.0);
1785 }
1786 }
1787
1788 #[test]
1789 fn classifier_multiple_features() {
1790 let mut model = FtrlClassifier::new(FtrlConfig {
1791 alpha: 0.5,
1792 beta: 1.0,
1793 l1: 0.0,
1794 l2: 0.0,
1795 max_features: None,
1796 new_feature_policy: NewFeaturePolicy::default(),
1797 })
1798 .unwrap();
1799 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1800 for _ in 0..1000 {
1801 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1802 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1803 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1804 let y = x0 + x1 > 0.0;
1806 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1807 model.learn(&sf, y).unwrap();
1808 }
1809 let weights = model.weights();
1810 assert!(weights.iter().any(|&(id, _)| id == 0));
1812 assert!(weights.iter().any(|&(id, _)| id == 1));
1813 let p_pos = model
1815 .predict_proba(
1816 &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1817 )
1818 .unwrap();
1819 let p_neg = model
1820 .predict_proba(
1821 &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1822 )
1823 .unwrap();
1824 assert!(p_pos > 0.8);
1825 assert!(p_neg < 0.2);
1826 }
1827
1828 #[test]
1829 fn log_loss_converges() {
1830 let mut model = FtrlClassifier::new(FtrlConfig {
1832 alpha: 0.5,
1833 beta: 1.0,
1834 l1: 0.0,
1835 l2: 0.0,
1836 max_features: None,
1837 new_feature_policy: NewFeaturePolicy::default(),
1838 })
1839 .unwrap();
1840 let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1841 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1842 let mut first_loss = 0.0;
1843 let mut last_loss = 0.0;
1844 for i in 0..1000 {
1845 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1846 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1847 let y = x0 > 0.0;
1848 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1849 let p = model.predict_proba(&sf).unwrap();
1850 let loss = loss_fn.loss(p, y);
1851 if i < 20 {
1852 first_loss += loss;
1853 }
1854 if i >= 980 {
1855 last_loss += loss;
1856 }
1857 model.learn(&sf, y).unwrap();
1858 }
1859 assert!(last_loss < first_loss, "log loss should decrease");
1860 }
1861
1862 #[test]
1867 fn classifier_overflow_does_not_mutate_state() {
1868 let mut model = FtrlClassifier::new(FtrlConfig {
1869 alpha: 0.1,
1870 beta: 1.0,
1871 l1: 0.0,
1872 l2: 0.0,
1873 max_features: None,
1874 new_feature_policy: NewFeaturePolicy::default(),
1875 })
1876 .unwrap();
1877 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1878 let result = model.learn(&sf, false);
1881 assert!(result.is_err(), "expected overflow error, got {result:?}");
1882 assert_eq!(model.samples_seen(), 0);
1883 assert_eq!(model.feature_count(), 0);
1884 assert!(model.params.is_empty());
1885 assert_eq!(model.intercept.z, 0.0);
1886 assert_eq!(model.intercept.n, 0.0);
1887 }
1888
1889 #[test]
1890 fn classifier_partial_update_is_atomic() {
1891 let mut model = FtrlClassifier::new(FtrlConfig {
1892 alpha: 0.1,
1893 beta: 1.0,
1894 l1: 0.0,
1895 l2: 0.0,
1896 max_features: None,
1897 new_feature_policy: NewFeaturePolicy::default(),
1898 })
1899 .unwrap();
1900 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1901 assert!(model.learn(&sf, false).is_err());
1902 assert!(!model.params.contains_key(&0));
1903 assert!(!model.params.contains_key(&1));
1904 assert_eq!(model.samples_seen(), 0);
1905 }
1906
1907 #[test]
1908 #[cfg(feature = "serde")]
1909 fn classifier_samples_seen_overflow_is_atomic() {
1910 let json = format!(
1911 "{{\"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\":{}}}",
1912 u64::MAX
1913 );
1914 let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1915 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1916 let result = model.learn(&sf, true);
1917 assert!(result.is_err(), "expected counter overflow");
1918 assert_eq!(model.samples_seen(), u64::MAX);
1919 assert_eq!(model.feature_count(), 0);
1920 assert_eq!(model.intercept.z, 0.0);
1921 assert_eq!(model.intercept.n, 0.0);
1922 }
1923
1924 #[test]
1929 fn classifier_max_features_reject_at_limit() {
1930 let mut model = FtrlClassifier::new(FtrlConfig {
1931 alpha: 0.5,
1932 beta: 1.0,
1933 l1: 0.0,
1934 l2: 0.0,
1935 max_features: Some(2),
1936 new_feature_policy: NewFeaturePolicy::Reject,
1937 })
1938 .unwrap();
1939 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1940 model.learn(&sf, true).unwrap();
1941 assert_eq!(model.feature_count(), 2);
1942 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1943 assert!(model.learn(&sf_new, true).is_err());
1944 assert_eq!(model.feature_count(), 2);
1945 assert_eq!(model.samples_seen(), 1);
1946 }
1947
1948 #[test]
1949 fn classifier_max_features_ignore_skips_new() {
1950 let mut model = FtrlClassifier::new(FtrlConfig {
1951 alpha: 0.5,
1952 beta: 1.0,
1953 l1: 0.0,
1954 l2: 0.0,
1955 max_features: Some(2),
1956 new_feature_policy: NewFeaturePolicy::Ignore,
1957 })
1958 .unwrap();
1959 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1960 model.learn(&sf, true).unwrap();
1961 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1962 model.learn(&sf_mixed, false).unwrap();
1963 assert_eq!(model.feature_count(), 2);
1964 assert!(!model.params.contains_key(&2));
1965 assert_eq!(model.samples_seen(), 2);
1966 }
1967
1968 #[test]
1969 fn classifier_max_features_multi_new_prejudge() {
1970 let mut model = FtrlClassifier::new(FtrlConfig {
1971 alpha: 0.5,
1972 beta: 1.0,
1973 l1: 0.0,
1974 l2: 0.0,
1975 max_features: Some(2),
1976 new_feature_policy: NewFeaturePolicy::Reject,
1977 })
1978 .unwrap();
1979 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1980 assert!(model.learn(&sf, true).is_err());
1981 assert_eq!(model.feature_count(), 0);
1982 assert_eq!(model.samples_seen(), 0);
1983 }
1984
1985 #[test]
1990 #[cfg(feature = "serde")]
1991 fn classifier_serde_rejects_negative_n() {
1992 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}";
1993 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1994 assert!(result.is_err(), "negative n must be rejected");
1995 }
1996
1997 #[test]
1998 #[cfg(feature = "serde")]
1999 fn classifier_serde_rejects_invalid_config() {
2000 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}";
2001 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2002 assert!(result.is_err(), "invalid alpha must be rejected");
2003 }
2004
2005 #[test]
2006 #[cfg(feature = "serde")]
2007 fn classifier_serde_accepts_missing_optional_fields() {
2008 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}";
2009 let model: FtrlClassifier = serde_json::from_str(json).unwrap();
2010 assert!(model.config().max_features.is_none());
2011 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
2012 }
2013
2014 #[test]
2019 fn regressor_gradient_squared_underflow_is_atomic() {
2020 let mut model = FtrlRegressor::new(FtrlConfig {
2026 alpha: 1.0,
2027 beta: 0.0,
2028 l1: 0.0,
2029 l2: 0.0,
2030 max_features: None,
2031 new_feature_policy: NewFeaturePolicy::default(),
2032 })
2033 .unwrap();
2034 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2035 let result = model.learn(&sf, -1e-200);
2036 assert!(result.is_err(), "expected underflow error, got {result:?}");
2037 assert_eq!(model.samples_seen(), 0);
2038 assert_eq!(model.feature_count(), 0);
2039 assert!(model.params.is_empty());
2040 assert_eq!(model.intercept.z, 0.0);
2041 assert_eq!(model.intercept.n, 0.0);
2042 }
2043
2044 #[test]
2045 fn classifier_gradient_squared_underflow_is_atomic() {
2046 let mut model = FtrlClassifier::new(FtrlConfig {
2050 alpha: 1.0,
2051 beta: 0.0,
2052 l1: 0.0,
2053 l2: 0.0,
2054 max_features: None,
2055 new_feature_policy: NewFeaturePolicy::default(),
2056 })
2057 .unwrap();
2058 let sf = SparseFeatures::from_sorted(vec![(0, 1e-200)]).unwrap();
2059 let result = model.learn(&sf, false);
2060 assert!(result.is_err(), "expected underflow error, got {result:?}");
2061 assert_eq!(model.samples_seen(), 0);
2062 assert_eq!(model.feature_count(), 0);
2063 }
2064
2065 #[test]
2066 fn regressor_boundary_config_predict_after_learn_always_finite() {
2067 let mut model = FtrlRegressor::new(FtrlConfig {
2071 alpha: 1.0,
2072 beta: 0.0,
2073 l1: 0.0,
2074 l2: 0.0,
2075 max_features: None,
2076 new_feature_policy: NewFeaturePolicy::default(),
2077 })
2078 .unwrap();
2079 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2080 for _ in 0..100 {
2081 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2082 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2083 let y = 2.0 * x0 - x1;
2084 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2085 model.learn(&sf, y).unwrap();
2086 let pred = model.predict(&sf);
2087 assert!(
2088 pred.is_ok(),
2089 "predict failed after successful learn: {pred:?}"
2090 );
2091 assert!(
2092 pred.unwrap().is_finite(),
2093 "predict must return finite value after successful learn"
2094 );
2095 }
2096 }
2097
2098 #[test]
2099 fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2100 let mut model = FtrlClassifier::new(FtrlConfig {
2101 alpha: 1.0,
2102 beta: 0.0,
2103 l1: 0.0,
2104 l2: 0.0,
2105 max_features: None,
2106 new_feature_policy: NewFeaturePolicy::default(),
2107 })
2108 .unwrap();
2109 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2110 for _ in 0..100 {
2111 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2112 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2113 let y = x0 > 0.0;
2114 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2115 model.learn(&sf, y).unwrap();
2116 let proba = model.predict_proba(&sf);
2117 assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2118 let p = proba.unwrap();
2119 assert!(p.is_finite(), "probability must be finite, got {p}");
2120 assert!(
2121 (0.0..=1.0).contains(&p),
2122 "probability must be in [0,1], got {p}"
2123 );
2124 }
2125 }
2126
2127 #[test]
2128 #[cfg(feature = "serde")]
2129 fn regressor_serde_rejects_n_zero_z_nonzero() {
2130 let json = "{\"config\":{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":0.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2134 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2135 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2136 }
2137
2138 #[test]
2139 #[cfg(feature = "serde")]
2140 fn classifier_serde_rejects_n_zero_z_nonzero() {
2141 let json = "{\"config\":{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":0.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2142 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2143 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2144 }
2145
2146 #[test]
2147 #[cfg(feature = "serde")]
2148 fn regressor_predict_dot_plus_intercept_overflow() {
2149 let z = -f64::MAX * 0.75;
2156 let json = format!(
2157 "{{\"config\":{{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{\"0\":{{\"z\":{0},\"n\":1.0}}}},\"intercept\":{{\"z\":{0},\"n\":1.0}},\"samples_seen\":1}}",
2158 z
2159 );
2160 let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2161 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2162 let result = model.predict(&sf);
2163 assert!(
2164 result.is_err(),
2165 "expected dot+intercept overflow error, got {result:?}"
2166 );
2167 }
2168
2169 #[test]
2170 fn regressor_ignore_skips_overflowing_new_feature() {
2171 let mut model = FtrlRegressor::new(FtrlConfig {
2172 alpha: 0.5,
2173 beta: 1.0,
2174 l1: 0.0,
2175 l2: 0.0,
2176 max_features: Some(1),
2177 new_feature_policy: NewFeaturePolicy::Ignore,
2178 })
2179 .unwrap();
2180 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2181 model.learn(&sf, 1.0).unwrap();
2182 assert_eq!(model.feature_count(), 1);
2183
2184 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2189 let result = model.learn(&sf_mixed, 1.0);
2190 assert!(
2191 result.is_ok(),
2192 "Ignore must skip overflowing new feature, got {result:?}"
2193 );
2194 assert_eq!(model.feature_count(), 1);
2195 assert!(!model.params.contains_key(&1));
2196 assert_eq!(model.samples_seen(), 2);
2197 assert!(model.predict(&sf).is_ok());
2198 }
2199
2200 #[test]
2201 fn classifier_ignore_skips_overflowing_new_feature() {
2202 let mut model = FtrlClassifier::new(FtrlConfig {
2203 alpha: 0.5,
2204 beta: 1.0,
2205 l1: 0.0,
2206 l2: 0.0,
2207 max_features: Some(1),
2208 new_feature_policy: NewFeaturePolicy::Ignore,
2209 })
2210 .unwrap();
2211 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2212 model.learn(&sf, true).unwrap();
2213 assert_eq!(model.feature_count(), 1);
2214
2215 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2216 let result = model.learn(&sf_mixed, true);
2217 assert!(
2218 result.is_ok(),
2219 "Ignore must skip overflowing new feature, got {result:?}"
2220 );
2221 assert_eq!(model.feature_count(), 1);
2222 assert!(!model.params.contains_key(&1));
2223 assert_eq!(model.samples_seen(), 2);
2224 assert!(model.predict_proba(&sf).is_ok());
2225 }
2226
2227 #[test]
2232 #[cfg(feature = "serde")]
2233 fn regressor_serde_rejects_config_dependent_zero_denominator() {
2234 let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1e-300}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2239 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2240 assert!(
2241 result.is_err(),
2242 "config-dependent zero denominator must be rejected"
2243 );
2244 }
2245
2246 #[test]
2247 #[cfg(feature = "serde")]
2248 fn classifier_serde_rejects_config_dependent_zero_denominator() {
2249 let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1e-300}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2250 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2251 assert!(
2252 result.is_err(),
2253 "config-dependent zero denominator must be rejected"
2254 );
2255 }
2256
2257 #[test]
2258 #[cfg(feature = "serde")]
2259 fn regressor_serde_rejects_intercept_zero_denominator() {
2260 let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":1.0,\"n\":1e-300},\"samples_seen\":0}";
2263 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2264 assert!(
2265 result.is_err(),
2266 "intercept zero denominator must be rejected"
2267 );
2268 }
2269
2270 #[test]
2271 #[cfg(feature = "serde")]
2272 fn classifier_serde_rejects_intercept_zero_denominator() {
2273 let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":1.0,\"n\":1e-300},\"samples_seen\":0}";
2274 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2275 assert!(
2276 result.is_err(),
2277 "intercept zero denominator must be rejected"
2278 );
2279 }
2280
2281 #[test]
2282 #[cfg(feature = "serde")]
2283 fn regressor_valid_boundary_state_roundtrips() {
2284 let mut model = FtrlRegressor::new(FtrlConfig {
2288 alpha: 1.0,
2289 beta: 0.0,
2290 l1: 0.0,
2291 l2: 0.0,
2292 max_features: None,
2293 new_feature_policy: NewFeaturePolicy::default(),
2294 })
2295 .unwrap();
2296 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2297 model.learn(&sf, 3.0).unwrap();
2298 model.learn(&sf, 5.0).unwrap();
2299 let json = serde_json::to_string(&model).unwrap();
2300 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2301 assert_eq!(restored.samples_seen(), model.samples_seen());
2302 assert_eq!(restored.feature_count(), model.feature_count());
2303 let p1 = model.predict(&sf).unwrap();
2304 let p2 = restored.predict(&sf).unwrap();
2305 assert!((p1 - p2).abs() < 1e-12);
2306 for (_, w) in restored.weights() {
2308 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2309 }
2310 assert!(restored.intercept().is_finite());
2311 }
2312
2313 #[test]
2314 #[cfg(feature = "serde")]
2315 fn classifier_valid_boundary_state_roundtrips() {
2316 let mut model = FtrlClassifier::new(FtrlConfig {
2317 alpha: 1.0,
2318 beta: 0.0,
2319 l1: 0.0,
2320 l2: 0.0,
2321 max_features: None,
2322 new_feature_policy: NewFeaturePolicy::default(),
2323 })
2324 .unwrap();
2325 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2326 model.learn(&sf, true).unwrap();
2327 model.learn(&sf, false).unwrap();
2328 let json = serde_json::to_string(&model).unwrap();
2329 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2330 assert_eq!(restored.samples_seen(), model.samples_seen());
2331 assert_eq!(restored.feature_count(), model.feature_count());
2332 let p1 = model.predict_proba(&sf).unwrap();
2333 let p2 = restored.predict_proba(&sf).unwrap();
2334 assert!((p1 - p2).abs() < 1e-12);
2335 for (_, w) in restored.weights() {
2336 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2337 }
2338 assert!(restored.intercept().is_finite());
2339 }
2340
2341 #[test]
2346 #[cfg(feature = "serde")]
2347 fn regressor_serde_rejects_params_above_max_features() {
2348 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":1,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1.0},\"1\":{\"z\":1.0,\"n\":1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":1}";
2351 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2352 let err = match result {
2353 Ok(_) => panic!("expected serde error, got Ok"),
2354 Err(e) => e,
2355 };
2356 let msg = err.to_string();
2357 assert!(
2358 msg.contains("max_features") && msg.contains("feature count"),
2359 "error must mention feature count / max_features, got: {msg}"
2360 );
2361 }
2362
2363 #[test]
2364 #[cfg(feature = "serde")]
2365 fn classifier_serde_rejects_params_above_max_features() {
2366 let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":1,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1.0},\"1\":{\"z\":1.0,\"n\":1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":1}";
2367 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2368 let err = match result {
2369 Ok(_) => panic!("expected serde error, got Ok"),
2370 Err(e) => e,
2371 };
2372 let msg = err.to_string();
2373 assert!(
2374 msg.contains("max_features") && msg.contains("feature count"),
2375 "error must mention feature count / max_features, got: {msg}"
2376 );
2377 }
2378
2379 #[test]
2380 #[cfg(feature = "serde")]
2381 fn regressor_serde_accepts_params_equal_to_max_features() {
2382 let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":2,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":3}";
2385 let model: FtrlRegressor =
2386 serde_json::from_str(json).expect("equal count must be accepted");
2387 assert_eq!(model.feature_count(), 2);
2388 assert_eq!(model.samples_seen(), 3);
2389 for (_, w) in model.weights() {
2390 assert!(w.is_finite(), "weight must be finite, got {w}");
2391 }
2392 assert!(model.intercept().is_finite());
2393 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2394 let pred = model.predict(&sf).expect("predict must succeed");
2395 assert!(pred.is_finite(), "prediction must be finite, got {pred}");
2396 }
2397
2398 #[test]
2399 #[cfg(feature = "serde")]
2400 fn classifier_serde_accepts_params_equal_to_max_features() {
2401 let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":2,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":3}";
2402 let model: FtrlClassifier =
2403 serde_json::from_str(json).expect("equal count must be accepted");
2404 assert_eq!(model.feature_count(), 2);
2405 assert_eq!(model.samples_seen(), 3);
2406 for (_, w) in model.weights() {
2407 assert!(w.is_finite(), "weight must be finite, got {w}");
2408 }
2409 assert!(model.intercept().is_finite());
2410 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2411 let p = model
2412 .predict_proba(&sf)
2413 .expect("predict_proba must succeed");
2414 assert!(p.is_finite(), "probability must be finite, got {p}");
2415 assert!(
2416 (0.0..=1.0).contains(&p),
2417 "probability must be in [0,1], got {p}"
2418 );
2419 }
2420
2421 #[test]
2422 #[cfg(feature = "serde")]
2423 fn regressor_serde_allows_unbounded_params_when_max_features_none() {
2424 let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0},\"2\":{\"z\":1.5,\"n\":3.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":6}";
2427 let model: FtrlRegressor =
2428 serde_json::from_str(json).expect("unbounded state must be accepted");
2429 assert_eq!(model.feature_count(), 3);
2430 for (_, w) in model.weights() {
2431 assert!(w.is_finite(), "weight must be finite, got {w}");
2432 }
2433 assert!(model.intercept().is_finite());
2434 }
2435
2436 #[test]
2437 #[cfg(feature = "serde")]
2438 fn classifier_serde_allows_unbounded_params_when_max_features_none() {
2439 let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0},\"2\":{\"z\":1.5,\"n\":3.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":6}";
2440 let model: FtrlClassifier =
2441 serde_json::from_str(json).expect("unbounded state must be accepted");
2442 assert_eq!(model.feature_count(), 3);
2443 for (_, w) in model.weights() {
2444 assert!(w.is_finite(), "weight must be finite, got {w}");
2445 }
2446 assert!(model.intercept().is_finite());
2447 }
2448}