1use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
45use crate::loss::log_loss::sigmoid;
46#[cfg(feature = "serde")]
47use crate::persistence::ValidateState;
48use crate::sparse::{FeatureId, SparseFeatures};
49use crate::traits::{SparseClassifier, SparseRegressor};
50use std::collections::BTreeMap;
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
58#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
59#[non_exhaustive]
60pub enum NewFeaturePolicy {
61 #[default]
65 Reject,
66 Ignore,
70}
71
72#[derive(Debug, Clone)]
78#[cfg_attr(feature = "serde", derive(serde::Serialize))]
79#[non_exhaustive]
80pub struct FtrlConfig {
81 pub alpha: f64,
83 pub beta: f64,
85 pub l1: f64,
87 pub l2: f64,
89 pub max_features: Option<usize>,
96 pub new_feature_policy: NewFeaturePolicy,
99}
100
101impl Default for FtrlConfig {
102 fn default() -> Self {
103 Self {
104 alpha: 0.1,
105 beta: 1.0,
106 l1: 1.0,
107 l2: 1.0,
108 max_features: None,
109 new_feature_policy: NewFeaturePolicy::default(),
110 }
111 }
112}
113
114impl FtrlConfig {
115 pub(crate) fn validate(&self) -> Result<(), RillError> {
117 ensure_finite("alpha", self.alpha)?;
118 ensure_finite("beta", self.beta)?;
119 ensure_finite("l1", self.l1)?;
120 ensure_finite("l2", self.l2)?;
121 if self.alpha <= 0.0 {
122 return Err(RillError::InvalidParameter {
123 name: "alpha",
124 value: self.alpha,
125 });
126 }
127 if self.beta < 0.0 {
128 return Err(RillError::InvalidParameter {
129 name: "beta",
130 value: self.beta,
131 });
132 }
133 if self.l1 < 0.0 {
134 return Err(RillError::InvalidParameter {
135 name: "l1",
136 value: self.l1,
137 });
138 }
139 if self.l2 < 0.0 {
140 return Err(RillError::InvalidParameter {
141 name: "l2",
142 value: self.l2,
143 });
144 }
145 if let Some(max_features) = self.max_features
146 && max_features == 0
147 {
148 return Err(RillError::InvalidParameter {
149 name: "max_features",
150 value: 0.0,
151 });
152 }
153 Ok(())
154 }
155}
156
157#[cfg(feature = "serde")]
158impl<'de> serde::Deserialize<'de> for FtrlConfig {
159 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
160 where
161 D: serde::Deserializer<'de>,
162 {
163 #[derive(serde::Deserialize)]
164 struct FtrlConfigState {
165 alpha: f64,
166 beta: f64,
167 l1: f64,
168 l2: f64,
169 #[serde(default)]
170 max_features: Option<usize>,
171 #[serde(default)]
172 new_feature_policy: NewFeaturePolicy,
173 }
174
175 let state = FtrlConfigState::deserialize(deserializer)?;
176 let config = FtrlConfig {
177 alpha: state.alpha,
178 beta: state.beta,
179 l1: state.l1,
180 l2: state.l2,
181 max_features: state.max_features,
182 new_feature_policy: state.new_feature_policy,
183 };
184 config.validate().map_err(serde::de::Error::custom)?;
185 Ok(config)
186 }
187}
188
189#[derive(Debug, Clone, Default)]
195#[cfg_attr(feature = "serde", derive(serde::Serialize))]
196pub struct FtrlParam {
197 z: f64,
199 n: f64,
201}
202
203impl FtrlParam {
204 fn weight(&self, config: &FtrlConfig) -> f64 {
208 if self.z.abs() <= config.l1 {
209 0.0
210 } else {
211 let sign = self.z.signum();
212 let numerator = -(self.z - sign * config.l1);
213 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
214 numerator / denominator
215 }
216 }
217
218 fn intercept_weight(&self, config: &FtrlConfig) -> f64 {
223 if self.n == 0.0 {
224 0.0
225 } else {
226 let numerator = -self.z;
227 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
228 numerator / denominator
229 }
230 }
231
232 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
251 fn weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
252 ensure_finite("ftrl_z", self.z)?;
253 ensure_finite("ftrl_n", self.n)?;
254 if self.n < 0.0 {
255 return Err(RillError::InvalidState(format!(
256 "ftrl n must be non-negative, got {}",
257 self.n
258 )));
259 }
260 if self.z.abs() <= config.l1 {
262 return Ok(0.0);
263 }
264 let sign = self.z.signum();
265 let numerator = -(self.z - sign * config.l1);
266 ensure_finite("ftrl_weight_numerator", numerator)?;
267 let sqrt_n = self.n.sqrt();
268 ensure_finite("ftrl_weight_sqrt_n", sqrt_n)?;
269 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
270 ensure_finite("ftrl_weight_denominator", denominator)?;
271 if denominator == 0.0 {
272 return Err(RillError::InvalidState(format!(
273 "ftrl weight denominator is zero (z={}, n={}, alpha={}, beta={}, l1={}, l2={})",
274 self.z, self.n, config.alpha, config.beta, config.l1, config.l2
275 )));
276 }
277 let weight = numerator / denominator;
278 ensure_finite("ftrl_weight", weight)?;
279 Ok(weight)
280 }
281
282 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
291 fn intercept_weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
292 ensure_finite("ftrl_z", self.z)?;
293 ensure_finite("ftrl_n", self.n)?;
294 if self.n < 0.0 {
295 return Err(RillError::InvalidState(format!(
296 "ftrl n must be non-negative, got {}",
297 self.n
298 )));
299 }
300 if self.n == 0.0 {
304 if self.z != 0.0 {
305 return Err(RillError::InvalidState(format!(
306 "ftrl intercept has n=0 but z={} (non-zero); cannot produce a finite weight",
307 self.z
308 )));
309 }
310 return Ok(0.0);
311 }
312 let numerator = -self.z;
313 ensure_finite("ftrl_intercept_numerator", numerator)?;
314 let sqrt_n = self.n.sqrt();
315 ensure_finite("ftrl_intercept_sqrt_n", sqrt_n)?;
316 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
317 ensure_finite("ftrl_intercept_denominator", denominator)?;
318 if denominator == 0.0 {
319 return Err(RillError::InvalidState(format!(
320 "ftrl intercept denominator is zero (z={}, n={}, alpha={}, beta={}, l2={})",
321 self.z, self.n, config.alpha, config.beta, config.l2
322 )));
323 }
324 let weight = numerator / denominator;
325 ensure_finite("ftrl_intercept_weight", weight)?;
326 Ok(weight)
327 }
328
329 fn next_updated(
340 &self,
341 gradient: f64,
342 weight: f64,
343 config: &FtrlConfig,
344 ) -> Result<(f64, f64), RillError> {
345 let gradient_sq = gradient * gradient;
346 ensure_finite("ftrl_gradient_squared", gradient_sq)?;
347 if gradient != 0.0 && gradient_sq == 0.0 {
354 return Err(RillError::NonFiniteValue {
355 field: "ftrl_gradient_squared",
356 value: gradient_sq,
357 });
358 }
359 let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
360 let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
361 ensure_finite("ftrl_sigma", sigma)?;
362 let sigma_w = sigma * weight;
363 ensure_finite("ftrl_sigma_weight", sigma_w)?;
364 let z_delta = gradient - sigma_w;
365 ensure_finite("ftrl_z_delta", z_delta)?;
366 let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
367 Ok((z_new, n_new))
368 }
369
370 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
382 fn validate(&self) -> Result<(), RillError> {
383 ensure_finite("ftrl_z", self.z)?;
384 ensure_finite("ftrl_n", self.n)?;
385 if self.n < 0.0 {
386 return Err(RillError::InvalidState(format!(
387 "ftrl n must be non-negative, got {0}",
388 self.n
389 )));
390 }
391 if self.n == 0.0 && self.z != 0.0 {
392 return Err(RillError::InvalidState(format!(
393 "ftrl param has n=0 but z={0} (non-zero); this state cannot \
394 produce a finite weight",
395 self.z
396 )));
397 }
398 Ok(())
399 }
400}
401
402#[cfg(feature = "serde")]
403impl<'de> serde::Deserialize<'de> for FtrlParam {
404 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
405 where
406 D: serde::Deserializer<'de>,
407 {
408 #[derive(serde::Deserialize)]
409 struct FtrlParamState {
410 z: f64,
411 n: f64,
412 }
413
414 let state = FtrlParamState::deserialize(deserializer)?;
415 let param = FtrlParam {
416 z: state.z,
417 n: state.n,
418 };
419 param.validate().map_err(serde::de::Error::custom)?;
420 Ok(param)
421 }
422}
423
424fn compute_dot(
431 params: &BTreeMap<FeatureId, FtrlParam>,
432 config: &FtrlConfig,
433 features: &SparseFeatures,
434) -> Result<f64, RillError> {
435 if features.is_empty() {
436 return Err(RillError::EmptyFeatures);
437 }
438 let mut dot = 0.0;
439 for &(id, value) in features.values() {
440 ensure_finite("sparse_value", value)?;
441 if let Some(param) = params.get(&id) {
442 let w = param.weight(config);
443 ensure_finite("ftrl_weight", w)?;
444 let contribution = w * value;
445 ensure_finite("ftrl_dot_contribution", contribution)?;
446 dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
447 }
448 }
449 Ok(dot)
450}
451
452#[derive(Debug, Clone)]
471#[cfg_attr(feature = "serde", derive(serde::Serialize))]
472pub struct FtrlRegressor {
473 config: FtrlConfig,
474 params: BTreeMap<FeatureId, FtrlParam>,
475 intercept: FtrlParam,
476 samples_seen: u64,
477}
478
479impl FtrlRegressor {
480 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
484 config.validate()?;
485 Ok(Self {
486 config,
487 params: BTreeMap::new(),
488 intercept: FtrlParam::default(),
489 samples_seen: 0,
490 })
491 }
492
493 pub const fn config(&self) -> &FtrlConfig {
495 &self.config
496 }
497
498 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
503 self.params
504 .iter()
505 .map(|(&id, param)| (id, param.weight(&self.config)))
506 .filter(|&(_, w)| w != 0.0)
507 .collect()
508 }
509
510 pub fn intercept(&self) -> f64 {
512 self.intercept.intercept_weight(&self.config)
513 }
514
515 pub fn feature_count(&self) -> usize {
517 self.params.len()
518 }
519
520 fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
522 let dot = compute_dot(&self.params, &self.config, features)?;
523 let intercept = self.intercept.intercept_weight(&self.config);
524 ensure_finite("ftrl_intercept", intercept)?;
525 checked_finite_add(dot, intercept, "ftrl_prediction")
526 }
527}
528
529#[cfg(feature = "serde")]
530impl<'de> serde::Deserialize<'de> for FtrlRegressor {
531 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
532 where
533 D: serde::Deserializer<'de>,
534 {
535 #[derive(serde::Deserialize)]
536 struct FtrlRegressorState {
537 config: FtrlConfig,
538 params: BTreeMap<FeatureId, FtrlParam>,
539 intercept: FtrlParam,
540 samples_seen: u64,
541 }
542
543 let state = FtrlRegressorState::deserialize(deserializer)?;
544 let model = FtrlRegressor {
545 config: state.config,
546 params: state.params,
547 intercept: state.intercept,
548 samples_seen: state.samples_seen,
549 };
550 model
553 .validate_invariants()
554 .map_err(serde::de::Error::custom)?;
555 Ok(model)
556 }
557}
558
559impl FtrlRegressor {
560 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
561 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
562 self.config.validate()?;
569 if let Some(max_features) = self.config.max_features
574 && self.params.len() > max_features
575 {
576 return Err(RillError::InvalidState(format!(
577 "FTRL stored feature count {} exceeds max_features {}",
578 self.params.len(),
579 max_features
580 )));
581 }
582 for (id, param) in &self.params {
583 param.validate()?;
584 param
585 .weight_checked(&self.config)
586 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
587 }
588 self.intercept.validate()?;
589 self.intercept
590 .intercept_weight_checked(&self.config)
591 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
592 Ok(())
593 }
594}
595
596impl SparseRegressor for FtrlRegressor {
597 fn samples_seen(&self) -> u64 {
598 self.samples_seen
599 }
600
601 fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
602 self.predict_inner(features)
603 }
604
605 fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
606 if features.is_empty() {
607 return Err(RillError::EmptyFeatures);
608 }
609 ensure_finite("target", target)?;
610
611 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
613
614 let prediction = self.predict_inner(features)?;
615 ensure_finite("ftrl_prediction", prediction)?;
616 let grad = prediction - target;
617 ensure_finite("ftrl_gradient", grad)?;
618
619 let new_ids_count = features
623 .values()
624 .iter()
625 .filter(|(id, _)| !self.params.contains_key(id))
626 .count();
627 let mut skip_new_features = false;
628 if let Some(max_features) = self.config.max_features {
629 let projected = self.params.len().saturating_add(new_ids_count);
630 if projected > max_features {
631 match self.config.new_feature_policy {
632 NewFeaturePolicy::Reject => {
633 return Err(RillError::InvalidState(format!(
634 "FTRL feature count {projected} exceeds max_features {max_features}"
635 )));
636 }
637 NewFeaturePolicy::Ignore => {
638 skip_new_features = true;
639 }
640 }
641 }
642 }
643
644 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
646 for &(id, value) in features.values() {
647 let is_new = !self.params.contains_key(&id);
652 if is_new && skip_new_features {
653 continue;
654 }
655
656 let g = grad * value;
657 ensure_finite("ftrl_feature_gradient", g)?;
658
659 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
660 let w = param.weight(&self.config);
661 param.next_updated(g, w, &self.config)?
662 } else {
663 let param = FtrlParam::default();
664 let w = param.weight(&self.config);
665 param.next_updated(g, w, &self.config)?
666 };
667 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
672 ensure_finite("ftrl_next_weight", next_w)?;
673 updates.push((id, new_z, new_n));
674 }
675
676 let w_b = self.intercept.intercept_weight(&self.config);
678 let (new_intercept_z, new_intercept_n) =
679 self.intercept.next_updated(grad, w_b, &self.config)?;
680 let next_intercept_w = FtrlParam {
682 z: new_intercept_z,
683 n: new_intercept_n,
684 }
685 .intercept_weight(&self.config);
686 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
687
688 for (id, new_z, new_n) in updates {
690 let param = self.params.entry(id).or_default();
691 param.z = new_z;
692 param.n = new_n;
693 }
694 self.intercept.z = new_intercept_z;
695 self.intercept.n = new_intercept_n;
696 self.samples_seen = next_samples_seen;
697
698 Ok(())
699 }
700
701 fn reset(&mut self) {
702 self.params.clear();
703 self.intercept = FtrlParam::default();
704 self.samples_seen = 0;
705 }
706}
707
708#[derive(Debug, Clone)]
731#[cfg_attr(feature = "serde", derive(serde::Serialize))]
732pub struct FtrlClassifier {
733 config: FtrlConfig,
734 params: BTreeMap<FeatureId, FtrlParam>,
735 intercept: FtrlParam,
736 samples_seen: u64,
737}
738
739impl FtrlClassifier {
740 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
744 config.validate()?;
745 Ok(Self {
746 config,
747 params: BTreeMap::new(),
748 intercept: FtrlParam::default(),
749 samples_seen: 0,
750 })
751 }
752
753 pub const fn config(&self) -> &FtrlConfig {
755 &self.config
756 }
757
758 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
763 self.params
764 .iter()
765 .map(|(&id, param)| (id, param.weight(&self.config)))
766 .filter(|&(_, w)| w != 0.0)
767 .collect()
768 }
769
770 pub fn intercept(&self) -> f64 {
772 self.intercept.intercept_weight(&self.config)
773 }
774
775 pub fn feature_count(&self) -> usize {
777 self.params.len()
778 }
779
780 fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
782 let dot = compute_dot(&self.params, &self.config, features)?;
783 let intercept = self.intercept.intercept_weight(&self.config);
784 ensure_finite("ftrl_intercept", intercept)?;
785 let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
786 Ok(sigmoid(logit))
787 }
788}
789
790#[cfg(feature = "serde")]
791impl<'de> serde::Deserialize<'de> for FtrlClassifier {
792 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
793 where
794 D: serde::Deserializer<'de>,
795 {
796 #[derive(serde::Deserialize)]
797 struct FtrlClassifierState {
798 config: FtrlConfig,
799 params: BTreeMap<FeatureId, FtrlParam>,
800 intercept: FtrlParam,
801 samples_seen: u64,
802 }
803
804 let state = FtrlClassifierState::deserialize(deserializer)?;
805 let model = FtrlClassifier {
806 config: state.config,
807 params: state.params,
808 intercept: state.intercept,
809 samples_seen: state.samples_seen,
810 };
811 model
812 .validate_invariants()
813 .map_err(serde::de::Error::custom)?;
814 Ok(model)
815 }
816}
817
818impl FtrlClassifier {
819 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
820 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
821 self.config.validate()?;
823 if let Some(max_features) = self.config.max_features
826 && self.params.len() > max_features
827 {
828 return Err(RillError::InvalidState(format!(
829 "FTRL stored feature count {} exceeds max_features {}",
830 self.params.len(),
831 max_features
832 )));
833 }
834 for (id, param) in &self.params {
835 param.validate()?;
836 param
837 .weight_checked(&self.config)
838 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
839 }
840 self.intercept.validate()?;
841 self.intercept
842 .intercept_weight_checked(&self.config)
843 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
844 Ok(())
845 }
846}
847
848#[cfg(feature = "serde")]
849impl ValidateState for FtrlConfig {
850 fn validate_state(&self) -> Result<(), RillError> {
851 FtrlConfig::validate(self)
852 }
853}
854
855#[cfg(feature = "serde")]
856impl ValidateState for FtrlRegressor {
857 fn validate_state(&self) -> Result<(), RillError> {
858 FtrlRegressor::validate_invariants(self)
859 }
860}
861
862#[cfg(feature = "serde")]
863impl ValidateState for FtrlClassifier {
864 fn validate_state(&self) -> Result<(), RillError> {
865 FtrlClassifier::validate_invariants(self)
866 }
867}
868
869impl SparseClassifier for FtrlClassifier {
870 fn samples_seen(&self) -> u64 {
871 self.samples_seen
872 }
873
874 fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
875 self.predict_proba_inner(features)
876 }
877
878 fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
879 if features.is_empty() {
880 return Err(RillError::EmptyFeatures);
881 }
882
883 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
884
885 let probability = self.predict_proba_inner(features)?;
886 ensure_finite("ftrl_probability", probability)?;
887 let y = if target { 1.0 } else { 0.0 };
888 let grad = probability - y;
889 ensure_finite("ftrl_gradient", grad)?;
890
891 let new_ids_count = features
892 .values()
893 .iter()
894 .filter(|(id, _)| !self.params.contains_key(id))
895 .count();
896 let mut skip_new_features = false;
897 if let Some(max_features) = self.config.max_features {
898 let projected = self.params.len().saturating_add(new_ids_count);
899 if projected > max_features {
900 match self.config.new_feature_policy {
901 NewFeaturePolicy::Reject => {
902 return Err(RillError::InvalidState(format!(
903 "FTRL feature count {projected} exceeds max_features {max_features}"
904 )));
905 }
906 NewFeaturePolicy::Ignore => {
907 skip_new_features = true;
908 }
909 }
910 }
911 }
912
913 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
914 for &(id, value) in features.values() {
915 let is_new = !self.params.contains_key(&id);
920 if is_new && skip_new_features {
921 continue;
922 }
923
924 let g = grad * value;
925 ensure_finite("ftrl_feature_gradient", g)?;
926
927 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
928 let w = param.weight(&self.config);
929 param.next_updated(g, w, &self.config)?
930 } else {
931 let param = FtrlParam::default();
932 let w = param.weight(&self.config);
933 param.next_updated(g, w, &self.config)?
934 };
935 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
937 ensure_finite("ftrl_next_weight", next_w)?;
938 updates.push((id, new_z, new_n));
939 }
940
941 let w_b = self.intercept.intercept_weight(&self.config);
942 let (new_intercept_z, new_intercept_n) =
943 self.intercept.next_updated(grad, w_b, &self.config)?;
944 let next_intercept_w = FtrlParam {
946 z: new_intercept_z,
947 n: new_intercept_n,
948 }
949 .intercept_weight(&self.config);
950 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
951
952 for (id, new_z, new_n) in updates {
953 let param = self.params.entry(id).or_default();
954 param.z = new_z;
955 param.n = new_n;
956 }
957 self.intercept.z = new_intercept_z;
958 self.intercept.n = new_intercept_n;
959 self.samples_seen = next_samples_seen;
960
961 Ok(())
962 }
963
964 fn reset(&mut self) {
965 self.params.clear();
966 self.intercept = FtrlParam::default();
967 self.samples_seen = 0;
968 }
969}
970
971#[cfg(test)]
972mod tests {
973 use super::*;
974 use rand::SeedableRng;
975
976 #[test]
981 fn cold_start_returns_zero() {
982 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
983 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
984 let pred = model.predict(&sf).unwrap();
985 assert!(pred.abs() < 1e-12);
986 }
987
988 #[test]
989 fn learn_linear_data_converges() {
990 let mut model = FtrlRegressor::new(FtrlConfig {
992 alpha: 0.5,
993 beta: 1.0,
994 l1: 0.0,
995 l2: 0.0,
996 max_features: None,
997 new_feature_policy: NewFeaturePolicy::default(),
998 })
999 .unwrap();
1000 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1001 let mut first_err = 0.0;
1002 let mut last_err = 0.0;
1003 for i in 0..500 {
1004 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1005 let y = 2.0 * x;
1006 let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
1007 let pred = model.predict(&sf).unwrap();
1008 let err = (pred - y).abs();
1009 if i < 10 {
1010 first_err += err;
1011 }
1012 if i >= 490 {
1013 last_err += err;
1014 }
1015 model.learn(&sf, y).unwrap();
1016 }
1017 assert!(last_err < first_err, "error should decrease");
1018 let weights = model.weights();
1019 assert_eq!(weights.len(), 1);
1020 assert!(
1021 (weights[0].1 - 2.0).abs() < 0.5,
1022 "weight should approach 2.0"
1023 );
1024 }
1025
1026 #[test]
1027 fn l1_produces_sparse_weights() {
1028 let mut model = FtrlRegressor::new(FtrlConfig {
1030 alpha: 0.1,
1031 beta: 1.0,
1032 l1: 100.0,
1033 l2: 0.0,
1034 max_features: None,
1035 new_feature_policy: NewFeaturePolicy::default(),
1036 })
1037 .unwrap();
1038 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
1039 for _ in 0..200 {
1040 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1041 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1042 let y = 0.5 * x1;
1043 let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
1044 model.learn(&sf, y).unwrap();
1045 }
1046 let weights = model.weights();
1047 assert!(
1049 weights.is_empty(),
1050 "weights should all be zero, got {weights:?}"
1051 );
1052 }
1053
1054 #[test]
1055 fn dynamic_features() {
1056 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1057 assert_eq!(model.feature_count(), 0);
1058 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1059 model.learn(&sf1, 1.0).unwrap();
1060 assert_eq!(model.feature_count(), 1);
1061 let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1063 model.learn(&sf2, 2.0).unwrap();
1064 assert_eq!(model.feature_count(), 2);
1065 assert!(model.params.contains_key(&0));
1067 assert!(model.params.contains_key(&5));
1068 }
1069
1070 #[test]
1071 fn predict_does_not_update_state() {
1072 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1073 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1074 let _ = model.predict(&sf).unwrap();
1075 assert_eq!(model.samples_seen(), 0);
1076 assert_eq!(model.feature_count(), 0);
1077 model.learn(&sf, 1.0).unwrap();
1079 let count_after_learn = model.feature_count();
1080 let _ = model.predict(&sf).unwrap();
1081 assert_eq!(model.feature_count(), count_after_learn);
1082 assert_eq!(model.samples_seen(), 1);
1083 }
1084
1085 #[test]
1086 fn non_finite_value_rejected() {
1087 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1088 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1090 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1091 assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1092 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1093 assert!(model.predict(&sf).is_ok());
1094 }
1095
1096 #[test]
1097 fn non_finite_target_rejected() {
1098 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1099 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1100 assert!(model.learn(&sf, f64::NAN).is_err());
1101 assert!(model.learn(&sf, f64::INFINITY).is_err());
1102 assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1103 assert_eq!(model.samples_seen(), 0);
1105 }
1106
1107 #[test]
1108 fn empty_features_rejected() {
1109 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1110 let sf = SparseFeatures::new();
1111 assert!(model.predict(&sf).is_err());
1112 assert!(model.learn(&sf, 1.0).is_err());
1113 }
1114
1115 #[test]
1116 fn reset_clears_state() {
1117 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1118 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1119 model.learn(&sf, 3.0).unwrap();
1120 model.learn(&sf, 3.0).unwrap();
1121 assert_eq!(model.samples_seen(), 2);
1122 assert_eq!(model.feature_count(), 2);
1123 model.reset();
1124 assert_eq!(model.samples_seen(), 0);
1125 assert_eq!(model.feature_count(), 0);
1126 assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1127 }
1128
1129 #[test]
1130 fn invalid_config_rejected() {
1131 assert!(
1132 FtrlRegressor::new(FtrlConfig {
1133 alpha: 0.0,
1134 ..FtrlConfig::default()
1135 })
1136 .is_err()
1137 );
1138 assert!(
1139 FtrlRegressor::new(FtrlConfig {
1140 alpha: -1.0,
1141 ..FtrlConfig::default()
1142 })
1143 .is_err()
1144 );
1145 assert!(
1146 FtrlRegressor::new(FtrlConfig {
1147 beta: -1.0,
1148 ..FtrlConfig::default()
1149 })
1150 .is_err()
1151 );
1152 assert!(
1153 FtrlRegressor::new(FtrlConfig {
1154 l1: -1.0,
1155 ..FtrlConfig::default()
1156 })
1157 .is_err()
1158 );
1159 assert!(
1160 FtrlRegressor::new(FtrlConfig {
1161 l2: -1.0,
1162 ..FtrlConfig::default()
1163 })
1164 .is_err()
1165 );
1166 assert!(
1167 FtrlRegressor::new(FtrlConfig {
1168 alpha: f64::NAN,
1169 ..FtrlConfig::default()
1170 })
1171 .is_err()
1172 );
1173 assert!(
1174 FtrlRegressor::new(FtrlConfig {
1175 max_features: Some(0),
1176 ..FtrlConfig::default()
1177 })
1178 .is_err()
1179 );
1180 }
1181
1182 #[test]
1183 #[cfg(feature = "serde")]
1184 fn serde_roundtrip() {
1185 let mut model = FtrlRegressor::new(FtrlConfig {
1186 alpha: 0.2,
1187 beta: 0.5,
1188 l1: 0.5,
1189 l2: 0.5,
1190 max_features: Some(100),
1191 new_feature_policy: NewFeaturePolicy::Reject,
1192 })
1193 .unwrap();
1194 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1195 model.learn(&sf, 5.0).unwrap();
1196 let json = serde_json::to_string(&model).unwrap();
1197 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1198 assert_eq!(restored.samples_seen(), model.samples_seen());
1199 assert_eq!(restored.feature_count(), model.feature_count());
1200 let pred_orig = model.predict(&sf).unwrap();
1201 let pred_restored = restored.predict(&sf).unwrap();
1202 assert!((pred_orig - pred_restored).abs() < 1e-12);
1203 }
1204
1205 #[test]
1206 fn weights_returns_nonzero_only() {
1207 let mut model = FtrlRegressor::new(FtrlConfig {
1208 alpha: 0.5,
1209 beta: 1.0,
1210 l1: 0.0,
1211 l2: 0.0,
1212 max_features: None,
1213 new_feature_policy: NewFeaturePolicy::default(),
1214 })
1215 .unwrap();
1216 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1218 for _ in 0..50 {
1219 model.learn(&sf, 1.0).unwrap();
1220 }
1221 let weights = model.weights();
1222 for &(_, w) in &weights {
1224 assert!(w != 0.0);
1225 }
1226 assert!(weights.iter().any(|&(id, _)| id == 0));
1228 }
1229
1230 #[test]
1231 fn multiple_features() {
1232 let mut model = FtrlRegressor::new(FtrlConfig {
1234 alpha: 0.5,
1235 beta: 1.0,
1236 l1: 0.0,
1237 l2: 0.0,
1238 max_features: None,
1239 new_feature_policy: NewFeaturePolicy::default(),
1240 })
1241 .unwrap();
1242 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1243 for _ in 0..500 {
1244 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1245 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1246 let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1247 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1248 model.learn(&sf, y).unwrap();
1249 }
1250 let weights = model.weights();
1251 assert_eq!(weights.len(), 2);
1252 let w0 = weights
1253 .iter()
1254 .find(|&&(id, _)| id == 0)
1255 .map(|&(_, w)| w)
1256 .unwrap();
1257 let w1 = weights
1258 .iter()
1259 .find(|&&(id, _)| id == 1)
1260 .map(|&(_, w)| w)
1261 .unwrap();
1262 assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1263 assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1264 assert!(
1265 (model.intercept() - 0.5).abs() < 0.5,
1266 "intercept should approach 0.5"
1267 );
1268 }
1269
1270 #[test]
1271 fn intercept_learned() {
1272 let mut model = FtrlRegressor::new(FtrlConfig {
1275 alpha: 0.5,
1276 beta: 1.0,
1277 l1: 0.0,
1278 l2: 0.0,
1279 max_features: None,
1280 new_feature_policy: NewFeaturePolicy::default(),
1281 })
1282 .unwrap();
1283 let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1284 for _ in 0..300 {
1285 model.learn(&sf, 3.0).unwrap();
1286 }
1287 let pred = model.predict(&sf).unwrap();
1288 assert!(
1289 (pred - 3.0).abs() < 0.5,
1290 "prediction should approach 3.0, got {pred}"
1291 );
1292 assert!(
1293 (model.intercept() - 3.0).abs() < 0.5,
1294 "intercept should approach 3.0"
1295 );
1296 assert!(model.weights().is_empty());
1298 }
1299
1300 #[test]
1301 fn high_dim_sparse() {
1302 let mut model = FtrlRegressor::new(FtrlConfig {
1305 alpha: 0.3,
1306 beta: 1.0,
1307 l1: 0.0,
1308 l2: 0.0,
1309 max_features: None,
1310 new_feature_policy: NewFeaturePolicy::default(),
1311 })
1312 .unwrap();
1313 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1314 let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1316 let mut first_err = 0.0;
1317 let mut last_err = 0.0;
1318 for i in 0..2000 {
1319 let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1320 for (j, &w) in true_w.iter().enumerate() {
1321 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1322 active.push((j as u64, x * w));
1323 }
1324 for k in 5..10 {
1326 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1327 active.push((k as u64 + 100, x));
1328 }
1329 active.sort_by_key(|(id, _)| *id);
1330 let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1331 let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1332 let pred = model.predict(&sf).unwrap();
1333 let err = (pred - y).abs();
1334 if i < 20 {
1335 first_err += err;
1336 }
1337 if i >= 1980 {
1338 last_err += err;
1339 }
1340 model.learn(&sf, y).unwrap();
1341 }
1342 assert!(
1343 last_err < first_err,
1344 "error should decrease in high-dim sparse"
1345 );
1346 }
1347
1348 #[test]
1353 fn regressor_overflow_does_not_mutate_state() {
1354 let mut model = FtrlRegressor::new(FtrlConfig {
1358 alpha: 0.1,
1359 beta: 1.0,
1360 l1: 0.0,
1361 l2: 0.0,
1362 max_features: None,
1363 new_feature_policy: NewFeaturePolicy::default(),
1364 })
1365 .unwrap();
1366 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1367 let result = model.learn(&sf, 1e100);
1368 assert!(result.is_err(), "expected overflow error, got {result:?}");
1369 assert_eq!(model.samples_seen(), 0);
1370 assert_eq!(model.feature_count(), 0);
1371 assert!(model.params.is_empty());
1372 assert_eq!(model.intercept.z, 0.0);
1373 assert_eq!(model.intercept.n, 0.0);
1374 }
1375
1376 #[test]
1377 fn regressor_partial_update_is_atomic() {
1378 let mut model = FtrlRegressor::new(FtrlConfig {
1381 alpha: 0.1,
1382 beta: 1.0,
1383 l1: 0.0,
1384 l2: 0.0,
1385 max_features: None,
1386 new_feature_policy: NewFeaturePolicy::default(),
1387 })
1388 .unwrap();
1389 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1390 assert!(model.learn(&sf, 1e100).is_err());
1391 assert!(!model.params.contains_key(&0));
1393 assert!(!model.params.contains_key(&1));
1394 assert_eq!(model.samples_seen(), 0);
1395 }
1396
1397 #[test]
1398 #[cfg(feature = "serde")]
1399 fn regressor_samples_seen_overflow_is_atomic() {
1400 let json = format!(
1401 "{{\"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\":{}}}",
1402 u64::MAX
1403 );
1404 let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1405 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1406 let result = model.learn(&sf, 1.0);
1407 assert!(result.is_err(), "expected counter overflow");
1408 assert_eq!(model.samples_seen(), u64::MAX);
1409 assert_eq!(model.feature_count(), 0);
1410 assert_eq!(model.intercept.z, 0.0);
1411 assert_eq!(model.intercept.n, 0.0);
1412 }
1413
1414 #[test]
1419 fn regressor_max_features_reject_at_limit() {
1420 let mut model = FtrlRegressor::new(FtrlConfig {
1421 alpha: 0.5,
1422 beta: 1.0,
1423 l1: 0.0,
1424 l2: 0.0,
1425 max_features: Some(2),
1426 new_feature_policy: NewFeaturePolicy::Reject,
1427 })
1428 .unwrap();
1429 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1431 model.learn(&sf, 1.0).unwrap();
1432 assert_eq!(model.feature_count(), 2);
1433 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1435 assert!(model.learn(&sf_new, 1.0).is_err());
1436 assert_eq!(model.feature_count(), 2);
1437 assert_eq!(model.samples_seen(), 1);
1438 let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1440 model.learn(&sf_existing, 1.0).unwrap();
1441 assert_eq!(model.feature_count(), 2);
1442 assert_eq!(model.samples_seen(), 2);
1443 }
1444
1445 #[test]
1446 fn regressor_max_features_ignore_skips_new() {
1447 let mut model = FtrlRegressor::new(FtrlConfig {
1448 alpha: 0.5,
1449 beta: 1.0,
1450 l1: 0.0,
1451 l2: 0.0,
1452 max_features: Some(2),
1453 new_feature_policy: NewFeaturePolicy::Ignore,
1454 })
1455 .unwrap();
1456 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1457 model.learn(&sf, 1.0).unwrap();
1458 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1462 model.learn(&sf_mixed, 1.0).unwrap();
1463 assert_eq!(model.feature_count(), 2);
1464 assert!(!model.params.contains_key(&2));
1465 assert_eq!(model.samples_seen(), 2);
1466 }
1467
1468 #[test]
1469 fn regressor_max_features_multi_new_prejudge() {
1470 let mut model = FtrlRegressor::new(FtrlConfig {
1471 alpha: 0.5,
1472 beta: 1.0,
1473 l1: 0.0,
1474 l2: 0.0,
1475 max_features: Some(2),
1476 new_feature_policy: NewFeaturePolicy::Reject,
1477 })
1478 .unwrap();
1479 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1482 assert!(model.learn(&sf, 1.0).is_err());
1483 assert_eq!(model.feature_count(), 0);
1484 assert_eq!(model.samples_seen(), 0);
1485 }
1486
1487 #[test]
1492 #[cfg(feature = "serde")]
1493 fn regressor_serde_rejects_negative_n() {
1494 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}";
1495 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1496 assert!(result.is_err(), "negative n must be rejected");
1497 }
1498
1499 #[test]
1500 #[cfg(feature = "serde")]
1501 fn regressor_serde_rejects_invalid_config() {
1502 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}";
1503 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1504 assert!(result.is_err(), "invalid alpha must be rejected");
1505 }
1506
1507 #[test]
1508 #[cfg(feature = "serde")]
1509 fn regressor_serde_accepts_missing_optional_fields() {
1510 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}";
1512 let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1513 assert!(model.config().max_features.is_none());
1514 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1515 }
1516
1517 #[test]
1522 fn cold_start_returns_0_5() {
1523 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1524 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1525 let p = model.predict_proba(&sf).unwrap();
1526 assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1527 }
1528
1529 #[test]
1530 fn learn_separable_data() {
1531 let mut model = FtrlClassifier::new(FtrlConfig {
1533 alpha: 0.5,
1534 beta: 1.0,
1535 l1: 0.0,
1536 l2: 0.0,
1537 max_features: None,
1538 new_feature_policy: NewFeaturePolicy::default(),
1539 })
1540 .unwrap();
1541 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1542 for _ in 0..1000 {
1543 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1544 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1545 let y = x0 > 0.0;
1546 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1547 model.learn(&sf, y).unwrap();
1548 }
1549 let p_pos = model
1550 .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1551 .unwrap();
1552 let p_neg = model
1553 .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1554 .unwrap();
1555 assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1556 assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1557 }
1558
1559 #[test]
1560 fn classifier_l1_produces_sparse_weights() {
1561 let mut model = FtrlClassifier::new(FtrlConfig {
1562 alpha: 0.1,
1563 beta: 1.0,
1564 l1: 100.0,
1565 l2: 0.0,
1566 max_features: None,
1567 new_feature_policy: NewFeaturePolicy::default(),
1568 })
1569 .unwrap();
1570 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1571 for _ in 0..200 {
1572 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1573 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1574 let y = x0 > 0.0;
1575 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1576 model.learn(&sf, y).unwrap();
1577 }
1578 let weights = model.weights();
1579 assert!(
1580 weights.is_empty(),
1581 "weights should all be zero with high L1, got {weights:?}"
1582 );
1583 }
1584
1585 #[test]
1586 fn classifier_dynamic_features() {
1587 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1588 assert_eq!(model.feature_count(), 0);
1589 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1590 model.learn(&sf1, true).unwrap();
1591 assert_eq!(model.feature_count(), 1);
1592 let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1593 model.learn(&sf2, false).unwrap();
1594 assert_eq!(model.feature_count(), 2);
1595 }
1596
1597 #[test]
1598 fn classifier_predict_does_not_update_state() {
1599 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1600 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1601 let _ = model.predict_proba(&sf).unwrap();
1602 assert_eq!(model.samples_seen(), 0);
1603 assert_eq!(model.feature_count(), 0);
1604 model.learn(&sf, true).unwrap();
1605 let count = model.feature_count();
1606 let _ = model.predict_proba(&sf).unwrap();
1607 assert_eq!(model.feature_count(), count);
1608 assert_eq!(model.samples_seen(), 1);
1609 }
1610
1611 #[test]
1612 fn classifier_non_finite_value_rejected() {
1613 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1614 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1615 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1616 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1617 assert!(model.predict_proba(&sf).is_ok());
1618 }
1619
1620 #[test]
1621 fn classifier_empty_features_rejected() {
1622 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1623 let sf = SparseFeatures::new();
1624 assert!(model.predict_proba(&sf).is_err());
1625 assert!(model.learn(&sf, true).is_err());
1626 }
1627
1628 #[test]
1629 fn classifier_reset_clears_state() {
1630 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1631 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1632 model.learn(&sf, true).unwrap();
1633 model.learn(&sf, false).unwrap();
1634 assert_eq!(model.samples_seen(), 2);
1635 assert!(model.feature_count() > 0);
1636 model.reset();
1637 assert_eq!(model.samples_seen(), 0);
1638 assert_eq!(model.feature_count(), 0);
1639 let p = model.predict_proba(&sf).unwrap();
1640 assert!((p - 0.5).abs() < 1e-12);
1641 }
1642
1643 #[test]
1644 fn classifier_invalid_config_rejected() {
1645 assert!(
1646 FtrlClassifier::new(FtrlConfig {
1647 alpha: 0.0,
1648 ..FtrlConfig::default()
1649 })
1650 .is_err()
1651 );
1652 assert!(
1653 FtrlClassifier::new(FtrlConfig {
1654 beta: -0.1,
1655 ..FtrlConfig::default()
1656 })
1657 .is_err()
1658 );
1659 assert!(
1660 FtrlClassifier::new(FtrlConfig {
1661 l1: -1.0,
1662 ..FtrlConfig::default()
1663 })
1664 .is_err()
1665 );
1666 assert!(
1667 FtrlClassifier::new(FtrlConfig {
1668 l2: -1.0,
1669 ..FtrlConfig::default()
1670 })
1671 .is_err()
1672 );
1673 assert!(
1674 FtrlClassifier::new(FtrlConfig {
1675 alpha: f64::INFINITY,
1676 ..FtrlConfig::default()
1677 })
1678 .is_err()
1679 );
1680 }
1681
1682 #[test]
1683 #[cfg(feature = "serde")]
1684 fn classifier_serde_roundtrip() {
1685 let mut model = FtrlClassifier::new(FtrlConfig {
1686 alpha: 0.3,
1687 beta: 0.5,
1688 l1: 0.1,
1689 l2: 0.2,
1690 max_features: Some(100),
1691 new_feature_policy: NewFeaturePolicy::Reject,
1692 })
1693 .unwrap();
1694 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1695 model.learn(&sf, true).unwrap();
1696 model.learn(&sf, false).unwrap();
1697 let json = serde_json::to_string(&model).unwrap();
1698 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1699 assert_eq!(restored.samples_seen(), model.samples_seen());
1700 assert_eq!(restored.feature_count(), model.feature_count());
1701 let p1 = model.predict_proba(&sf).unwrap();
1702 let p2 = restored.predict_proba(&sf).unwrap();
1703 assert!((p1 - p2).abs() < 1e-12);
1704 }
1705
1706 #[test]
1707 fn predict_proba_in_range() {
1708 let mut model = FtrlClassifier::new(FtrlConfig {
1709 alpha: 0.5,
1710 beta: 1.0,
1711 l1: 0.0,
1712 l2: 0.0,
1713 max_features: None,
1714 new_feature_policy: NewFeaturePolicy::default(),
1715 })
1716 .unwrap();
1717 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1718 for _ in 0..200 {
1719 let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1720 let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1721 let y = x0 > 0.0;
1722 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1723 model.learn(&sf, y).unwrap();
1724 let p = model.predict_proba(&sf).unwrap();
1725 assert!(
1726 (0.0..=1.0).contains(&p),
1727 "probability must be in [0,1], got {p}"
1728 );
1729 }
1730 }
1731
1732 #[test]
1733 fn learn_improves_accuracy() {
1734 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1735 let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1737 .map(|_| {
1738 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1739 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1740 let y = x0 + x1 > 0.0;
1741 (
1742 SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1743 y,
1744 )
1745 })
1746 .collect();
1747
1748 let mut model = FtrlClassifier::new(FtrlConfig {
1749 alpha: 0.5,
1750 beta: 1.0,
1751 l1: 0.0,
1752 l2: 0.0,
1753 max_features: None,
1754 new_feature_policy: NewFeaturePolicy::default(),
1755 })
1756 .unwrap();
1757
1758 let acc_before: f64 = test_set
1760 .iter()
1761 .map(|(sf, y)| {
1762 let pred = model.predict(sf).unwrap();
1763 if pred == *y { 1.0 } else { 0.0 }
1764 })
1765 .sum::<f64>()
1766 / test_set.len() as f64;
1767
1768 for _ in 0..1000 {
1770 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1771 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1772 let y = x0 + x1 > 0.0;
1773 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1774 model.learn(&sf, y).unwrap();
1775 }
1776
1777 let acc_after: f64 = test_set
1778 .iter()
1779 .map(|(sf, y)| {
1780 let pred = model.predict(sf).unwrap();
1781 if pred == *y { 1.0 } else { 0.0 }
1782 })
1783 .sum::<f64>()
1784 / test_set.len() as f64;
1785
1786 assert!(
1787 acc_after > acc_before,
1788 "accuracy should improve: {acc_before} -> {acc_after}"
1789 );
1790 }
1791
1792 #[test]
1793 fn classifier_weights_returns_nonzero_only() {
1794 let mut model = FtrlClassifier::new(FtrlConfig {
1795 alpha: 0.5,
1796 beta: 1.0,
1797 l1: 0.0,
1798 l2: 0.0,
1799 max_features: None,
1800 new_feature_policy: NewFeaturePolicy::default(),
1801 })
1802 .unwrap();
1803 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1804 for _ in 0..50 {
1805 model.learn(&sf, true).unwrap();
1806 }
1807 let weights = model.weights();
1808 for &(_, w) in &weights {
1809 assert!(w != 0.0);
1810 }
1811 }
1812
1813 #[test]
1814 fn classifier_multiple_features() {
1815 let mut model = FtrlClassifier::new(FtrlConfig {
1816 alpha: 0.5,
1817 beta: 1.0,
1818 l1: 0.0,
1819 l2: 0.0,
1820 max_features: None,
1821 new_feature_policy: NewFeaturePolicy::default(),
1822 })
1823 .unwrap();
1824 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1825 for _ in 0..1000 {
1826 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1827 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1828 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1829 let y = x0 + x1 > 0.0;
1831 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1832 model.learn(&sf, y).unwrap();
1833 }
1834 let weights = model.weights();
1835 assert!(weights.iter().any(|&(id, _)| id == 0));
1837 assert!(weights.iter().any(|&(id, _)| id == 1));
1838 let p_pos = model
1840 .predict_proba(
1841 &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1842 )
1843 .unwrap();
1844 let p_neg = model
1845 .predict_proba(
1846 &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1847 )
1848 .unwrap();
1849 assert!(p_pos > 0.8);
1850 assert!(p_neg < 0.2);
1851 }
1852
1853 #[test]
1854 fn log_loss_converges() {
1855 let mut model = FtrlClassifier::new(FtrlConfig {
1857 alpha: 0.5,
1858 beta: 1.0,
1859 l1: 0.0,
1860 l2: 0.0,
1861 max_features: None,
1862 new_feature_policy: NewFeaturePolicy::default(),
1863 })
1864 .unwrap();
1865 let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1866 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1867 let mut first_loss = 0.0;
1868 let mut last_loss = 0.0;
1869 for i in 0..1000 {
1870 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1871 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1872 let y = x0 > 0.0;
1873 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1874 let p = model.predict_proba(&sf).unwrap();
1875 let loss = loss_fn.loss(p, y);
1876 if i < 20 {
1877 first_loss += loss;
1878 }
1879 if i >= 980 {
1880 last_loss += loss;
1881 }
1882 model.learn(&sf, y).unwrap();
1883 }
1884 assert!(last_loss < first_loss, "log loss should decrease");
1885 }
1886
1887 #[test]
1892 fn classifier_overflow_does_not_mutate_state() {
1893 let mut model = FtrlClassifier::new(FtrlConfig {
1894 alpha: 0.1,
1895 beta: 1.0,
1896 l1: 0.0,
1897 l2: 0.0,
1898 max_features: None,
1899 new_feature_policy: NewFeaturePolicy::default(),
1900 })
1901 .unwrap();
1902 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1903 let result = model.learn(&sf, false);
1906 assert!(result.is_err(), "expected overflow error, got {result:?}");
1907 assert_eq!(model.samples_seen(), 0);
1908 assert_eq!(model.feature_count(), 0);
1909 assert!(model.params.is_empty());
1910 assert_eq!(model.intercept.z, 0.0);
1911 assert_eq!(model.intercept.n, 0.0);
1912 }
1913
1914 #[test]
1915 fn classifier_partial_update_is_atomic() {
1916 let mut model = FtrlClassifier::new(FtrlConfig {
1917 alpha: 0.1,
1918 beta: 1.0,
1919 l1: 0.0,
1920 l2: 0.0,
1921 max_features: None,
1922 new_feature_policy: NewFeaturePolicy::default(),
1923 })
1924 .unwrap();
1925 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1926 assert!(model.learn(&sf, false).is_err());
1927 assert!(!model.params.contains_key(&0));
1928 assert!(!model.params.contains_key(&1));
1929 assert_eq!(model.samples_seen(), 0);
1930 }
1931
1932 #[test]
1933 #[cfg(feature = "serde")]
1934 fn classifier_samples_seen_overflow_is_atomic() {
1935 let json = format!(
1936 "{{\"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\":{}}}",
1937 u64::MAX
1938 );
1939 let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1940 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1941 let result = model.learn(&sf, true);
1942 assert!(result.is_err(), "expected counter overflow");
1943 assert_eq!(model.samples_seen(), u64::MAX);
1944 assert_eq!(model.feature_count(), 0);
1945 assert_eq!(model.intercept.z, 0.0);
1946 assert_eq!(model.intercept.n, 0.0);
1947 }
1948
1949 #[test]
1954 fn classifier_max_features_reject_at_limit() {
1955 let mut model = FtrlClassifier::new(FtrlConfig {
1956 alpha: 0.5,
1957 beta: 1.0,
1958 l1: 0.0,
1959 l2: 0.0,
1960 max_features: Some(2),
1961 new_feature_policy: NewFeaturePolicy::Reject,
1962 })
1963 .unwrap();
1964 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1965 model.learn(&sf, true).unwrap();
1966 assert_eq!(model.feature_count(), 2);
1967 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1968 assert!(model.learn(&sf_new, true).is_err());
1969 assert_eq!(model.feature_count(), 2);
1970 assert_eq!(model.samples_seen(), 1);
1971 }
1972
1973 #[test]
1974 fn classifier_max_features_ignore_skips_new() {
1975 let mut model = FtrlClassifier::new(FtrlConfig {
1976 alpha: 0.5,
1977 beta: 1.0,
1978 l1: 0.0,
1979 l2: 0.0,
1980 max_features: Some(2),
1981 new_feature_policy: NewFeaturePolicy::Ignore,
1982 })
1983 .unwrap();
1984 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1985 model.learn(&sf, true).unwrap();
1986 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1987 model.learn(&sf_mixed, false).unwrap();
1988 assert_eq!(model.feature_count(), 2);
1989 assert!(!model.params.contains_key(&2));
1990 assert_eq!(model.samples_seen(), 2);
1991 }
1992
1993 #[test]
1994 fn classifier_max_features_multi_new_prejudge() {
1995 let mut model = FtrlClassifier::new(FtrlConfig {
1996 alpha: 0.5,
1997 beta: 1.0,
1998 l1: 0.0,
1999 l2: 0.0,
2000 max_features: Some(2),
2001 new_feature_policy: NewFeaturePolicy::Reject,
2002 })
2003 .unwrap();
2004 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
2005 assert!(model.learn(&sf, true).is_err());
2006 assert_eq!(model.feature_count(), 0);
2007 assert_eq!(model.samples_seen(), 0);
2008 }
2009
2010 #[test]
2015 #[cfg(feature = "serde")]
2016 fn classifier_serde_rejects_negative_n() {
2017 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}";
2018 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2019 assert!(result.is_err(), "negative n must be rejected");
2020 }
2021
2022 #[test]
2023 #[cfg(feature = "serde")]
2024 fn classifier_serde_rejects_invalid_config() {
2025 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}";
2026 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2027 assert!(result.is_err(), "invalid alpha must be rejected");
2028 }
2029
2030 #[test]
2031 #[cfg(feature = "serde")]
2032 fn classifier_serde_accepts_missing_optional_fields() {
2033 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}";
2034 let model: FtrlClassifier = serde_json::from_str(json).unwrap();
2035 assert!(model.config().max_features.is_none());
2036 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
2037 }
2038
2039 #[test]
2044 fn regressor_gradient_squared_underflow_is_atomic() {
2045 let mut model = FtrlRegressor::new(FtrlConfig {
2051 alpha: 1.0,
2052 beta: 0.0,
2053 l1: 0.0,
2054 l2: 0.0,
2055 max_features: None,
2056 new_feature_policy: NewFeaturePolicy::default(),
2057 })
2058 .unwrap();
2059 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2060 let result = model.learn(&sf, -1e-200);
2061 assert!(result.is_err(), "expected underflow error, got {result:?}");
2062 assert_eq!(model.samples_seen(), 0);
2063 assert_eq!(model.feature_count(), 0);
2064 assert!(model.params.is_empty());
2065 assert_eq!(model.intercept.z, 0.0);
2066 assert_eq!(model.intercept.n, 0.0);
2067 }
2068
2069 #[test]
2070 fn classifier_gradient_squared_underflow_is_atomic() {
2071 let mut model = FtrlClassifier::new(FtrlConfig {
2075 alpha: 1.0,
2076 beta: 0.0,
2077 l1: 0.0,
2078 l2: 0.0,
2079 max_features: None,
2080 new_feature_policy: NewFeaturePolicy::default(),
2081 })
2082 .unwrap();
2083 let sf = SparseFeatures::from_sorted(vec![(0, 1e-200)]).unwrap();
2084 let result = model.learn(&sf, false);
2085 assert!(result.is_err(), "expected underflow error, got {result:?}");
2086 assert_eq!(model.samples_seen(), 0);
2087 assert_eq!(model.feature_count(), 0);
2088 }
2089
2090 #[test]
2091 fn regressor_boundary_config_predict_after_learn_always_finite() {
2092 let mut model = FtrlRegressor::new(FtrlConfig {
2096 alpha: 1.0,
2097 beta: 0.0,
2098 l1: 0.0,
2099 l2: 0.0,
2100 max_features: None,
2101 new_feature_policy: NewFeaturePolicy::default(),
2102 })
2103 .unwrap();
2104 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2105 for _ in 0..100 {
2106 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2107 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2108 let y = 2.0 * x0 - x1;
2109 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2110 model.learn(&sf, y).unwrap();
2111 let pred = model.predict(&sf);
2112 assert!(
2113 pred.is_ok(),
2114 "predict failed after successful learn: {pred:?}"
2115 );
2116 assert!(
2117 pred.unwrap().is_finite(),
2118 "predict must return finite value after successful learn"
2119 );
2120 }
2121 }
2122
2123 #[test]
2124 fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2125 let mut model = FtrlClassifier::new(FtrlConfig {
2126 alpha: 1.0,
2127 beta: 0.0,
2128 l1: 0.0,
2129 l2: 0.0,
2130 max_features: None,
2131 new_feature_policy: NewFeaturePolicy::default(),
2132 })
2133 .unwrap();
2134 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2135 for _ in 0..100 {
2136 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2137 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2138 let y = x0 > 0.0;
2139 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2140 model.learn(&sf, y).unwrap();
2141 let proba = model.predict_proba(&sf);
2142 assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2143 let p = proba.unwrap();
2144 assert!(p.is_finite(), "probability must be finite, got {p}");
2145 assert!(
2146 (0.0..=1.0).contains(&p),
2147 "probability must be in [0,1], got {p}"
2148 );
2149 }
2150 }
2151
2152 #[test]
2153 #[cfg(feature = "serde")]
2154 fn regressor_serde_rejects_n_zero_z_nonzero() {
2155 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}";
2159 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2160 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2161 }
2162
2163 #[test]
2164 #[cfg(feature = "serde")]
2165 fn classifier_serde_rejects_n_zero_z_nonzero() {
2166 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}";
2167 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2168 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2169 }
2170
2171 #[test]
2172 #[cfg(feature = "serde")]
2173 fn regressor_predict_dot_plus_intercept_overflow() {
2174 let z = -f64::MAX * 0.75;
2181 let json = format!(
2182 "{{\"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}}",
2183 z
2184 );
2185 let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2186 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2187 let result = model.predict(&sf);
2188 assert!(
2189 result.is_err(),
2190 "expected dot+intercept overflow error, got {result:?}"
2191 );
2192 }
2193
2194 #[test]
2195 fn regressor_ignore_skips_overflowing_new_feature() {
2196 let mut model = FtrlRegressor::new(FtrlConfig {
2197 alpha: 0.5,
2198 beta: 1.0,
2199 l1: 0.0,
2200 l2: 0.0,
2201 max_features: Some(1),
2202 new_feature_policy: NewFeaturePolicy::Ignore,
2203 })
2204 .unwrap();
2205 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2206 model.learn(&sf, 1.0).unwrap();
2207 assert_eq!(model.feature_count(), 1);
2208
2209 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2214 let result = model.learn(&sf_mixed, 1.0);
2215 assert!(
2216 result.is_ok(),
2217 "Ignore must skip overflowing new feature, got {result:?}"
2218 );
2219 assert_eq!(model.feature_count(), 1);
2220 assert!(!model.params.contains_key(&1));
2221 assert_eq!(model.samples_seen(), 2);
2222 assert!(model.predict(&sf).is_ok());
2223 }
2224
2225 #[test]
2226 fn classifier_ignore_skips_overflowing_new_feature() {
2227 let mut model = FtrlClassifier::new(FtrlConfig {
2228 alpha: 0.5,
2229 beta: 1.0,
2230 l1: 0.0,
2231 l2: 0.0,
2232 max_features: Some(1),
2233 new_feature_policy: NewFeaturePolicy::Ignore,
2234 })
2235 .unwrap();
2236 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2237 model.learn(&sf, true).unwrap();
2238 assert_eq!(model.feature_count(), 1);
2239
2240 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2241 let result = model.learn(&sf_mixed, true);
2242 assert!(
2243 result.is_ok(),
2244 "Ignore must skip overflowing new feature, got {result:?}"
2245 );
2246 assert_eq!(model.feature_count(), 1);
2247 assert!(!model.params.contains_key(&1));
2248 assert_eq!(model.samples_seen(), 2);
2249 assert!(model.predict_proba(&sf).is_ok());
2250 }
2251
2252 #[test]
2257 #[cfg(feature = "serde")]
2258 fn regressor_serde_rejects_config_dependent_zero_denominator() {
2259 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}";
2264 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2265 assert!(
2266 result.is_err(),
2267 "config-dependent zero denominator must be rejected"
2268 );
2269 }
2270
2271 #[test]
2272 #[cfg(feature = "serde")]
2273 fn classifier_serde_rejects_config_dependent_zero_denominator() {
2274 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}";
2275 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2276 assert!(
2277 result.is_err(),
2278 "config-dependent zero denominator must be rejected"
2279 );
2280 }
2281
2282 #[test]
2283 #[cfg(feature = "serde")]
2284 fn regressor_serde_rejects_intercept_zero_denominator() {
2285 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}";
2288 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2289 assert!(
2290 result.is_err(),
2291 "intercept zero denominator must be rejected"
2292 );
2293 }
2294
2295 #[test]
2296 #[cfg(feature = "serde")]
2297 fn classifier_serde_rejects_intercept_zero_denominator() {
2298 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}";
2299 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2300 assert!(
2301 result.is_err(),
2302 "intercept zero denominator must be rejected"
2303 );
2304 }
2305
2306 #[test]
2307 #[cfg(feature = "serde")]
2308 fn regressor_valid_boundary_state_roundtrips() {
2309 let mut model = FtrlRegressor::new(FtrlConfig {
2313 alpha: 1.0,
2314 beta: 0.0,
2315 l1: 0.0,
2316 l2: 0.0,
2317 max_features: None,
2318 new_feature_policy: NewFeaturePolicy::default(),
2319 })
2320 .unwrap();
2321 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2322 model.learn(&sf, 3.0).unwrap();
2323 model.learn(&sf, 5.0).unwrap();
2324 let json = serde_json::to_string(&model).unwrap();
2325 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2326 assert_eq!(restored.samples_seen(), model.samples_seen());
2327 assert_eq!(restored.feature_count(), model.feature_count());
2328 let p1 = model.predict(&sf).unwrap();
2329 let p2 = restored.predict(&sf).unwrap();
2330 assert!((p1 - p2).abs() < 1e-12);
2331 for (_, w) in restored.weights() {
2333 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2334 }
2335 assert!(restored.intercept().is_finite());
2336 }
2337
2338 #[test]
2339 #[cfg(feature = "serde")]
2340 fn classifier_valid_boundary_state_roundtrips() {
2341 let mut model = FtrlClassifier::new(FtrlConfig {
2342 alpha: 1.0,
2343 beta: 0.0,
2344 l1: 0.0,
2345 l2: 0.0,
2346 max_features: None,
2347 new_feature_policy: NewFeaturePolicy::default(),
2348 })
2349 .unwrap();
2350 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2351 model.learn(&sf, true).unwrap();
2352 model.learn(&sf, false).unwrap();
2353 let json = serde_json::to_string(&model).unwrap();
2354 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2355 assert_eq!(restored.samples_seen(), model.samples_seen());
2356 assert_eq!(restored.feature_count(), model.feature_count());
2357 let p1 = model.predict_proba(&sf).unwrap();
2358 let p2 = restored.predict_proba(&sf).unwrap();
2359 assert!((p1 - p2).abs() < 1e-12);
2360 for (_, w) in restored.weights() {
2361 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2362 }
2363 assert!(restored.intercept().is_finite());
2364 }
2365
2366 #[test]
2371 #[cfg(feature = "serde")]
2372 fn regressor_serde_rejects_params_above_max_features() {
2373 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}";
2376 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2377 let err = match result {
2378 Ok(_) => panic!("expected serde error, got Ok"),
2379 Err(e) => e,
2380 };
2381 let msg = err.to_string();
2382 assert!(
2383 msg.contains("max_features") && msg.contains("feature count"),
2384 "error must mention feature count / max_features, got: {msg}"
2385 );
2386 }
2387
2388 #[test]
2389 #[cfg(feature = "serde")]
2390 fn classifier_serde_rejects_params_above_max_features() {
2391 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}";
2392 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2393 let err = match result {
2394 Ok(_) => panic!("expected serde error, got Ok"),
2395 Err(e) => e,
2396 };
2397 let msg = err.to_string();
2398 assert!(
2399 msg.contains("max_features") && msg.contains("feature count"),
2400 "error must mention feature count / max_features, got: {msg}"
2401 );
2402 }
2403
2404 #[test]
2405 #[cfg(feature = "serde")]
2406 fn regressor_serde_accepts_params_equal_to_max_features() {
2407 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}";
2410 let model: FtrlRegressor =
2411 serde_json::from_str(json).expect("equal count must be accepted");
2412 assert_eq!(model.feature_count(), 2);
2413 assert_eq!(model.samples_seen(), 3);
2414 for (_, w) in model.weights() {
2415 assert!(w.is_finite(), "weight must be finite, got {w}");
2416 }
2417 assert!(model.intercept().is_finite());
2418 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2419 let pred = model.predict(&sf).expect("predict must succeed");
2420 assert!(pred.is_finite(), "prediction must be finite, got {pred}");
2421 }
2422
2423 #[test]
2424 #[cfg(feature = "serde")]
2425 fn classifier_serde_accepts_params_equal_to_max_features() {
2426 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}";
2427 let model: FtrlClassifier =
2428 serde_json::from_str(json).expect("equal count must be accepted");
2429 assert_eq!(model.feature_count(), 2);
2430 assert_eq!(model.samples_seen(), 3);
2431 for (_, w) in model.weights() {
2432 assert!(w.is_finite(), "weight must be finite, got {w}");
2433 }
2434 assert!(model.intercept().is_finite());
2435 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2436 let p = model
2437 .predict_proba(&sf)
2438 .expect("predict_proba must succeed");
2439 assert!(p.is_finite(), "probability must be finite, got {p}");
2440 assert!(
2441 (0.0..=1.0).contains(&p),
2442 "probability must be in [0,1], got {p}"
2443 );
2444 }
2445
2446 #[test]
2447 #[cfg(feature = "serde")]
2448 fn regressor_serde_allows_unbounded_params_when_max_features_none() {
2449 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}";
2452 let model: FtrlRegressor =
2453 serde_json::from_str(json).expect("unbounded state must be accepted");
2454 assert_eq!(model.feature_count(), 3);
2455 for (_, w) in model.weights() {
2456 assert!(w.is_finite(), "weight must be finite, got {w}");
2457 }
2458 assert!(model.intercept().is_finite());
2459 }
2460
2461 #[test]
2462 #[cfg(feature = "serde")]
2463 fn classifier_serde_allows_unbounded_params_when_max_features_none() {
2464 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}";
2465 let model: FtrlClassifier =
2466 serde_json::from_str(json).expect("unbounded state must be accepted");
2467 assert_eq!(model.feature_count(), 3);
2468 for (_, w) in model.weights() {
2469 assert!(w.is_finite(), "weight must be finite, got {w}");
2470 }
2471 assert!(model.intercept().is_finite());
2472 }
2473}