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 for (id, param) in &self.params {
566 param.validate()?;
567 param
568 .weight_checked(&self.config)
569 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
570 }
571 self.intercept.validate()?;
572 self.intercept
573 .intercept_weight_checked(&self.config)
574 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
575 Ok(())
576 }
577}
578
579impl SparseRegressor for FtrlRegressor {
580 fn samples_seen(&self) -> u64 {
581 self.samples_seen
582 }
583
584 fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
585 self.predict_inner(features)
586 }
587
588 fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
589 if features.is_empty() {
590 return Err(RillError::EmptyFeatures);
591 }
592 ensure_finite("target", target)?;
593
594 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
596
597 let prediction = self.predict_inner(features)?;
598 ensure_finite("ftrl_prediction", prediction)?;
599 let grad = prediction - target;
600 ensure_finite("ftrl_gradient", grad)?;
601
602 let new_ids_count = features
606 .values()
607 .iter()
608 .filter(|(id, _)| !self.params.contains_key(id))
609 .count();
610 let mut skip_new_features = false;
611 if let Some(max_features) = self.config.max_features {
612 let projected = self.params.len().saturating_add(new_ids_count);
613 if projected > max_features {
614 match self.config.new_feature_policy {
615 NewFeaturePolicy::Reject => {
616 return Err(RillError::InvalidState(format!(
617 "FTRL feature count {projected} exceeds max_features {max_features}"
618 )));
619 }
620 NewFeaturePolicy::Ignore => {
621 skip_new_features = true;
622 }
623 }
624 }
625 }
626
627 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
629 for &(id, value) in features.values() {
630 let is_new = !self.params.contains_key(&id);
635 if is_new && skip_new_features {
636 continue;
637 }
638
639 let g = grad * value;
640 ensure_finite("ftrl_feature_gradient", g)?;
641
642 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
643 let w = param.weight(&self.config);
644 param.next_updated(g, w, &self.config)?
645 } else {
646 let param = FtrlParam::default();
647 let w = param.weight(&self.config);
648 param.next_updated(g, w, &self.config)?
649 };
650 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
655 ensure_finite("ftrl_next_weight", next_w)?;
656 updates.push((id, new_z, new_n));
657 }
658
659 let w_b = self.intercept.intercept_weight(&self.config);
661 let (new_intercept_z, new_intercept_n) =
662 self.intercept.next_updated(grad, w_b, &self.config)?;
663 let next_intercept_w = FtrlParam {
665 z: new_intercept_z,
666 n: new_intercept_n,
667 }
668 .intercept_weight(&self.config);
669 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
670
671 for (id, new_z, new_n) in updates {
673 let param = self.params.entry(id).or_default();
674 param.z = new_z;
675 param.n = new_n;
676 }
677 self.intercept.z = new_intercept_z;
678 self.intercept.n = new_intercept_n;
679 self.samples_seen = next_samples_seen;
680
681 Ok(())
682 }
683
684 fn reset(&mut self) {
685 self.params.clear();
686 self.intercept = FtrlParam::default();
687 self.samples_seen = 0;
688 }
689}
690
691#[derive(Debug, Clone)]
714#[cfg_attr(feature = "serde", derive(serde::Serialize))]
715pub struct FtrlClassifier {
716 config: FtrlConfig,
717 params: BTreeMap<FeatureId, FtrlParam>,
718 intercept: FtrlParam,
719 samples_seen: u64,
720}
721
722impl FtrlClassifier {
723 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
727 config.validate()?;
728 Ok(Self {
729 config,
730 params: BTreeMap::new(),
731 intercept: FtrlParam::default(),
732 samples_seen: 0,
733 })
734 }
735
736 pub const fn config(&self) -> &FtrlConfig {
738 &self.config
739 }
740
741 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
746 self.params
747 .iter()
748 .map(|(&id, param)| (id, param.weight(&self.config)))
749 .filter(|&(_, w)| w != 0.0)
750 .collect()
751 }
752
753 pub fn intercept(&self) -> f64 {
755 self.intercept.intercept_weight(&self.config)
756 }
757
758 pub fn feature_count(&self) -> usize {
760 self.params.len()
761 }
762
763 fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
765 let dot = compute_dot(&self.params, &self.config, features)?;
766 let intercept = self.intercept.intercept_weight(&self.config);
767 ensure_finite("ftrl_intercept", intercept)?;
768 let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
769 Ok(sigmoid(logit))
770 }
771}
772
773#[cfg(feature = "serde")]
774impl<'de> serde::Deserialize<'de> for FtrlClassifier {
775 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
776 where
777 D: serde::Deserializer<'de>,
778 {
779 #[derive(serde::Deserialize)]
780 struct FtrlClassifierState {
781 config: FtrlConfig,
782 params: BTreeMap<FeatureId, FtrlParam>,
783 intercept: FtrlParam,
784 samples_seen: u64,
785 }
786
787 let state = FtrlClassifierState::deserialize(deserializer)?;
788 let model = FtrlClassifier {
789 config: state.config,
790 params: state.params,
791 intercept: state.intercept,
792 samples_seen: state.samples_seen,
793 };
794 model
795 .validate_invariants()
796 .map_err(serde::de::Error::custom)?;
797 Ok(model)
798 }
799}
800
801impl FtrlClassifier {
802 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
803 fn validate_invariants(&self) -> Result<(), RillError> {
804 self.config.validate()?;
806 for (id, param) in &self.params {
807 param.validate()?;
808 param
809 .weight_checked(&self.config)
810 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
811 }
812 self.intercept.validate()?;
813 self.intercept
814 .intercept_weight_checked(&self.config)
815 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
816 Ok(())
817 }
818}
819
820impl SparseClassifier for FtrlClassifier {
821 fn samples_seen(&self) -> u64 {
822 self.samples_seen
823 }
824
825 fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
826 self.predict_proba_inner(features)
827 }
828
829 fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
830 if features.is_empty() {
831 return Err(RillError::EmptyFeatures);
832 }
833
834 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
835
836 let probability = self.predict_proba_inner(features)?;
837 ensure_finite("ftrl_probability", probability)?;
838 let y = if target { 1.0 } else { 0.0 };
839 let grad = probability - y;
840 ensure_finite("ftrl_gradient", grad)?;
841
842 let new_ids_count = features
843 .values()
844 .iter()
845 .filter(|(id, _)| !self.params.contains_key(id))
846 .count();
847 let mut skip_new_features = false;
848 if let Some(max_features) = self.config.max_features {
849 let projected = self.params.len().saturating_add(new_ids_count);
850 if projected > max_features {
851 match self.config.new_feature_policy {
852 NewFeaturePolicy::Reject => {
853 return Err(RillError::InvalidState(format!(
854 "FTRL feature count {projected} exceeds max_features {max_features}"
855 )));
856 }
857 NewFeaturePolicy::Ignore => {
858 skip_new_features = true;
859 }
860 }
861 }
862 }
863
864 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
865 for &(id, value) in features.values() {
866 let is_new = !self.params.contains_key(&id);
871 if is_new && skip_new_features {
872 continue;
873 }
874
875 let g = grad * value;
876 ensure_finite("ftrl_feature_gradient", g)?;
877
878 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
879 let w = param.weight(&self.config);
880 param.next_updated(g, w, &self.config)?
881 } else {
882 let param = FtrlParam::default();
883 let w = param.weight(&self.config);
884 param.next_updated(g, w, &self.config)?
885 };
886 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
888 ensure_finite("ftrl_next_weight", next_w)?;
889 updates.push((id, new_z, new_n));
890 }
891
892 let w_b = self.intercept.intercept_weight(&self.config);
893 let (new_intercept_z, new_intercept_n) =
894 self.intercept.next_updated(grad, w_b, &self.config)?;
895 let next_intercept_w = FtrlParam {
897 z: new_intercept_z,
898 n: new_intercept_n,
899 }
900 .intercept_weight(&self.config);
901 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
902
903 for (id, new_z, new_n) in updates {
904 let param = self.params.entry(id).or_default();
905 param.z = new_z;
906 param.n = new_n;
907 }
908 self.intercept.z = new_intercept_z;
909 self.intercept.n = new_intercept_n;
910 self.samples_seen = next_samples_seen;
911
912 Ok(())
913 }
914
915 fn reset(&mut self) {
916 self.params.clear();
917 self.intercept = FtrlParam::default();
918 self.samples_seen = 0;
919 }
920}
921
922#[cfg(test)]
923mod tests {
924 use super::*;
925 use rand::SeedableRng;
926
927 #[test]
932 fn cold_start_returns_zero() {
933 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
934 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
935 let pred = model.predict(&sf).unwrap();
936 assert!(pred.abs() < 1e-12);
937 }
938
939 #[test]
940 fn learn_linear_data_converges() {
941 let mut model = FtrlRegressor::new(FtrlConfig {
943 alpha: 0.5,
944 beta: 1.0,
945 l1: 0.0,
946 l2: 0.0,
947 max_features: None,
948 new_feature_policy: NewFeaturePolicy::default(),
949 })
950 .unwrap();
951 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
952 let mut first_err = 0.0;
953 let mut last_err = 0.0;
954 for i in 0..500 {
955 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
956 let y = 2.0 * x;
957 let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
958 let pred = model.predict(&sf).unwrap();
959 let err = (pred - y).abs();
960 if i < 10 {
961 first_err += err;
962 }
963 if i >= 490 {
964 last_err += err;
965 }
966 model.learn(&sf, y).unwrap();
967 }
968 assert!(last_err < first_err, "error should decrease");
969 let weights = model.weights();
970 assert_eq!(weights.len(), 1);
971 assert!(
972 (weights[0].1 - 2.0).abs() < 0.5,
973 "weight should approach 2.0"
974 );
975 }
976
977 #[test]
978 fn l1_produces_sparse_weights() {
979 let mut model = FtrlRegressor::new(FtrlConfig {
981 alpha: 0.1,
982 beta: 1.0,
983 l1: 100.0,
984 l2: 0.0,
985 max_features: None,
986 new_feature_policy: NewFeaturePolicy::default(),
987 })
988 .unwrap();
989 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
990 for _ in 0..200 {
991 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
992 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
993 let y = 0.5 * x1;
994 let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
995 model.learn(&sf, y).unwrap();
996 }
997 let weights = model.weights();
998 assert!(
1000 weights.is_empty(),
1001 "weights should all be zero, got {weights:?}"
1002 );
1003 }
1004
1005 #[test]
1006 fn dynamic_features() {
1007 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1008 assert_eq!(model.feature_count(), 0);
1009 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1010 model.learn(&sf1, 1.0).unwrap();
1011 assert_eq!(model.feature_count(), 1);
1012 let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1014 model.learn(&sf2, 2.0).unwrap();
1015 assert_eq!(model.feature_count(), 2);
1016 assert!(model.params.contains_key(&0));
1018 assert!(model.params.contains_key(&5));
1019 }
1020
1021 #[test]
1022 fn predict_does_not_update_state() {
1023 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1024 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1025 let _ = model.predict(&sf).unwrap();
1026 assert_eq!(model.samples_seen(), 0);
1027 assert_eq!(model.feature_count(), 0);
1028 model.learn(&sf, 1.0).unwrap();
1030 let count_after_learn = model.feature_count();
1031 let _ = model.predict(&sf).unwrap();
1032 assert_eq!(model.feature_count(), count_after_learn);
1033 assert_eq!(model.samples_seen(), 1);
1034 }
1035
1036 #[test]
1037 fn non_finite_value_rejected() {
1038 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1039 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1041 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1042 assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1043 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1044 assert!(model.predict(&sf).is_ok());
1045 }
1046
1047 #[test]
1048 fn non_finite_target_rejected() {
1049 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1050 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1051 assert!(model.learn(&sf, f64::NAN).is_err());
1052 assert!(model.learn(&sf, f64::INFINITY).is_err());
1053 assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1054 assert_eq!(model.samples_seen(), 0);
1056 }
1057
1058 #[test]
1059 fn empty_features_rejected() {
1060 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1061 let sf = SparseFeatures::new();
1062 assert!(model.predict(&sf).is_err());
1063 assert!(model.learn(&sf, 1.0).is_err());
1064 }
1065
1066 #[test]
1067 fn reset_clears_state() {
1068 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1069 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1070 model.learn(&sf, 3.0).unwrap();
1071 model.learn(&sf, 3.0).unwrap();
1072 assert_eq!(model.samples_seen(), 2);
1073 assert_eq!(model.feature_count(), 2);
1074 model.reset();
1075 assert_eq!(model.samples_seen(), 0);
1076 assert_eq!(model.feature_count(), 0);
1077 assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1078 }
1079
1080 #[test]
1081 fn invalid_config_rejected() {
1082 assert!(
1083 FtrlRegressor::new(FtrlConfig {
1084 alpha: 0.0,
1085 ..FtrlConfig::default()
1086 })
1087 .is_err()
1088 );
1089 assert!(
1090 FtrlRegressor::new(FtrlConfig {
1091 alpha: -1.0,
1092 ..FtrlConfig::default()
1093 })
1094 .is_err()
1095 );
1096 assert!(
1097 FtrlRegressor::new(FtrlConfig {
1098 beta: -1.0,
1099 ..FtrlConfig::default()
1100 })
1101 .is_err()
1102 );
1103 assert!(
1104 FtrlRegressor::new(FtrlConfig {
1105 l1: -1.0,
1106 ..FtrlConfig::default()
1107 })
1108 .is_err()
1109 );
1110 assert!(
1111 FtrlRegressor::new(FtrlConfig {
1112 l2: -1.0,
1113 ..FtrlConfig::default()
1114 })
1115 .is_err()
1116 );
1117 assert!(
1118 FtrlRegressor::new(FtrlConfig {
1119 alpha: f64::NAN,
1120 ..FtrlConfig::default()
1121 })
1122 .is_err()
1123 );
1124 assert!(
1125 FtrlRegressor::new(FtrlConfig {
1126 max_features: Some(0),
1127 ..FtrlConfig::default()
1128 })
1129 .is_err()
1130 );
1131 }
1132
1133 #[test]
1134 #[cfg(feature = "serde")]
1135 fn serde_roundtrip() {
1136 let mut model = FtrlRegressor::new(FtrlConfig {
1137 alpha: 0.2,
1138 beta: 0.5,
1139 l1: 0.5,
1140 l2: 0.5,
1141 max_features: Some(100),
1142 new_feature_policy: NewFeaturePolicy::Reject,
1143 })
1144 .unwrap();
1145 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1146 model.learn(&sf, 5.0).unwrap();
1147 let json = serde_json::to_string(&model).unwrap();
1148 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1149 assert_eq!(restored.samples_seen(), model.samples_seen());
1150 assert_eq!(restored.feature_count(), model.feature_count());
1151 let pred_orig = model.predict(&sf).unwrap();
1152 let pred_restored = restored.predict(&sf).unwrap();
1153 assert!((pred_orig - pred_restored).abs() < 1e-12);
1154 }
1155
1156 #[test]
1157 fn weights_returns_nonzero_only() {
1158 let mut model = FtrlRegressor::new(FtrlConfig {
1159 alpha: 0.5,
1160 beta: 1.0,
1161 l1: 0.0,
1162 l2: 0.0,
1163 max_features: None,
1164 new_feature_policy: NewFeaturePolicy::default(),
1165 })
1166 .unwrap();
1167 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1169 for _ in 0..50 {
1170 model.learn(&sf, 1.0).unwrap();
1171 }
1172 let weights = model.weights();
1173 for &(_, w) in &weights {
1175 assert!(w != 0.0);
1176 }
1177 assert!(weights.iter().any(|&(id, _)| id == 0));
1179 }
1180
1181 #[test]
1182 fn multiple_features() {
1183 let mut model = FtrlRegressor::new(FtrlConfig {
1185 alpha: 0.5,
1186 beta: 1.0,
1187 l1: 0.0,
1188 l2: 0.0,
1189 max_features: None,
1190 new_feature_policy: NewFeaturePolicy::default(),
1191 })
1192 .unwrap();
1193 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1194 for _ in 0..500 {
1195 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1196 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1197 let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1198 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1199 model.learn(&sf, y).unwrap();
1200 }
1201 let weights = model.weights();
1202 assert_eq!(weights.len(), 2);
1203 let w0 = weights
1204 .iter()
1205 .find(|&&(id, _)| id == 0)
1206 .map(|&(_, w)| w)
1207 .unwrap();
1208 let w1 = weights
1209 .iter()
1210 .find(|&&(id, _)| id == 1)
1211 .map(|&(_, w)| w)
1212 .unwrap();
1213 assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1214 assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1215 assert!(
1216 (model.intercept() - 0.5).abs() < 0.5,
1217 "intercept should approach 0.5"
1218 );
1219 }
1220
1221 #[test]
1222 fn intercept_learned() {
1223 let mut model = FtrlRegressor::new(FtrlConfig {
1226 alpha: 0.5,
1227 beta: 1.0,
1228 l1: 0.0,
1229 l2: 0.0,
1230 max_features: None,
1231 new_feature_policy: NewFeaturePolicy::default(),
1232 })
1233 .unwrap();
1234 let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1235 for _ in 0..300 {
1236 model.learn(&sf, 3.0).unwrap();
1237 }
1238 let pred = model.predict(&sf).unwrap();
1239 assert!(
1240 (pred - 3.0).abs() < 0.5,
1241 "prediction should approach 3.0, got {pred}"
1242 );
1243 assert!(
1244 (model.intercept() - 3.0).abs() < 0.5,
1245 "intercept should approach 3.0"
1246 );
1247 assert!(model.weights().is_empty());
1249 }
1250
1251 #[test]
1252 fn high_dim_sparse() {
1253 let mut model = FtrlRegressor::new(FtrlConfig {
1256 alpha: 0.3,
1257 beta: 1.0,
1258 l1: 0.0,
1259 l2: 0.0,
1260 max_features: None,
1261 new_feature_policy: NewFeaturePolicy::default(),
1262 })
1263 .unwrap();
1264 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1265 let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1267 let mut first_err = 0.0;
1268 let mut last_err = 0.0;
1269 for i in 0..2000 {
1270 let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1271 for (j, &w) in true_w.iter().enumerate() {
1272 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1273 active.push((j as u64, x * w));
1274 }
1275 for k in 5..10 {
1277 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1278 active.push((k as u64 + 100, x));
1279 }
1280 active.sort_by_key(|(id, _)| *id);
1281 let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1282 let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1283 let pred = model.predict(&sf).unwrap();
1284 let err = (pred - y).abs();
1285 if i < 20 {
1286 first_err += err;
1287 }
1288 if i >= 1980 {
1289 last_err += err;
1290 }
1291 model.learn(&sf, y).unwrap();
1292 }
1293 assert!(
1294 last_err < first_err,
1295 "error should decrease in high-dim sparse"
1296 );
1297 }
1298
1299 #[test]
1304 fn regressor_overflow_does_not_mutate_state() {
1305 let mut model = FtrlRegressor::new(FtrlConfig {
1309 alpha: 0.1,
1310 beta: 1.0,
1311 l1: 0.0,
1312 l2: 0.0,
1313 max_features: None,
1314 new_feature_policy: NewFeaturePolicy::default(),
1315 })
1316 .unwrap();
1317 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1318 let result = model.learn(&sf, 1e100);
1319 assert!(result.is_err(), "expected overflow error, got {result:?}");
1320 assert_eq!(model.samples_seen(), 0);
1321 assert_eq!(model.feature_count(), 0);
1322 assert!(model.params.is_empty());
1323 assert_eq!(model.intercept.z, 0.0);
1324 assert_eq!(model.intercept.n, 0.0);
1325 }
1326
1327 #[test]
1328 fn regressor_partial_update_is_atomic() {
1329 let mut model = FtrlRegressor::new(FtrlConfig {
1332 alpha: 0.1,
1333 beta: 1.0,
1334 l1: 0.0,
1335 l2: 0.0,
1336 max_features: None,
1337 new_feature_policy: NewFeaturePolicy::default(),
1338 })
1339 .unwrap();
1340 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1341 assert!(model.learn(&sf, 1e100).is_err());
1342 assert!(!model.params.contains_key(&0));
1344 assert!(!model.params.contains_key(&1));
1345 assert_eq!(model.samples_seen(), 0);
1346 }
1347
1348 #[test]
1349 #[cfg(feature = "serde")]
1350 fn regressor_samples_seen_overflow_is_atomic() {
1351 let json = format!(
1352 "{{\"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\":{}}}",
1353 u64::MAX
1354 );
1355 let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1356 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1357 let result = model.learn(&sf, 1.0);
1358 assert!(result.is_err(), "expected counter overflow");
1359 assert_eq!(model.samples_seen(), u64::MAX);
1360 assert_eq!(model.feature_count(), 0);
1361 assert_eq!(model.intercept.z, 0.0);
1362 assert_eq!(model.intercept.n, 0.0);
1363 }
1364
1365 #[test]
1370 fn regressor_max_features_reject_at_limit() {
1371 let mut model = FtrlRegressor::new(FtrlConfig {
1372 alpha: 0.5,
1373 beta: 1.0,
1374 l1: 0.0,
1375 l2: 0.0,
1376 max_features: Some(2),
1377 new_feature_policy: NewFeaturePolicy::Reject,
1378 })
1379 .unwrap();
1380 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1382 model.learn(&sf, 1.0).unwrap();
1383 assert_eq!(model.feature_count(), 2);
1384 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1386 assert!(model.learn(&sf_new, 1.0).is_err());
1387 assert_eq!(model.feature_count(), 2);
1388 assert_eq!(model.samples_seen(), 1);
1389 let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1391 model.learn(&sf_existing, 1.0).unwrap();
1392 assert_eq!(model.feature_count(), 2);
1393 assert_eq!(model.samples_seen(), 2);
1394 }
1395
1396 #[test]
1397 fn regressor_max_features_ignore_skips_new() {
1398 let mut model = FtrlRegressor::new(FtrlConfig {
1399 alpha: 0.5,
1400 beta: 1.0,
1401 l1: 0.0,
1402 l2: 0.0,
1403 max_features: Some(2),
1404 new_feature_policy: NewFeaturePolicy::Ignore,
1405 })
1406 .unwrap();
1407 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1408 model.learn(&sf, 1.0).unwrap();
1409 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1413 model.learn(&sf_mixed, 1.0).unwrap();
1414 assert_eq!(model.feature_count(), 2);
1415 assert!(!model.params.contains_key(&2));
1416 assert_eq!(model.samples_seen(), 2);
1417 }
1418
1419 #[test]
1420 fn regressor_max_features_multi_new_prejudge() {
1421 let mut model = FtrlRegressor::new(FtrlConfig {
1422 alpha: 0.5,
1423 beta: 1.0,
1424 l1: 0.0,
1425 l2: 0.0,
1426 max_features: Some(2),
1427 new_feature_policy: NewFeaturePolicy::Reject,
1428 })
1429 .unwrap();
1430 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1433 assert!(model.learn(&sf, 1.0).is_err());
1434 assert_eq!(model.feature_count(), 0);
1435 assert_eq!(model.samples_seen(), 0);
1436 }
1437
1438 #[test]
1443 #[cfg(feature = "serde")]
1444 fn regressor_serde_rejects_negative_n() {
1445 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}";
1446 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1447 assert!(result.is_err(), "negative n must be rejected");
1448 }
1449
1450 #[test]
1451 #[cfg(feature = "serde")]
1452 fn regressor_serde_rejects_invalid_config() {
1453 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}";
1454 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1455 assert!(result.is_err(), "invalid alpha must be rejected");
1456 }
1457
1458 #[test]
1459 #[cfg(feature = "serde")]
1460 fn regressor_serde_accepts_missing_optional_fields() {
1461 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}";
1463 let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1464 assert!(model.config().max_features.is_none());
1465 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1466 }
1467
1468 #[test]
1473 fn cold_start_returns_0_5() {
1474 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1475 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1476 let p = model.predict_proba(&sf).unwrap();
1477 assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1478 }
1479
1480 #[test]
1481 fn learn_separable_data() {
1482 let mut model = FtrlClassifier::new(FtrlConfig {
1484 alpha: 0.5,
1485 beta: 1.0,
1486 l1: 0.0,
1487 l2: 0.0,
1488 max_features: None,
1489 new_feature_policy: NewFeaturePolicy::default(),
1490 })
1491 .unwrap();
1492 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1493 for _ in 0..1000 {
1494 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1495 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1496 let y = x0 > 0.0;
1497 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1498 model.learn(&sf, y).unwrap();
1499 }
1500 let p_pos = model
1501 .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1502 .unwrap();
1503 let p_neg = model
1504 .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1505 .unwrap();
1506 assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1507 assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1508 }
1509
1510 #[test]
1511 fn classifier_l1_produces_sparse_weights() {
1512 let mut model = FtrlClassifier::new(FtrlConfig {
1513 alpha: 0.1,
1514 beta: 1.0,
1515 l1: 100.0,
1516 l2: 0.0,
1517 max_features: None,
1518 new_feature_policy: NewFeaturePolicy::default(),
1519 })
1520 .unwrap();
1521 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1522 for _ in 0..200 {
1523 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1524 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1525 let y = x0 > 0.0;
1526 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1527 model.learn(&sf, y).unwrap();
1528 }
1529 let weights = model.weights();
1530 assert!(
1531 weights.is_empty(),
1532 "weights should all be zero with high L1, got {weights:?}"
1533 );
1534 }
1535
1536 #[test]
1537 fn classifier_dynamic_features() {
1538 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1539 assert_eq!(model.feature_count(), 0);
1540 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1541 model.learn(&sf1, true).unwrap();
1542 assert_eq!(model.feature_count(), 1);
1543 let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1544 model.learn(&sf2, false).unwrap();
1545 assert_eq!(model.feature_count(), 2);
1546 }
1547
1548 #[test]
1549 fn classifier_predict_does_not_update_state() {
1550 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1551 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1552 let _ = model.predict_proba(&sf).unwrap();
1553 assert_eq!(model.samples_seen(), 0);
1554 assert_eq!(model.feature_count(), 0);
1555 model.learn(&sf, true).unwrap();
1556 let count = model.feature_count();
1557 let _ = model.predict_proba(&sf).unwrap();
1558 assert_eq!(model.feature_count(), count);
1559 assert_eq!(model.samples_seen(), 1);
1560 }
1561
1562 #[test]
1563 fn classifier_non_finite_value_rejected() {
1564 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1565 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1566 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1567 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1568 assert!(model.predict_proba(&sf).is_ok());
1569 }
1570
1571 #[test]
1572 fn classifier_empty_features_rejected() {
1573 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1574 let sf = SparseFeatures::new();
1575 assert!(model.predict_proba(&sf).is_err());
1576 assert!(model.learn(&sf, true).is_err());
1577 }
1578
1579 #[test]
1580 fn classifier_reset_clears_state() {
1581 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1582 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1583 model.learn(&sf, true).unwrap();
1584 model.learn(&sf, false).unwrap();
1585 assert_eq!(model.samples_seen(), 2);
1586 assert!(model.feature_count() > 0);
1587 model.reset();
1588 assert_eq!(model.samples_seen(), 0);
1589 assert_eq!(model.feature_count(), 0);
1590 let p = model.predict_proba(&sf).unwrap();
1591 assert!((p - 0.5).abs() < 1e-12);
1592 }
1593
1594 #[test]
1595 fn classifier_invalid_config_rejected() {
1596 assert!(
1597 FtrlClassifier::new(FtrlConfig {
1598 alpha: 0.0,
1599 ..FtrlConfig::default()
1600 })
1601 .is_err()
1602 );
1603 assert!(
1604 FtrlClassifier::new(FtrlConfig {
1605 beta: -0.1,
1606 ..FtrlConfig::default()
1607 })
1608 .is_err()
1609 );
1610 assert!(
1611 FtrlClassifier::new(FtrlConfig {
1612 l1: -1.0,
1613 ..FtrlConfig::default()
1614 })
1615 .is_err()
1616 );
1617 assert!(
1618 FtrlClassifier::new(FtrlConfig {
1619 l2: -1.0,
1620 ..FtrlConfig::default()
1621 })
1622 .is_err()
1623 );
1624 assert!(
1625 FtrlClassifier::new(FtrlConfig {
1626 alpha: f64::INFINITY,
1627 ..FtrlConfig::default()
1628 })
1629 .is_err()
1630 );
1631 }
1632
1633 #[test]
1634 #[cfg(feature = "serde")]
1635 fn classifier_serde_roundtrip() {
1636 let mut model = FtrlClassifier::new(FtrlConfig {
1637 alpha: 0.3,
1638 beta: 0.5,
1639 l1: 0.1,
1640 l2: 0.2,
1641 max_features: Some(100),
1642 new_feature_policy: NewFeaturePolicy::Reject,
1643 })
1644 .unwrap();
1645 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1646 model.learn(&sf, true).unwrap();
1647 model.learn(&sf, false).unwrap();
1648 let json = serde_json::to_string(&model).unwrap();
1649 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1650 assert_eq!(restored.samples_seen(), model.samples_seen());
1651 assert_eq!(restored.feature_count(), model.feature_count());
1652 let p1 = model.predict_proba(&sf).unwrap();
1653 let p2 = restored.predict_proba(&sf).unwrap();
1654 assert!((p1 - p2).abs() < 1e-12);
1655 }
1656
1657 #[test]
1658 fn predict_proba_in_range() {
1659 let mut model = FtrlClassifier::new(FtrlConfig {
1660 alpha: 0.5,
1661 beta: 1.0,
1662 l1: 0.0,
1663 l2: 0.0,
1664 max_features: None,
1665 new_feature_policy: NewFeaturePolicy::default(),
1666 })
1667 .unwrap();
1668 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1669 for _ in 0..200 {
1670 let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1671 let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1672 let y = x0 > 0.0;
1673 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1674 model.learn(&sf, y).unwrap();
1675 let p = model.predict_proba(&sf).unwrap();
1676 assert!(
1677 (0.0..=1.0).contains(&p),
1678 "probability must be in [0,1], got {p}"
1679 );
1680 }
1681 }
1682
1683 #[test]
1684 fn learn_improves_accuracy() {
1685 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1686 let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1688 .map(|_| {
1689 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1690 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1691 let y = x0 + x1 > 0.0;
1692 (
1693 SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1694 y,
1695 )
1696 })
1697 .collect();
1698
1699 let mut model = FtrlClassifier::new(FtrlConfig {
1700 alpha: 0.5,
1701 beta: 1.0,
1702 l1: 0.0,
1703 l2: 0.0,
1704 max_features: None,
1705 new_feature_policy: NewFeaturePolicy::default(),
1706 })
1707 .unwrap();
1708
1709 let acc_before: f64 = test_set
1711 .iter()
1712 .map(|(sf, y)| {
1713 let pred = model.predict(sf).unwrap();
1714 if pred == *y { 1.0 } else { 0.0 }
1715 })
1716 .sum::<f64>()
1717 / test_set.len() as f64;
1718
1719 for _ in 0..1000 {
1721 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1722 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1723 let y = x0 + x1 > 0.0;
1724 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1725 model.learn(&sf, y).unwrap();
1726 }
1727
1728 let acc_after: f64 = test_set
1729 .iter()
1730 .map(|(sf, y)| {
1731 let pred = model.predict(sf).unwrap();
1732 if pred == *y { 1.0 } else { 0.0 }
1733 })
1734 .sum::<f64>()
1735 / test_set.len() as f64;
1736
1737 assert!(
1738 acc_after > acc_before,
1739 "accuracy should improve: {acc_before} -> {acc_after}"
1740 );
1741 }
1742
1743 #[test]
1744 fn classifier_weights_returns_nonzero_only() {
1745 let mut model = FtrlClassifier::new(FtrlConfig {
1746 alpha: 0.5,
1747 beta: 1.0,
1748 l1: 0.0,
1749 l2: 0.0,
1750 max_features: None,
1751 new_feature_policy: NewFeaturePolicy::default(),
1752 })
1753 .unwrap();
1754 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1755 for _ in 0..50 {
1756 model.learn(&sf, true).unwrap();
1757 }
1758 let weights = model.weights();
1759 for &(_, w) in &weights {
1760 assert!(w != 0.0);
1761 }
1762 }
1763
1764 #[test]
1765 fn classifier_multiple_features() {
1766 let mut model = FtrlClassifier::new(FtrlConfig {
1767 alpha: 0.5,
1768 beta: 1.0,
1769 l1: 0.0,
1770 l2: 0.0,
1771 max_features: None,
1772 new_feature_policy: NewFeaturePolicy::default(),
1773 })
1774 .unwrap();
1775 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1776 for _ in 0..1000 {
1777 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1778 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1779 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1780 let y = x0 + x1 > 0.0;
1782 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1783 model.learn(&sf, y).unwrap();
1784 }
1785 let weights = model.weights();
1786 assert!(weights.iter().any(|&(id, _)| id == 0));
1788 assert!(weights.iter().any(|&(id, _)| id == 1));
1789 let p_pos = model
1791 .predict_proba(
1792 &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1793 )
1794 .unwrap();
1795 let p_neg = model
1796 .predict_proba(
1797 &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1798 )
1799 .unwrap();
1800 assert!(p_pos > 0.8);
1801 assert!(p_neg < 0.2);
1802 }
1803
1804 #[test]
1805 fn log_loss_converges() {
1806 let mut model = FtrlClassifier::new(FtrlConfig {
1808 alpha: 0.5,
1809 beta: 1.0,
1810 l1: 0.0,
1811 l2: 0.0,
1812 max_features: None,
1813 new_feature_policy: NewFeaturePolicy::default(),
1814 })
1815 .unwrap();
1816 let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1817 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1818 let mut first_loss = 0.0;
1819 let mut last_loss = 0.0;
1820 for i in 0..1000 {
1821 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1822 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1823 let y = x0 > 0.0;
1824 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1825 let p = model.predict_proba(&sf).unwrap();
1826 let loss = loss_fn.loss(p, y);
1827 if i < 20 {
1828 first_loss += loss;
1829 }
1830 if i >= 980 {
1831 last_loss += loss;
1832 }
1833 model.learn(&sf, y).unwrap();
1834 }
1835 assert!(last_loss < first_loss, "log loss should decrease");
1836 }
1837
1838 #[test]
1843 fn classifier_overflow_does_not_mutate_state() {
1844 let mut model = FtrlClassifier::new(FtrlConfig {
1845 alpha: 0.1,
1846 beta: 1.0,
1847 l1: 0.0,
1848 l2: 0.0,
1849 max_features: None,
1850 new_feature_policy: NewFeaturePolicy::default(),
1851 })
1852 .unwrap();
1853 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1854 let result = model.learn(&sf, false);
1857 assert!(result.is_err(), "expected overflow error, got {result:?}");
1858 assert_eq!(model.samples_seen(), 0);
1859 assert_eq!(model.feature_count(), 0);
1860 assert!(model.params.is_empty());
1861 assert_eq!(model.intercept.z, 0.0);
1862 assert_eq!(model.intercept.n, 0.0);
1863 }
1864
1865 #[test]
1866 fn classifier_partial_update_is_atomic() {
1867 let mut model = FtrlClassifier::new(FtrlConfig {
1868 alpha: 0.1,
1869 beta: 1.0,
1870 l1: 0.0,
1871 l2: 0.0,
1872 max_features: None,
1873 new_feature_policy: NewFeaturePolicy::default(),
1874 })
1875 .unwrap();
1876 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1877 assert!(model.learn(&sf, false).is_err());
1878 assert!(!model.params.contains_key(&0));
1879 assert!(!model.params.contains_key(&1));
1880 assert_eq!(model.samples_seen(), 0);
1881 }
1882
1883 #[test]
1884 #[cfg(feature = "serde")]
1885 fn classifier_samples_seen_overflow_is_atomic() {
1886 let json = format!(
1887 "{{\"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\":{}}}",
1888 u64::MAX
1889 );
1890 let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1891 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1892 let result = model.learn(&sf, true);
1893 assert!(result.is_err(), "expected counter overflow");
1894 assert_eq!(model.samples_seen(), u64::MAX);
1895 assert_eq!(model.feature_count(), 0);
1896 assert_eq!(model.intercept.z, 0.0);
1897 assert_eq!(model.intercept.n, 0.0);
1898 }
1899
1900 #[test]
1905 fn classifier_max_features_reject_at_limit() {
1906 let mut model = FtrlClassifier::new(FtrlConfig {
1907 alpha: 0.5,
1908 beta: 1.0,
1909 l1: 0.0,
1910 l2: 0.0,
1911 max_features: Some(2),
1912 new_feature_policy: NewFeaturePolicy::Reject,
1913 })
1914 .unwrap();
1915 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1916 model.learn(&sf, true).unwrap();
1917 assert_eq!(model.feature_count(), 2);
1918 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1919 assert!(model.learn(&sf_new, true).is_err());
1920 assert_eq!(model.feature_count(), 2);
1921 assert_eq!(model.samples_seen(), 1);
1922 }
1923
1924 #[test]
1925 fn classifier_max_features_ignore_skips_new() {
1926 let mut model = FtrlClassifier::new(FtrlConfig {
1927 alpha: 0.5,
1928 beta: 1.0,
1929 l1: 0.0,
1930 l2: 0.0,
1931 max_features: Some(2),
1932 new_feature_policy: NewFeaturePolicy::Ignore,
1933 })
1934 .unwrap();
1935 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1936 model.learn(&sf, true).unwrap();
1937 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1938 model.learn(&sf_mixed, false).unwrap();
1939 assert_eq!(model.feature_count(), 2);
1940 assert!(!model.params.contains_key(&2));
1941 assert_eq!(model.samples_seen(), 2);
1942 }
1943
1944 #[test]
1945 fn classifier_max_features_multi_new_prejudge() {
1946 let mut model = FtrlClassifier::new(FtrlConfig {
1947 alpha: 0.5,
1948 beta: 1.0,
1949 l1: 0.0,
1950 l2: 0.0,
1951 max_features: Some(2),
1952 new_feature_policy: NewFeaturePolicy::Reject,
1953 })
1954 .unwrap();
1955 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1956 assert!(model.learn(&sf, true).is_err());
1957 assert_eq!(model.feature_count(), 0);
1958 assert_eq!(model.samples_seen(), 0);
1959 }
1960
1961 #[test]
1966 #[cfg(feature = "serde")]
1967 fn classifier_serde_rejects_negative_n() {
1968 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}";
1969 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1970 assert!(result.is_err(), "negative n must be rejected");
1971 }
1972
1973 #[test]
1974 #[cfg(feature = "serde")]
1975 fn classifier_serde_rejects_invalid_config() {
1976 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}";
1977 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1978 assert!(result.is_err(), "invalid alpha must be rejected");
1979 }
1980
1981 #[test]
1982 #[cfg(feature = "serde")]
1983 fn classifier_serde_accepts_missing_optional_fields() {
1984 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}";
1985 let model: FtrlClassifier = serde_json::from_str(json).unwrap();
1986 assert!(model.config().max_features.is_none());
1987 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1988 }
1989
1990 #[test]
1995 fn regressor_gradient_squared_underflow_is_atomic() {
1996 let mut model = FtrlRegressor::new(FtrlConfig {
2002 alpha: 1.0,
2003 beta: 0.0,
2004 l1: 0.0,
2005 l2: 0.0,
2006 max_features: None,
2007 new_feature_policy: NewFeaturePolicy::default(),
2008 })
2009 .unwrap();
2010 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2011 let result = model.learn(&sf, -1e-200);
2012 assert!(result.is_err(), "expected underflow error, got {result:?}");
2013 assert_eq!(model.samples_seen(), 0);
2014 assert_eq!(model.feature_count(), 0);
2015 assert!(model.params.is_empty());
2016 assert_eq!(model.intercept.z, 0.0);
2017 assert_eq!(model.intercept.n, 0.0);
2018 }
2019
2020 #[test]
2021 fn classifier_gradient_squared_underflow_is_atomic() {
2022 let mut model = FtrlClassifier::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, 1e-200)]).unwrap();
2035 let result = model.learn(&sf, false);
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 }
2040
2041 #[test]
2042 fn regressor_boundary_config_predict_after_learn_always_finite() {
2043 let mut model = FtrlRegressor::new(FtrlConfig {
2047 alpha: 1.0,
2048 beta: 0.0,
2049 l1: 0.0,
2050 l2: 0.0,
2051 max_features: None,
2052 new_feature_policy: NewFeaturePolicy::default(),
2053 })
2054 .unwrap();
2055 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2056 for _ in 0..100 {
2057 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2058 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2059 let y = 2.0 * x0 - x1;
2060 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2061 model.learn(&sf, y).unwrap();
2062 let pred = model.predict(&sf);
2063 assert!(
2064 pred.is_ok(),
2065 "predict failed after successful learn: {pred:?}"
2066 );
2067 assert!(
2068 pred.unwrap().is_finite(),
2069 "predict must return finite value after successful learn"
2070 );
2071 }
2072 }
2073
2074 #[test]
2075 fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2076 let mut model = FtrlClassifier::new(FtrlConfig {
2077 alpha: 1.0,
2078 beta: 0.0,
2079 l1: 0.0,
2080 l2: 0.0,
2081 max_features: None,
2082 new_feature_policy: NewFeaturePolicy::default(),
2083 })
2084 .unwrap();
2085 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2086 for _ in 0..100 {
2087 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2088 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2089 let y = x0 > 0.0;
2090 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2091 model.learn(&sf, y).unwrap();
2092 let proba = model.predict_proba(&sf);
2093 assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2094 let p = proba.unwrap();
2095 assert!(p.is_finite(), "probability must be finite, got {p}");
2096 assert!(
2097 (0.0..=1.0).contains(&p),
2098 "probability must be in [0,1], got {p}"
2099 );
2100 }
2101 }
2102
2103 #[test]
2104 #[cfg(feature = "serde")]
2105 fn regressor_serde_rejects_n_zero_z_nonzero() {
2106 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}";
2110 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2111 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2112 }
2113
2114 #[test]
2115 #[cfg(feature = "serde")]
2116 fn classifier_serde_rejects_n_zero_z_nonzero() {
2117 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}";
2118 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2119 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2120 }
2121
2122 #[test]
2123 #[cfg(feature = "serde")]
2124 fn regressor_predict_dot_plus_intercept_overflow() {
2125 let z = -f64::MAX * 0.75;
2132 let json = format!(
2133 "{{\"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}}",
2134 z
2135 );
2136 let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2137 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2138 let result = model.predict(&sf);
2139 assert!(
2140 result.is_err(),
2141 "expected dot+intercept overflow error, got {result:?}"
2142 );
2143 }
2144
2145 #[test]
2146 fn regressor_ignore_skips_overflowing_new_feature() {
2147 let mut model = FtrlRegressor::new(FtrlConfig {
2148 alpha: 0.5,
2149 beta: 1.0,
2150 l1: 0.0,
2151 l2: 0.0,
2152 max_features: Some(1),
2153 new_feature_policy: NewFeaturePolicy::Ignore,
2154 })
2155 .unwrap();
2156 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2157 model.learn(&sf, 1.0).unwrap();
2158 assert_eq!(model.feature_count(), 1);
2159
2160 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2165 let result = model.learn(&sf_mixed, 1.0);
2166 assert!(
2167 result.is_ok(),
2168 "Ignore must skip overflowing new feature, got {result:?}"
2169 );
2170 assert_eq!(model.feature_count(), 1);
2171 assert!(!model.params.contains_key(&1));
2172 assert_eq!(model.samples_seen(), 2);
2173 assert!(model.predict(&sf).is_ok());
2174 }
2175
2176 #[test]
2177 fn classifier_ignore_skips_overflowing_new_feature() {
2178 let mut model = FtrlClassifier::new(FtrlConfig {
2179 alpha: 0.5,
2180 beta: 1.0,
2181 l1: 0.0,
2182 l2: 0.0,
2183 max_features: Some(1),
2184 new_feature_policy: NewFeaturePolicy::Ignore,
2185 })
2186 .unwrap();
2187 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2188 model.learn(&sf, true).unwrap();
2189 assert_eq!(model.feature_count(), 1);
2190
2191 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2192 let result = model.learn(&sf_mixed, true);
2193 assert!(
2194 result.is_ok(),
2195 "Ignore must skip overflowing new feature, got {result:?}"
2196 );
2197 assert_eq!(model.feature_count(), 1);
2198 assert!(!model.params.contains_key(&1));
2199 assert_eq!(model.samples_seen(), 2);
2200 assert!(model.predict_proba(&sf).is_ok());
2201 }
2202
2203 #[test]
2208 #[cfg(feature = "serde")]
2209 fn regressor_serde_rejects_config_dependent_zero_denominator() {
2210 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}";
2215 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2216 assert!(
2217 result.is_err(),
2218 "config-dependent zero denominator must be rejected"
2219 );
2220 }
2221
2222 #[test]
2223 #[cfg(feature = "serde")]
2224 fn classifier_serde_rejects_config_dependent_zero_denominator() {
2225 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}";
2226 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2227 assert!(
2228 result.is_err(),
2229 "config-dependent zero denominator must be rejected"
2230 );
2231 }
2232
2233 #[test]
2234 #[cfg(feature = "serde")]
2235 fn regressor_serde_rejects_intercept_zero_denominator() {
2236 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}";
2239 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2240 assert!(
2241 result.is_err(),
2242 "intercept zero denominator must be rejected"
2243 );
2244 }
2245
2246 #[test]
2247 #[cfg(feature = "serde")]
2248 fn classifier_serde_rejects_intercept_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\":{},\"intercept\":{\"z\":1.0,\"n\":1e-300},\"samples_seen\":0}";
2250 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2251 assert!(
2252 result.is_err(),
2253 "intercept zero denominator must be rejected"
2254 );
2255 }
2256
2257 #[test]
2258 #[cfg(feature = "serde")]
2259 fn regressor_valid_boundary_state_roundtrips() {
2260 let mut model = FtrlRegressor::new(FtrlConfig {
2264 alpha: 1.0,
2265 beta: 0.0,
2266 l1: 0.0,
2267 l2: 0.0,
2268 max_features: None,
2269 new_feature_policy: NewFeaturePolicy::default(),
2270 })
2271 .unwrap();
2272 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2273 model.learn(&sf, 3.0).unwrap();
2274 model.learn(&sf, 5.0).unwrap();
2275 let json = serde_json::to_string(&model).unwrap();
2276 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2277 assert_eq!(restored.samples_seen(), model.samples_seen());
2278 assert_eq!(restored.feature_count(), model.feature_count());
2279 let p1 = model.predict(&sf).unwrap();
2280 let p2 = restored.predict(&sf).unwrap();
2281 assert!((p1 - p2).abs() < 1e-12);
2282 for (_, w) in restored.weights() {
2284 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2285 }
2286 assert!(restored.intercept().is_finite());
2287 }
2288
2289 #[test]
2290 #[cfg(feature = "serde")]
2291 fn classifier_valid_boundary_state_roundtrips() {
2292 let mut model = FtrlClassifier::new(FtrlConfig {
2293 alpha: 1.0,
2294 beta: 0.0,
2295 l1: 0.0,
2296 l2: 0.0,
2297 max_features: None,
2298 new_feature_policy: NewFeaturePolicy::default(),
2299 })
2300 .unwrap();
2301 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2302 model.learn(&sf, true).unwrap();
2303 model.learn(&sf, false).unwrap();
2304 let json = serde_json::to_string(&model).unwrap();
2305 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2306 assert_eq!(restored.samples_seen(), model.samples_seen());
2307 assert_eq!(restored.feature_count(), model.feature_count());
2308 let p1 = model.predict_proba(&sf).unwrap();
2309 let p2 = restored.predict_proba(&sf).unwrap();
2310 assert!((p1 - p2).abs() < 1e-12);
2311 for (_, w) in restored.weights() {
2312 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2313 }
2314 assert!(restored.intercept().is_finite());
2315 }
2316}