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, Copy, PartialEq)]
79#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
80pub struct FtrlResourceDiagnostics {
81 pub current_features: usize,
83 pub configured_max: Option<usize>,
85 pub saturation: Option<f64>,
87 pub new_features_rejected: bool,
89}
90
91impl FtrlResourceDiagnostics {
92 fn from_model(
93 current_features: usize,
94 configured_max: Option<usize>,
95 policy: NewFeaturePolicy,
96 ) -> Self {
97 let saturation = configured_max.map(|max| current_features as f64 / max as f64);
98 let new_features_rejected = configured_max
99 .is_some_and(|max| current_features >= max && policy == NewFeaturePolicy::Reject);
100 Self {
101 current_features,
102 configured_max,
103 saturation,
104 new_features_rejected,
105 }
106 }
107}
108
109#[derive(Debug, Clone)]
115#[cfg_attr(feature = "serde", derive(serde::Serialize))]
116#[non_exhaustive]
117pub struct FtrlConfig {
118 pub alpha: f64,
120 pub beta: f64,
122 pub l1: f64,
124 pub l2: f64,
126 pub max_features: Option<usize>,
133 pub new_feature_policy: NewFeaturePolicy,
136}
137
138impl Default for FtrlConfig {
139 fn default() -> Self {
140 Self {
141 alpha: 0.1,
142 beta: 1.0,
143 l1: 1.0,
144 l2: 1.0,
145 max_features: None,
146 new_feature_policy: NewFeaturePolicy::default(),
147 }
148 }
149}
150
151impl FtrlConfig {
152 pub(crate) fn validate(&self) -> Result<(), RillError> {
154 ensure_finite("alpha", self.alpha)?;
155 ensure_finite("beta", self.beta)?;
156 ensure_finite("l1", self.l1)?;
157 ensure_finite("l2", self.l2)?;
158 if self.alpha <= 0.0 {
159 return Err(RillError::InvalidParameter {
160 name: "alpha",
161 value: self.alpha,
162 });
163 }
164 if self.beta < 0.0 {
165 return Err(RillError::InvalidParameter {
166 name: "beta",
167 value: self.beta,
168 });
169 }
170 if self.l1 < 0.0 {
171 return Err(RillError::InvalidParameter {
172 name: "l1",
173 value: self.l1,
174 });
175 }
176 if self.l2 < 0.0 {
177 return Err(RillError::InvalidParameter {
178 name: "l2",
179 value: self.l2,
180 });
181 }
182 if let Some(max_features) = self.max_features
183 && max_features == 0
184 {
185 return Err(RillError::InvalidParameter {
186 name: "max_features",
187 value: 0.0,
188 });
189 }
190 Ok(())
191 }
192}
193
194#[cfg(feature = "serde")]
195impl<'de> serde::Deserialize<'de> for FtrlConfig {
196 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
197 where
198 D: serde::Deserializer<'de>,
199 {
200 #[derive(serde::Deserialize)]
201 struct FtrlConfigState {
202 alpha: f64,
203 beta: f64,
204 l1: f64,
205 l2: f64,
206 #[serde(default)]
207 max_features: Option<usize>,
208 #[serde(default)]
209 new_feature_policy: NewFeaturePolicy,
210 }
211
212 let state = FtrlConfigState::deserialize(deserializer)?;
213 let config = FtrlConfig {
214 alpha: state.alpha,
215 beta: state.beta,
216 l1: state.l1,
217 l2: state.l2,
218 max_features: state.max_features,
219 new_feature_policy: state.new_feature_policy,
220 };
221 config.validate().map_err(serde::de::Error::custom)?;
222 Ok(config)
223 }
224}
225
226#[derive(Debug, Clone, Default)]
232#[cfg_attr(feature = "serde", derive(serde::Serialize))]
233pub struct FtrlParam {
234 z: f64,
236 n: f64,
238}
239
240impl FtrlParam {
241 fn weight(&self, config: &FtrlConfig) -> f64 {
245 if self.z.abs() <= config.l1 {
246 0.0
247 } else {
248 let sign = self.z.signum();
249 let numerator = -(self.z - sign * config.l1);
250 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
251 numerator / denominator
252 }
253 }
254
255 fn intercept_weight(&self, config: &FtrlConfig) -> f64 {
260 if self.n == 0.0 {
261 0.0
262 } else {
263 let numerator = -self.z;
264 let denominator = config.l2 + (config.beta + self.n.sqrt()) / config.alpha;
265 numerator / denominator
266 }
267 }
268
269 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
288 fn weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
289 ensure_finite("ftrl_z", self.z)?;
290 ensure_finite("ftrl_n", self.n)?;
291 if self.n < 0.0 {
292 return Err(RillError::InvalidState(format!(
293 "ftrl n must be non-negative, got {}",
294 self.n
295 )));
296 }
297 if self.z.abs() <= config.l1 {
299 return Ok(0.0);
300 }
301 let sign = self.z.signum();
302 let numerator = -(self.z - sign * config.l1);
303 ensure_finite("ftrl_weight_numerator", numerator)?;
304 let sqrt_n = self.n.sqrt();
305 ensure_finite("ftrl_weight_sqrt_n", sqrt_n)?;
306 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
307 ensure_finite("ftrl_weight_denominator", denominator)?;
308 if denominator == 0.0 {
309 return Err(RillError::InvalidState(format!(
310 "ftrl weight denominator is zero (z={}, n={}, alpha={}, beta={}, l1={}, l2={})",
311 self.z, self.n, config.alpha, config.beta, config.l1, config.l2
312 )));
313 }
314 let weight = numerator / denominator;
315 ensure_finite("ftrl_weight", weight)?;
316 Ok(weight)
317 }
318
319 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
328 fn intercept_weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
329 ensure_finite("ftrl_z", self.z)?;
330 ensure_finite("ftrl_n", self.n)?;
331 if self.n < 0.0 {
332 return Err(RillError::InvalidState(format!(
333 "ftrl n must be non-negative, got {}",
334 self.n
335 )));
336 }
337 if self.n == 0.0 {
341 if self.z != 0.0 {
342 return Err(RillError::InvalidState(format!(
343 "ftrl intercept has n=0 but z={} (non-zero); cannot produce a finite weight",
344 self.z
345 )));
346 }
347 return Ok(0.0);
348 }
349 let numerator = -self.z;
350 ensure_finite("ftrl_intercept_numerator", numerator)?;
351 let sqrt_n = self.n.sqrt();
352 ensure_finite("ftrl_intercept_sqrt_n", sqrt_n)?;
353 let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
354 ensure_finite("ftrl_intercept_denominator", denominator)?;
355 if denominator == 0.0 {
356 return Err(RillError::InvalidState(format!(
357 "ftrl intercept denominator is zero (z={}, n={}, alpha={}, beta={}, l2={})",
358 self.z, self.n, config.alpha, config.beta, config.l2
359 )));
360 }
361 let weight = numerator / denominator;
362 ensure_finite("ftrl_intercept_weight", weight)?;
363 Ok(weight)
364 }
365
366 fn next_updated(
377 &self,
378 gradient: f64,
379 weight: f64,
380 config: &FtrlConfig,
381 ) -> Result<(f64, f64), RillError> {
382 let gradient_sq = gradient * gradient;
383 ensure_finite("ftrl_gradient_squared", gradient_sq)?;
384 if gradient != 0.0 && gradient_sq == 0.0 {
391 return Err(RillError::NonFiniteValue {
392 field: "ftrl_gradient_squared",
393 value: gradient_sq,
394 });
395 }
396 let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
397 let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
398 ensure_finite("ftrl_sigma", sigma)?;
399 let sigma_w = sigma * weight;
400 ensure_finite("ftrl_sigma_weight", sigma_w)?;
401 let z_delta = gradient - sigma_w;
402 ensure_finite("ftrl_z_delta", z_delta)?;
403 let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
404 Ok((z_new, n_new))
405 }
406
407 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
419 fn validate(&self) -> Result<(), RillError> {
420 ensure_finite("ftrl_z", self.z)?;
421 ensure_finite("ftrl_n", self.n)?;
422 if self.n < 0.0 {
423 return Err(RillError::InvalidState(format!(
424 "ftrl n must be non-negative, got {0}",
425 self.n
426 )));
427 }
428 if self.n == 0.0 && self.z != 0.0 {
429 return Err(RillError::InvalidState(format!(
430 "ftrl param has n=0 but z={0} (non-zero); this state cannot \
431 produce a finite weight",
432 self.z
433 )));
434 }
435 Ok(())
436 }
437}
438
439#[cfg(feature = "serde")]
440impl<'de> serde::Deserialize<'de> for FtrlParam {
441 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
442 where
443 D: serde::Deserializer<'de>,
444 {
445 #[derive(serde::Deserialize)]
446 struct FtrlParamState {
447 z: f64,
448 n: f64,
449 }
450
451 let state = FtrlParamState::deserialize(deserializer)?;
452 let param = FtrlParam {
453 z: state.z,
454 n: state.n,
455 };
456 param.validate().map_err(serde::de::Error::custom)?;
457 Ok(param)
458 }
459}
460
461fn compute_dot(
468 params: &BTreeMap<FeatureId, FtrlParam>,
469 config: &FtrlConfig,
470 features: &SparseFeatures,
471) -> Result<f64, RillError> {
472 if features.is_empty() {
473 return Err(RillError::EmptyFeatures);
474 }
475 let mut dot = 0.0;
476 for &(id, value) in features.values() {
477 ensure_finite("sparse_value", value)?;
478 if let Some(param) = params.get(&id) {
479 let w = param.weight(config);
480 ensure_finite("ftrl_weight", w)?;
481 let contribution = w * value;
482 ensure_finite("ftrl_dot_contribution", contribution)?;
483 dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
484 }
485 }
486 Ok(dot)
487}
488
489#[derive(Debug, Clone)]
508#[cfg_attr(feature = "serde", derive(serde::Serialize))]
509pub struct FtrlRegressor {
510 config: FtrlConfig,
511 params: BTreeMap<FeatureId, FtrlParam>,
512 intercept: FtrlParam,
513 samples_seen: u64,
514}
515
516impl FtrlRegressor {
517 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
521 config.validate()?;
522 Ok(Self {
523 config,
524 params: BTreeMap::new(),
525 intercept: FtrlParam::default(),
526 samples_seen: 0,
527 })
528 }
529
530 pub const fn config(&self) -> &FtrlConfig {
532 &self.config
533 }
534
535 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
540 self.params
541 .iter()
542 .map(|(&id, param)| (id, param.weight(&self.config)))
543 .filter(|&(_, w)| w != 0.0)
544 .collect()
545 }
546
547 pub fn intercept(&self) -> f64 {
549 self.intercept.intercept_weight(&self.config)
550 }
551
552 pub fn feature_count(&self) -> usize {
554 self.params.len()
555 }
556
557 pub fn resource_diagnostics(&self) -> FtrlResourceDiagnostics {
559 FtrlResourceDiagnostics::from_model(
560 self.params.len(),
561 self.config.max_features,
562 self.config.new_feature_policy,
563 )
564 }
565
566 fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
568 let dot = compute_dot(&self.params, &self.config, features)?;
569 let intercept = self.intercept.intercept_weight(&self.config);
570 ensure_finite("ftrl_intercept", intercept)?;
571 checked_finite_add(dot, intercept, "ftrl_prediction")
572 }
573}
574
575#[cfg(feature = "serde")]
576impl<'de> serde::Deserialize<'de> for FtrlRegressor {
577 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
578 where
579 D: serde::Deserializer<'de>,
580 {
581 #[derive(serde::Deserialize)]
582 struct FtrlRegressorState {
583 config: FtrlConfig,
584 params: BTreeMap<FeatureId, FtrlParam>,
585 intercept: FtrlParam,
586 samples_seen: u64,
587 }
588
589 let state = FtrlRegressorState::deserialize(deserializer)?;
590 let model = FtrlRegressor {
591 config: state.config,
592 params: state.params,
593 intercept: state.intercept,
594 samples_seen: state.samples_seen,
595 };
596 model
599 .validate_invariants()
600 .map_err(serde::de::Error::custom)?;
601 Ok(model)
602 }
603}
604
605impl FtrlRegressor {
606 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
607 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
608 self.config.validate()?;
615 if let Some(max_features) = self.config.max_features
620 && self.params.len() > max_features
621 {
622 return Err(RillError::InvalidState(format!(
623 "FTRL stored feature count {} exceeds max_features {}",
624 self.params.len(),
625 max_features
626 )));
627 }
628 for (id, param) in &self.params {
629 param.validate()?;
630 param
631 .weight_checked(&self.config)
632 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
633 }
634 self.intercept.validate()?;
635 self.intercept
636 .intercept_weight_checked(&self.config)
637 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
638 Ok(())
639 }
640}
641
642impl SparseRegressor for FtrlRegressor {
643 fn samples_seen(&self) -> u64 {
644 self.samples_seen
645 }
646
647 fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
648 self.predict_inner(features)
649 }
650
651 fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
652 if features.is_empty() {
653 return Err(RillError::EmptyFeatures);
654 }
655 ensure_finite("target", target)?;
656
657 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
659
660 let prediction = self.predict_inner(features)?;
661 ensure_finite("ftrl_prediction", prediction)?;
662 let grad = prediction - target;
663 ensure_finite("ftrl_gradient", grad)?;
664
665 let new_ids_count = features
669 .values()
670 .iter()
671 .filter(|(id, _)| !self.params.contains_key(id))
672 .count();
673 let mut skip_new_features = false;
674 if let Some(max_features) = self.config.max_features {
675 let projected = self.params.len().saturating_add(new_ids_count);
676 if projected > max_features {
677 match self.config.new_feature_policy {
678 NewFeaturePolicy::Reject => {
679 return Err(RillError::InvalidState(format!(
680 "FTRL feature count {projected} exceeds max_features {max_features}"
681 )));
682 }
683 NewFeaturePolicy::Ignore => {
684 skip_new_features = true;
685 }
686 }
687 }
688 }
689
690 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
692 for &(id, value) in features.values() {
693 let is_new = !self.params.contains_key(&id);
698 if is_new && skip_new_features {
699 continue;
700 }
701
702 let g = grad * value;
703 ensure_finite("ftrl_feature_gradient", g)?;
704
705 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
706 let w = param.weight(&self.config);
707 param.next_updated(g, w, &self.config)?
708 } else {
709 let param = FtrlParam::default();
710 let w = param.weight(&self.config);
711 param.next_updated(g, w, &self.config)?
712 };
713 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
718 ensure_finite("ftrl_next_weight", next_w)?;
719 updates.push((id, new_z, new_n));
720 }
721
722 let w_b = self.intercept.intercept_weight(&self.config);
724 let (new_intercept_z, new_intercept_n) =
725 self.intercept.next_updated(grad, w_b, &self.config)?;
726 let next_intercept_w = FtrlParam {
728 z: new_intercept_z,
729 n: new_intercept_n,
730 }
731 .intercept_weight(&self.config);
732 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
733
734 for (id, new_z, new_n) in updates {
736 let param = self.params.entry(id).or_default();
737 param.z = new_z;
738 param.n = new_n;
739 }
740 self.intercept.z = new_intercept_z;
741 self.intercept.n = new_intercept_n;
742 self.samples_seen = next_samples_seen;
743
744 Ok(())
745 }
746
747 fn reset(&mut self) {
748 self.params.clear();
749 self.intercept = FtrlParam::default();
750 self.samples_seen = 0;
751 }
752}
753
754#[derive(Debug, Clone)]
777#[cfg_attr(feature = "serde", derive(serde::Serialize))]
778pub struct FtrlClassifier {
779 config: FtrlConfig,
780 params: BTreeMap<FeatureId, FtrlParam>,
781 intercept: FtrlParam,
782 samples_seen: u64,
783}
784
785impl FtrlClassifier {
786 pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
790 config.validate()?;
791 Ok(Self {
792 config,
793 params: BTreeMap::new(),
794 intercept: FtrlParam::default(),
795 samples_seen: 0,
796 })
797 }
798
799 pub const fn config(&self) -> &FtrlConfig {
801 &self.config
802 }
803
804 pub fn weights(&self) -> Vec<(FeatureId, f64)> {
809 self.params
810 .iter()
811 .map(|(&id, param)| (id, param.weight(&self.config)))
812 .filter(|&(_, w)| w != 0.0)
813 .collect()
814 }
815
816 pub fn intercept(&self) -> f64 {
818 self.intercept.intercept_weight(&self.config)
819 }
820
821 pub fn feature_count(&self) -> usize {
823 self.params.len()
824 }
825
826 pub fn resource_diagnostics(&self) -> FtrlResourceDiagnostics {
828 FtrlResourceDiagnostics::from_model(
829 self.params.len(),
830 self.config.max_features,
831 self.config.new_feature_policy,
832 )
833 }
834
835 fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
837 let dot = compute_dot(&self.params, &self.config, features)?;
838 let intercept = self.intercept.intercept_weight(&self.config);
839 ensure_finite("ftrl_intercept", intercept)?;
840 let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
841 Ok(sigmoid(logit))
842 }
843}
844
845#[cfg(feature = "serde")]
846impl<'de> serde::Deserialize<'de> for FtrlClassifier {
847 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
848 where
849 D: serde::Deserializer<'de>,
850 {
851 #[derive(serde::Deserialize)]
852 struct FtrlClassifierState {
853 config: FtrlConfig,
854 params: BTreeMap<FeatureId, FtrlParam>,
855 intercept: FtrlParam,
856 samples_seen: u64,
857 }
858
859 let state = FtrlClassifierState::deserialize(deserializer)?;
860 let model = FtrlClassifier {
861 config: state.config,
862 params: state.params,
863 intercept: state.intercept,
864 samples_seen: state.samples_seen,
865 };
866 model
867 .validate_invariants()
868 .map_err(serde::de::Error::custom)?;
869 Ok(model)
870 }
871}
872
873impl FtrlClassifier {
874 #[cfg_attr(not(feature = "serde"), allow(dead_code))]
875 pub(crate) fn validate_invariants(&self) -> Result<(), RillError> {
876 self.config.validate()?;
878 if let Some(max_features) = self.config.max_features
881 && self.params.len() > max_features
882 {
883 return Err(RillError::InvalidState(format!(
884 "FTRL stored feature count {} exceeds max_features {}",
885 self.params.len(),
886 max_features
887 )));
888 }
889 for (id, param) in &self.params {
890 param.validate()?;
891 param
892 .weight_checked(&self.config)
893 .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
894 }
895 self.intercept.validate()?;
896 self.intercept
897 .intercept_weight_checked(&self.config)
898 .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
899 Ok(())
900 }
901}
902
903#[cfg(feature = "serde")]
904impl ValidateState for FtrlConfig {
905 fn validate_state(&self) -> Result<(), RillError> {
906 FtrlConfig::validate(self)
907 }
908}
909
910#[cfg(feature = "serde")]
911impl ValidateState for FtrlRegressor {
912 fn validate_state(&self) -> Result<(), RillError> {
913 FtrlRegressor::validate_invariants(self)
914 }
915}
916
917#[cfg(feature = "serde")]
918impl ValidateState for FtrlClassifier {
919 fn validate_state(&self) -> Result<(), RillError> {
920 FtrlClassifier::validate_invariants(self)
921 }
922}
923
924impl SparseClassifier for FtrlClassifier {
925 fn samples_seen(&self) -> u64 {
926 self.samples_seen
927 }
928
929 fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
930 self.predict_proba_inner(features)
931 }
932
933 fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
934 if features.is_empty() {
935 return Err(RillError::EmptyFeatures);
936 }
937
938 let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
939
940 let probability = self.predict_proba_inner(features)?;
941 ensure_finite("ftrl_probability", probability)?;
942 let y = if target { 1.0 } else { 0.0 };
943 let grad = probability - y;
944 ensure_finite("ftrl_gradient", grad)?;
945
946 let new_ids_count = features
947 .values()
948 .iter()
949 .filter(|(id, _)| !self.params.contains_key(id))
950 .count();
951 let mut skip_new_features = false;
952 if let Some(max_features) = self.config.max_features {
953 let projected = self.params.len().saturating_add(new_ids_count);
954 if projected > max_features {
955 match self.config.new_feature_policy {
956 NewFeaturePolicy::Reject => {
957 return Err(RillError::InvalidState(format!(
958 "FTRL feature count {projected} exceeds max_features {max_features}"
959 )));
960 }
961 NewFeaturePolicy::Ignore => {
962 skip_new_features = true;
963 }
964 }
965 }
966 }
967
968 let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
969 for &(id, value) in features.values() {
970 let is_new = !self.params.contains_key(&id);
975 if is_new && skip_new_features {
976 continue;
977 }
978
979 let g = grad * value;
980 ensure_finite("ftrl_feature_gradient", g)?;
981
982 let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
983 let w = param.weight(&self.config);
984 param.next_updated(g, w, &self.config)?
985 } else {
986 let param = FtrlParam::default();
987 let w = param.weight(&self.config);
988 param.next_updated(g, w, &self.config)?
989 };
990 let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
992 ensure_finite("ftrl_next_weight", next_w)?;
993 updates.push((id, new_z, new_n));
994 }
995
996 let w_b = self.intercept.intercept_weight(&self.config);
997 let (new_intercept_z, new_intercept_n) =
998 self.intercept.next_updated(grad, w_b, &self.config)?;
999 let next_intercept_w = FtrlParam {
1001 z: new_intercept_z,
1002 n: new_intercept_n,
1003 }
1004 .intercept_weight(&self.config);
1005 ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
1006
1007 for (id, new_z, new_n) in updates {
1008 let param = self.params.entry(id).or_default();
1009 param.z = new_z;
1010 param.n = new_n;
1011 }
1012 self.intercept.z = new_intercept_z;
1013 self.intercept.n = new_intercept_n;
1014 self.samples_seen = next_samples_seen;
1015
1016 Ok(())
1017 }
1018
1019 fn reset(&mut self) {
1020 self.params.clear();
1021 self.intercept = FtrlParam::default();
1022 self.samples_seen = 0;
1023 }
1024}
1025
1026#[cfg(test)]
1027mod tests {
1028 use super::*;
1029 use rand::SeedableRng;
1030
1031 #[test]
1036 fn cold_start_returns_zero() {
1037 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1038 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1039 let pred = model.predict(&sf).unwrap();
1040 assert!(pred.abs() < 1e-12);
1041 }
1042
1043 #[test]
1044 fn learn_linear_data_converges() {
1045 let mut model = FtrlRegressor::new(FtrlConfig {
1047 alpha: 0.5,
1048 beta: 1.0,
1049 l1: 0.0,
1050 l2: 0.0,
1051 max_features: None,
1052 new_feature_policy: NewFeaturePolicy::default(),
1053 })
1054 .unwrap();
1055 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
1056 let mut first_err = 0.0;
1057 let mut last_err = 0.0;
1058 for i in 0..500 {
1059 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1060 let y = 2.0 * x;
1061 let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
1062 let pred = model.predict(&sf).unwrap();
1063 let err = (pred - y).abs();
1064 if i < 10 {
1065 first_err += err;
1066 }
1067 if i >= 490 {
1068 last_err += err;
1069 }
1070 model.learn(&sf, y).unwrap();
1071 }
1072 assert!(last_err < first_err, "error should decrease");
1073 let weights = model.weights();
1074 assert_eq!(weights.len(), 1);
1075 assert!(
1076 (weights[0].1 - 2.0).abs() < 0.5,
1077 "weight should approach 2.0"
1078 );
1079 }
1080
1081 #[test]
1082 fn l1_produces_sparse_weights() {
1083 let mut model = FtrlRegressor::new(FtrlConfig {
1085 alpha: 0.1,
1086 beta: 1.0,
1087 l1: 100.0,
1088 l2: 0.0,
1089 max_features: None,
1090 new_feature_policy: NewFeaturePolicy::default(),
1091 })
1092 .unwrap();
1093 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
1094 for _ in 0..200 {
1095 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1096 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1097 let y = 0.5 * x1;
1098 let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
1099 model.learn(&sf, y).unwrap();
1100 }
1101 let weights = model.weights();
1102 assert!(
1104 weights.is_empty(),
1105 "weights should all be zero, got {weights:?}"
1106 );
1107 }
1108
1109 #[test]
1110 fn dynamic_features() {
1111 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1112 assert_eq!(model.feature_count(), 0);
1113 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1114 model.learn(&sf1, 1.0).unwrap();
1115 assert_eq!(model.feature_count(), 1);
1116 let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1118 model.learn(&sf2, 2.0).unwrap();
1119 assert_eq!(model.feature_count(), 2);
1120 assert!(model.params.contains_key(&0));
1122 assert!(model.params.contains_key(&5));
1123 }
1124
1125 #[test]
1126 fn predict_does_not_update_state() {
1127 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1128 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1129 let _ = model.predict(&sf).unwrap();
1130 assert_eq!(model.samples_seen(), 0);
1131 assert_eq!(model.feature_count(), 0);
1132 model.learn(&sf, 1.0).unwrap();
1134 let count_after_learn = model.feature_count();
1135 let _ = model.predict(&sf).unwrap();
1136 assert_eq!(model.feature_count(), count_after_learn);
1137 assert_eq!(model.samples_seen(), 1);
1138 }
1139
1140 #[test]
1141 fn non_finite_value_rejected() {
1142 let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1143 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1145 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1146 assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1147 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1148 assert!(model.predict(&sf).is_ok());
1149 }
1150
1151 #[test]
1152 fn non_finite_target_rejected() {
1153 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1154 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1155 assert!(model.learn(&sf, f64::NAN).is_err());
1156 assert!(model.learn(&sf, f64::INFINITY).is_err());
1157 assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1158 assert_eq!(model.samples_seen(), 0);
1160 }
1161
1162 #[test]
1163 fn empty_features_rejected() {
1164 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1165 let sf = SparseFeatures::new();
1166 assert!(model.predict(&sf).is_err());
1167 assert!(model.learn(&sf, 1.0).is_err());
1168 }
1169
1170 #[test]
1171 fn reset_clears_state() {
1172 let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1173 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1174 model.learn(&sf, 3.0).unwrap();
1175 model.learn(&sf, 3.0).unwrap();
1176 assert_eq!(model.samples_seen(), 2);
1177 assert_eq!(model.feature_count(), 2);
1178 model.reset();
1179 assert_eq!(model.samples_seen(), 0);
1180 assert_eq!(model.feature_count(), 0);
1181 assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1182 }
1183
1184 #[test]
1185 fn invalid_config_rejected() {
1186 assert!(
1187 FtrlRegressor::new(FtrlConfig {
1188 alpha: 0.0,
1189 ..FtrlConfig::default()
1190 })
1191 .is_err()
1192 );
1193 assert!(
1194 FtrlRegressor::new(FtrlConfig {
1195 alpha: -1.0,
1196 ..FtrlConfig::default()
1197 })
1198 .is_err()
1199 );
1200 assert!(
1201 FtrlRegressor::new(FtrlConfig {
1202 beta: -1.0,
1203 ..FtrlConfig::default()
1204 })
1205 .is_err()
1206 );
1207 assert!(
1208 FtrlRegressor::new(FtrlConfig {
1209 l1: -1.0,
1210 ..FtrlConfig::default()
1211 })
1212 .is_err()
1213 );
1214 assert!(
1215 FtrlRegressor::new(FtrlConfig {
1216 l2: -1.0,
1217 ..FtrlConfig::default()
1218 })
1219 .is_err()
1220 );
1221 assert!(
1222 FtrlRegressor::new(FtrlConfig {
1223 alpha: f64::NAN,
1224 ..FtrlConfig::default()
1225 })
1226 .is_err()
1227 );
1228 assert!(
1229 FtrlRegressor::new(FtrlConfig {
1230 max_features: Some(0),
1231 ..FtrlConfig::default()
1232 })
1233 .is_err()
1234 );
1235 }
1236
1237 #[test]
1238 #[cfg(feature = "serde")]
1239 fn serde_roundtrip() {
1240 let mut model = FtrlRegressor::new(FtrlConfig {
1241 alpha: 0.2,
1242 beta: 0.5,
1243 l1: 0.5,
1244 l2: 0.5,
1245 max_features: Some(100),
1246 new_feature_policy: NewFeaturePolicy::Reject,
1247 })
1248 .unwrap();
1249 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1250 model.learn(&sf, 5.0).unwrap();
1251 let json = serde_json::to_string(&model).unwrap();
1252 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1253 assert_eq!(restored.samples_seen(), model.samples_seen());
1254 assert_eq!(restored.feature_count(), model.feature_count());
1255 let pred_orig = model.predict(&sf).unwrap();
1256 let pred_restored = restored.predict(&sf).unwrap();
1257 assert!((pred_orig - pred_restored).abs() < 1e-12);
1258 }
1259
1260 #[test]
1261 fn weights_returns_nonzero_only() {
1262 let mut model = FtrlRegressor::new(FtrlConfig {
1263 alpha: 0.5,
1264 beta: 1.0,
1265 l1: 0.0,
1266 l2: 0.0,
1267 max_features: None,
1268 new_feature_policy: NewFeaturePolicy::default(),
1269 })
1270 .unwrap();
1271 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1273 for _ in 0..50 {
1274 model.learn(&sf, 1.0).unwrap();
1275 }
1276 let weights = model.weights();
1277 for &(_, w) in &weights {
1279 assert!(w != 0.0);
1280 }
1281 assert!(weights.iter().any(|&(id, _)| id == 0));
1283 }
1284
1285 #[test]
1286 fn multiple_features() {
1287 let mut model = FtrlRegressor::new(FtrlConfig {
1289 alpha: 0.5,
1290 beta: 1.0,
1291 l1: 0.0,
1292 l2: 0.0,
1293 max_features: None,
1294 new_feature_policy: NewFeaturePolicy::default(),
1295 })
1296 .unwrap();
1297 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1298 for _ in 0..500 {
1299 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1300 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1301 let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1302 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1303 model.learn(&sf, y).unwrap();
1304 }
1305 let weights = model.weights();
1306 assert_eq!(weights.len(), 2);
1307 let w0 = weights
1308 .iter()
1309 .find(|&&(id, _)| id == 0)
1310 .map(|&(_, w)| w)
1311 .unwrap();
1312 let w1 = weights
1313 .iter()
1314 .find(|&&(id, _)| id == 1)
1315 .map(|&(_, w)| w)
1316 .unwrap();
1317 assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1318 assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1319 assert!(
1320 (model.intercept() - 0.5).abs() < 0.5,
1321 "intercept should approach 0.5"
1322 );
1323 }
1324
1325 #[test]
1326 fn intercept_learned() {
1327 let mut model = FtrlRegressor::new(FtrlConfig {
1330 alpha: 0.5,
1331 beta: 1.0,
1332 l1: 0.0,
1333 l2: 0.0,
1334 max_features: None,
1335 new_feature_policy: NewFeaturePolicy::default(),
1336 })
1337 .unwrap();
1338 let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1339 for _ in 0..300 {
1340 model.learn(&sf, 3.0).unwrap();
1341 }
1342 let pred = model.predict(&sf).unwrap();
1343 assert!(
1344 (pred - 3.0).abs() < 0.5,
1345 "prediction should approach 3.0, got {pred}"
1346 );
1347 assert!(
1348 (model.intercept() - 3.0).abs() < 0.5,
1349 "intercept should approach 3.0"
1350 );
1351 assert!(model.weights().is_empty());
1353 }
1354
1355 #[test]
1356 fn high_dim_sparse() {
1357 let mut model = FtrlRegressor::new(FtrlConfig {
1360 alpha: 0.3,
1361 beta: 1.0,
1362 l1: 0.0,
1363 l2: 0.0,
1364 max_features: None,
1365 new_feature_policy: NewFeaturePolicy::default(),
1366 })
1367 .unwrap();
1368 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1369 let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1371 let mut first_err = 0.0;
1372 let mut last_err = 0.0;
1373 for i in 0..2000 {
1374 let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1375 for (j, &w) in true_w.iter().enumerate() {
1376 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1377 active.push((j as u64, x * w));
1378 }
1379 for k in 5..10 {
1381 let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1382 active.push((k as u64 + 100, x));
1383 }
1384 active.sort_by_key(|(id, _)| *id);
1385 let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1386 let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1387 let pred = model.predict(&sf).unwrap();
1388 let err = (pred - y).abs();
1389 if i < 20 {
1390 first_err += err;
1391 }
1392 if i >= 1980 {
1393 last_err += err;
1394 }
1395 model.learn(&sf, y).unwrap();
1396 }
1397 assert!(
1398 last_err < first_err,
1399 "error should decrease in high-dim sparse"
1400 );
1401 }
1402
1403 #[test]
1408 fn regressor_overflow_does_not_mutate_state() {
1409 let mut model = FtrlRegressor::new(FtrlConfig {
1413 alpha: 0.1,
1414 beta: 1.0,
1415 l1: 0.0,
1416 l2: 0.0,
1417 max_features: None,
1418 new_feature_policy: NewFeaturePolicy::default(),
1419 })
1420 .unwrap();
1421 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1422 let result = model.learn(&sf, 1e100);
1423 assert!(result.is_err(), "expected overflow error, got {result:?}");
1424 assert_eq!(model.samples_seen(), 0);
1425 assert_eq!(model.feature_count(), 0);
1426 assert!(model.params.is_empty());
1427 assert_eq!(model.intercept.z, 0.0);
1428 assert_eq!(model.intercept.n, 0.0);
1429 }
1430
1431 #[test]
1432 fn regressor_partial_update_is_atomic() {
1433 let mut model = FtrlRegressor::new(FtrlConfig {
1436 alpha: 0.1,
1437 beta: 1.0,
1438 l1: 0.0,
1439 l2: 0.0,
1440 max_features: None,
1441 new_feature_policy: NewFeaturePolicy::default(),
1442 })
1443 .unwrap();
1444 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1445 assert!(model.learn(&sf, 1e100).is_err());
1446 assert!(!model.params.contains_key(&0));
1448 assert!(!model.params.contains_key(&1));
1449 assert_eq!(model.samples_seen(), 0);
1450 }
1451
1452 #[test]
1453 #[cfg(feature = "serde")]
1454 fn regressor_samples_seen_overflow_is_atomic() {
1455 let json = format!(
1456 "{{\"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\":{}}}",
1457 u64::MAX
1458 );
1459 let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1460 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1461 let result = model.learn(&sf, 1.0);
1462 assert!(result.is_err(), "expected counter overflow");
1463 assert_eq!(model.samples_seen(), u64::MAX);
1464 assert_eq!(model.feature_count(), 0);
1465 assert_eq!(model.intercept.z, 0.0);
1466 assert_eq!(model.intercept.n, 0.0);
1467 }
1468
1469 #[test]
1474 fn regressor_max_features_reject_at_limit() {
1475 let mut model = FtrlRegressor::new(FtrlConfig {
1476 alpha: 0.5,
1477 beta: 1.0,
1478 l1: 0.0,
1479 l2: 0.0,
1480 max_features: Some(2),
1481 new_feature_policy: NewFeaturePolicy::Reject,
1482 })
1483 .unwrap();
1484 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1486 model.learn(&sf, 1.0).unwrap();
1487 assert_eq!(model.feature_count(), 2);
1488 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1490 assert!(model.learn(&sf_new, 1.0).is_err());
1491 assert_eq!(model.feature_count(), 2);
1492 assert_eq!(model.samples_seen(), 1);
1493 let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1495 model.learn(&sf_existing, 1.0).unwrap();
1496 assert_eq!(model.feature_count(), 2);
1497 assert_eq!(model.samples_seen(), 2);
1498 }
1499
1500 #[test]
1501 fn regressor_max_features_ignore_skips_new() {
1502 let mut model = FtrlRegressor::new(FtrlConfig {
1503 alpha: 0.5,
1504 beta: 1.0,
1505 l1: 0.0,
1506 l2: 0.0,
1507 max_features: Some(2),
1508 new_feature_policy: NewFeaturePolicy::Ignore,
1509 })
1510 .unwrap();
1511 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1512 model.learn(&sf, 1.0).unwrap();
1513 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1517 model.learn(&sf_mixed, 1.0).unwrap();
1518 assert_eq!(model.feature_count(), 2);
1519 assert!(!model.params.contains_key(&2));
1520 assert_eq!(model.samples_seen(), 2);
1521 }
1522
1523 #[test]
1524 fn regressor_max_features_multi_new_prejudge() {
1525 let mut model = FtrlRegressor::new(FtrlConfig {
1526 alpha: 0.5,
1527 beta: 1.0,
1528 l1: 0.0,
1529 l2: 0.0,
1530 max_features: Some(2),
1531 new_feature_policy: NewFeaturePolicy::Reject,
1532 })
1533 .unwrap();
1534 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1537 assert!(model.learn(&sf, 1.0).is_err());
1538 assert_eq!(model.feature_count(), 0);
1539 assert_eq!(model.samples_seen(), 0);
1540 }
1541
1542 #[test]
1547 #[cfg(feature = "serde")]
1548 fn regressor_serde_rejects_negative_n() {
1549 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}";
1550 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1551 assert!(result.is_err(), "negative n must be rejected");
1552 }
1553
1554 #[test]
1555 #[cfg(feature = "serde")]
1556 fn regressor_serde_rejects_invalid_config() {
1557 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}";
1558 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1559 assert!(result.is_err(), "invalid alpha must be rejected");
1560 }
1561
1562 #[test]
1563 #[cfg(feature = "serde")]
1564 fn regressor_serde_accepts_missing_optional_fields() {
1565 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}";
1567 let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1568 assert!(model.config().max_features.is_none());
1569 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1570 }
1571
1572 #[test]
1577 fn cold_start_returns_0_5() {
1578 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1579 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1580 let p = model.predict_proba(&sf).unwrap();
1581 assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1582 }
1583
1584 #[test]
1585 fn learn_separable_data() {
1586 let mut model = FtrlClassifier::new(FtrlConfig {
1588 alpha: 0.5,
1589 beta: 1.0,
1590 l1: 0.0,
1591 l2: 0.0,
1592 max_features: None,
1593 new_feature_policy: NewFeaturePolicy::default(),
1594 })
1595 .unwrap();
1596 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1597 for _ in 0..1000 {
1598 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1599 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1600 let y = x0 > 0.0;
1601 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1602 model.learn(&sf, y).unwrap();
1603 }
1604 let p_pos = model
1605 .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1606 .unwrap();
1607 let p_neg = model
1608 .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1609 .unwrap();
1610 assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1611 assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1612 }
1613
1614 #[test]
1615 fn classifier_l1_produces_sparse_weights() {
1616 let mut model = FtrlClassifier::new(FtrlConfig {
1617 alpha: 0.1,
1618 beta: 1.0,
1619 l1: 100.0,
1620 l2: 0.0,
1621 max_features: None,
1622 new_feature_policy: NewFeaturePolicy::default(),
1623 })
1624 .unwrap();
1625 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1626 for _ in 0..200 {
1627 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1628 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1629 let y = x0 > 0.0;
1630 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1631 model.learn(&sf, y).unwrap();
1632 }
1633 let weights = model.weights();
1634 assert!(
1635 weights.is_empty(),
1636 "weights should all be zero with high L1, got {weights:?}"
1637 );
1638 }
1639
1640 #[test]
1641 fn classifier_dynamic_features() {
1642 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1643 assert_eq!(model.feature_count(), 0);
1644 let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1645 model.learn(&sf1, true).unwrap();
1646 assert_eq!(model.feature_count(), 1);
1647 let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1648 model.learn(&sf2, false).unwrap();
1649 assert_eq!(model.feature_count(), 2);
1650 }
1651
1652 #[test]
1653 fn classifier_predict_does_not_update_state() {
1654 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1655 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1656 let _ = model.predict_proba(&sf).unwrap();
1657 assert_eq!(model.samples_seen(), 0);
1658 assert_eq!(model.feature_count(), 0);
1659 model.learn(&sf, true).unwrap();
1660 let count = model.feature_count();
1661 let _ = model.predict_proba(&sf).unwrap();
1662 assert_eq!(model.feature_count(), count);
1663 assert_eq!(model.samples_seen(), 1);
1664 }
1665
1666 #[test]
1667 fn classifier_non_finite_value_rejected() {
1668 let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1669 assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1670 assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1671 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1672 assert!(model.predict_proba(&sf).is_ok());
1673 }
1674
1675 #[test]
1676 fn classifier_empty_features_rejected() {
1677 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1678 let sf = SparseFeatures::new();
1679 assert!(model.predict_proba(&sf).is_err());
1680 assert!(model.learn(&sf, true).is_err());
1681 }
1682
1683 #[test]
1684 fn classifier_reset_clears_state() {
1685 let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1686 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1687 model.learn(&sf, true).unwrap();
1688 model.learn(&sf, false).unwrap();
1689 assert_eq!(model.samples_seen(), 2);
1690 assert!(model.feature_count() > 0);
1691 model.reset();
1692 assert_eq!(model.samples_seen(), 0);
1693 assert_eq!(model.feature_count(), 0);
1694 let p = model.predict_proba(&sf).unwrap();
1695 assert!((p - 0.5).abs() < 1e-12);
1696 }
1697
1698 #[test]
1699 fn classifier_invalid_config_rejected() {
1700 assert!(
1701 FtrlClassifier::new(FtrlConfig {
1702 alpha: 0.0,
1703 ..FtrlConfig::default()
1704 })
1705 .is_err()
1706 );
1707 assert!(
1708 FtrlClassifier::new(FtrlConfig {
1709 beta: -0.1,
1710 ..FtrlConfig::default()
1711 })
1712 .is_err()
1713 );
1714 assert!(
1715 FtrlClassifier::new(FtrlConfig {
1716 l1: -1.0,
1717 ..FtrlConfig::default()
1718 })
1719 .is_err()
1720 );
1721 assert!(
1722 FtrlClassifier::new(FtrlConfig {
1723 l2: -1.0,
1724 ..FtrlConfig::default()
1725 })
1726 .is_err()
1727 );
1728 assert!(
1729 FtrlClassifier::new(FtrlConfig {
1730 alpha: f64::INFINITY,
1731 ..FtrlConfig::default()
1732 })
1733 .is_err()
1734 );
1735 }
1736
1737 #[test]
1738 #[cfg(feature = "serde")]
1739 fn classifier_serde_roundtrip() {
1740 let mut model = FtrlClassifier::new(FtrlConfig {
1741 alpha: 0.3,
1742 beta: 0.5,
1743 l1: 0.1,
1744 l2: 0.2,
1745 max_features: Some(100),
1746 new_feature_policy: NewFeaturePolicy::Reject,
1747 })
1748 .unwrap();
1749 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1750 model.learn(&sf, true).unwrap();
1751 model.learn(&sf, false).unwrap();
1752 let json = serde_json::to_string(&model).unwrap();
1753 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1754 assert_eq!(restored.samples_seen(), model.samples_seen());
1755 assert_eq!(restored.feature_count(), model.feature_count());
1756 let p1 = model.predict_proba(&sf).unwrap();
1757 let p2 = restored.predict_proba(&sf).unwrap();
1758 assert!((p1 - p2).abs() < 1e-12);
1759 }
1760
1761 #[test]
1762 fn predict_proba_in_range() {
1763 let mut model = FtrlClassifier::new(FtrlConfig {
1764 alpha: 0.5,
1765 beta: 1.0,
1766 l1: 0.0,
1767 l2: 0.0,
1768 max_features: None,
1769 new_feature_policy: NewFeaturePolicy::default(),
1770 })
1771 .unwrap();
1772 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1773 for _ in 0..200 {
1774 let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1775 let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1776 let y = x0 > 0.0;
1777 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1778 model.learn(&sf, y).unwrap();
1779 let p = model.predict_proba(&sf).unwrap();
1780 assert!(
1781 (0.0..=1.0).contains(&p),
1782 "probability must be in [0,1], got {p}"
1783 );
1784 }
1785 }
1786
1787 #[test]
1788 fn learn_improves_accuracy() {
1789 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1790 let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1792 .map(|_| {
1793 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1794 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1795 let y = x0 + x1 > 0.0;
1796 (
1797 SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1798 y,
1799 )
1800 })
1801 .collect();
1802
1803 let mut model = FtrlClassifier::new(FtrlConfig {
1804 alpha: 0.5,
1805 beta: 1.0,
1806 l1: 0.0,
1807 l2: 0.0,
1808 max_features: None,
1809 new_feature_policy: NewFeaturePolicy::default(),
1810 })
1811 .unwrap();
1812
1813 let acc_before: f64 = test_set
1815 .iter()
1816 .map(|(sf, y)| {
1817 let pred = model.predict(sf).unwrap();
1818 if pred == *y { 1.0 } else { 0.0 }
1819 })
1820 .sum::<f64>()
1821 / test_set.len() as f64;
1822
1823 for _ in 0..1000 {
1825 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1826 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1827 let y = x0 + x1 > 0.0;
1828 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1829 model.learn(&sf, y).unwrap();
1830 }
1831
1832 let acc_after: f64 = test_set
1833 .iter()
1834 .map(|(sf, y)| {
1835 let pred = model.predict(sf).unwrap();
1836 if pred == *y { 1.0 } else { 0.0 }
1837 })
1838 .sum::<f64>()
1839 / test_set.len() as f64;
1840
1841 assert!(
1842 acc_after > acc_before,
1843 "accuracy should improve: {acc_before} -> {acc_after}"
1844 );
1845 }
1846
1847 #[test]
1848 fn classifier_weights_returns_nonzero_only() {
1849 let mut model = FtrlClassifier::new(FtrlConfig {
1850 alpha: 0.5,
1851 beta: 1.0,
1852 l1: 0.0,
1853 l2: 0.0,
1854 max_features: None,
1855 new_feature_policy: NewFeaturePolicy::default(),
1856 })
1857 .unwrap();
1858 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1859 for _ in 0..50 {
1860 model.learn(&sf, true).unwrap();
1861 }
1862 let weights = model.weights();
1863 for &(_, w) in &weights {
1864 assert!(w != 0.0);
1865 }
1866 }
1867
1868 #[test]
1869 fn classifier_multiple_features() {
1870 let mut model = FtrlClassifier::new(FtrlConfig {
1871 alpha: 0.5,
1872 beta: 1.0,
1873 l1: 0.0,
1874 l2: 0.0,
1875 max_features: None,
1876 new_feature_policy: NewFeaturePolicy::default(),
1877 })
1878 .unwrap();
1879 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1880 for _ in 0..1000 {
1881 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1882 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1883 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1884 let y = x0 + x1 > 0.0;
1886 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1887 model.learn(&sf, y).unwrap();
1888 }
1889 let weights = model.weights();
1890 assert!(weights.iter().any(|&(id, _)| id == 0));
1892 assert!(weights.iter().any(|&(id, _)| id == 1));
1893 let p_pos = model
1895 .predict_proba(
1896 &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1897 )
1898 .unwrap();
1899 let p_neg = model
1900 .predict_proba(
1901 &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1902 )
1903 .unwrap();
1904 assert!(p_pos > 0.8);
1905 assert!(p_neg < 0.2);
1906 }
1907
1908 #[test]
1909 fn log_loss_converges() {
1910 let mut model = FtrlClassifier::new(FtrlConfig {
1912 alpha: 0.5,
1913 beta: 1.0,
1914 l1: 0.0,
1915 l2: 0.0,
1916 max_features: None,
1917 new_feature_policy: NewFeaturePolicy::default(),
1918 })
1919 .unwrap();
1920 let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1921 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1922 let mut first_loss = 0.0;
1923 let mut last_loss = 0.0;
1924 for i in 0..1000 {
1925 let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1926 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1927 let y = x0 > 0.0;
1928 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1929 let p = model.predict_proba(&sf).unwrap();
1930 let loss = loss_fn.loss(p, y);
1931 if i < 20 {
1932 first_loss += loss;
1933 }
1934 if i >= 980 {
1935 last_loss += loss;
1936 }
1937 model.learn(&sf, y).unwrap();
1938 }
1939 assert!(last_loss < first_loss, "log loss should decrease");
1940 }
1941
1942 #[test]
1947 fn classifier_overflow_does_not_mutate_state() {
1948 let mut model = FtrlClassifier::new(FtrlConfig {
1949 alpha: 0.1,
1950 beta: 1.0,
1951 l1: 0.0,
1952 l2: 0.0,
1953 max_features: None,
1954 new_feature_policy: NewFeaturePolicy::default(),
1955 })
1956 .unwrap();
1957 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1958 let result = model.learn(&sf, false);
1961 assert!(result.is_err(), "expected overflow error, got {result:?}");
1962 assert_eq!(model.samples_seen(), 0);
1963 assert_eq!(model.feature_count(), 0);
1964 assert!(model.params.is_empty());
1965 assert_eq!(model.intercept.z, 0.0);
1966 assert_eq!(model.intercept.n, 0.0);
1967 }
1968
1969 #[test]
1970 fn classifier_partial_update_is_atomic() {
1971 let mut model = FtrlClassifier::new(FtrlConfig {
1972 alpha: 0.1,
1973 beta: 1.0,
1974 l1: 0.0,
1975 l2: 0.0,
1976 max_features: None,
1977 new_feature_policy: NewFeaturePolicy::default(),
1978 })
1979 .unwrap();
1980 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1981 assert!(model.learn(&sf, false).is_err());
1982 assert!(!model.params.contains_key(&0));
1983 assert!(!model.params.contains_key(&1));
1984 assert_eq!(model.samples_seen(), 0);
1985 }
1986
1987 #[test]
1988 #[cfg(feature = "serde")]
1989 fn classifier_samples_seen_overflow_is_atomic() {
1990 let json = format!(
1991 "{{\"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\":{}}}",
1992 u64::MAX
1993 );
1994 let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1995 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1996 let result = model.learn(&sf, true);
1997 assert!(result.is_err(), "expected counter overflow");
1998 assert_eq!(model.samples_seen(), u64::MAX);
1999 assert_eq!(model.feature_count(), 0);
2000 assert_eq!(model.intercept.z, 0.0);
2001 assert_eq!(model.intercept.n, 0.0);
2002 }
2003
2004 #[test]
2009 fn classifier_max_features_reject_at_limit() {
2010 let mut model = FtrlClassifier::new(FtrlConfig {
2011 alpha: 0.5,
2012 beta: 1.0,
2013 l1: 0.0,
2014 l2: 0.0,
2015 max_features: Some(2),
2016 new_feature_policy: NewFeaturePolicy::Reject,
2017 })
2018 .unwrap();
2019 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2020 model.learn(&sf, true).unwrap();
2021 assert_eq!(model.feature_count(), 2);
2022 let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
2023 assert!(model.learn(&sf_new, true).is_err());
2024 assert_eq!(model.feature_count(), 2);
2025 assert_eq!(model.samples_seen(), 1);
2026 }
2027
2028 #[test]
2029 fn classifier_max_features_ignore_skips_new() {
2030 let mut model = FtrlClassifier::new(FtrlConfig {
2031 alpha: 0.5,
2032 beta: 1.0,
2033 l1: 0.0,
2034 l2: 0.0,
2035 max_features: Some(2),
2036 new_feature_policy: NewFeaturePolicy::Ignore,
2037 })
2038 .unwrap();
2039 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2040 model.learn(&sf, true).unwrap();
2041 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
2042 model.learn(&sf_mixed, false).unwrap();
2043 assert_eq!(model.feature_count(), 2);
2044 assert!(!model.params.contains_key(&2));
2045 assert_eq!(model.samples_seen(), 2);
2046 }
2047
2048 #[test]
2049 fn classifier_max_features_multi_new_prejudge() {
2050 let mut model = FtrlClassifier::new(FtrlConfig {
2051 alpha: 0.5,
2052 beta: 1.0,
2053 l1: 0.0,
2054 l2: 0.0,
2055 max_features: Some(2),
2056 new_feature_policy: NewFeaturePolicy::Reject,
2057 })
2058 .unwrap();
2059 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
2060 assert!(model.learn(&sf, true).is_err());
2061 assert_eq!(model.feature_count(), 0);
2062 assert_eq!(model.samples_seen(), 0);
2063 }
2064
2065 #[test]
2070 #[cfg(feature = "serde")]
2071 fn classifier_serde_rejects_negative_n() {
2072 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}";
2073 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2074 assert!(result.is_err(), "negative n must be rejected");
2075 }
2076
2077 #[test]
2078 #[cfg(feature = "serde")]
2079 fn classifier_serde_rejects_invalid_config() {
2080 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}";
2081 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2082 assert!(result.is_err(), "invalid alpha must be rejected");
2083 }
2084
2085 #[test]
2086 #[cfg(feature = "serde")]
2087 fn classifier_serde_accepts_missing_optional_fields() {
2088 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}";
2089 let model: FtrlClassifier = serde_json::from_str(json).unwrap();
2090 assert!(model.config().max_features.is_none());
2091 assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
2092 }
2093
2094 #[test]
2099 fn regressor_gradient_squared_underflow_is_atomic() {
2100 let mut model = FtrlRegressor::new(FtrlConfig {
2106 alpha: 1.0,
2107 beta: 0.0,
2108 l1: 0.0,
2109 l2: 0.0,
2110 max_features: None,
2111 new_feature_policy: NewFeaturePolicy::default(),
2112 })
2113 .unwrap();
2114 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2115 let result = model.learn(&sf, -1e-200);
2116 assert!(result.is_err(), "expected underflow error, got {result:?}");
2117 assert_eq!(model.samples_seen(), 0);
2118 assert_eq!(model.feature_count(), 0);
2119 assert!(model.params.is_empty());
2120 assert_eq!(model.intercept.z, 0.0);
2121 assert_eq!(model.intercept.n, 0.0);
2122 }
2123
2124 #[test]
2125 fn classifier_gradient_squared_underflow_is_atomic() {
2126 let mut model = FtrlClassifier::new(FtrlConfig {
2130 alpha: 1.0,
2131 beta: 0.0,
2132 l1: 0.0,
2133 l2: 0.0,
2134 max_features: None,
2135 new_feature_policy: NewFeaturePolicy::default(),
2136 })
2137 .unwrap();
2138 let sf = SparseFeatures::from_sorted(vec![(0, 1e-200)]).unwrap();
2139 let result = model.learn(&sf, false);
2140 assert!(result.is_err(), "expected underflow error, got {result:?}");
2141 assert_eq!(model.samples_seen(), 0);
2142 assert_eq!(model.feature_count(), 0);
2143 }
2144
2145 #[test]
2146 fn regressor_boundary_config_predict_after_learn_always_finite() {
2147 let mut model = FtrlRegressor::new(FtrlConfig {
2151 alpha: 1.0,
2152 beta: 0.0,
2153 l1: 0.0,
2154 l2: 0.0,
2155 max_features: None,
2156 new_feature_policy: NewFeaturePolicy::default(),
2157 })
2158 .unwrap();
2159 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2160 for _ in 0..100 {
2161 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2162 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2163 let y = 2.0 * x0 - x1;
2164 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2165 model.learn(&sf, y).unwrap();
2166 let pred = model.predict(&sf);
2167 assert!(
2168 pred.is_ok(),
2169 "predict failed after successful learn: {pred:?}"
2170 );
2171 assert!(
2172 pred.unwrap().is_finite(),
2173 "predict must return finite value after successful learn"
2174 );
2175 }
2176 }
2177
2178 #[test]
2179 fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2180 let mut model = FtrlClassifier::new(FtrlConfig {
2181 alpha: 1.0,
2182 beta: 0.0,
2183 l1: 0.0,
2184 l2: 0.0,
2185 max_features: None,
2186 new_feature_policy: NewFeaturePolicy::default(),
2187 })
2188 .unwrap();
2189 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2190 for _ in 0..100 {
2191 let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2192 let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2193 let y = x0 > 0.0;
2194 let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2195 model.learn(&sf, y).unwrap();
2196 let proba = model.predict_proba(&sf);
2197 assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2198 let p = proba.unwrap();
2199 assert!(p.is_finite(), "probability must be finite, got {p}");
2200 assert!(
2201 (0.0..=1.0).contains(&p),
2202 "probability must be in [0,1], got {p}"
2203 );
2204 }
2205 }
2206
2207 #[test]
2208 #[cfg(feature = "serde")]
2209 fn regressor_serde_rejects_n_zero_z_nonzero() {
2210 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}";
2214 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2215 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2216 }
2217
2218 #[test]
2219 #[cfg(feature = "serde")]
2220 fn classifier_serde_rejects_n_zero_z_nonzero() {
2221 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}";
2222 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2223 assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2224 }
2225
2226 #[test]
2227 #[cfg(feature = "serde")]
2228 fn regressor_predict_dot_plus_intercept_overflow() {
2229 let z = -f64::MAX * 0.75;
2236 let json = format!(
2237 "{{\"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}}",
2238 z
2239 );
2240 let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2241 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2242 let result = model.predict(&sf);
2243 assert!(
2244 result.is_err(),
2245 "expected dot+intercept overflow error, got {result:?}"
2246 );
2247 }
2248
2249 #[test]
2250 fn regressor_ignore_skips_overflowing_new_feature() {
2251 let mut model = FtrlRegressor::new(FtrlConfig {
2252 alpha: 0.5,
2253 beta: 1.0,
2254 l1: 0.0,
2255 l2: 0.0,
2256 max_features: Some(1),
2257 new_feature_policy: NewFeaturePolicy::Ignore,
2258 })
2259 .unwrap();
2260 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2261 model.learn(&sf, 1.0).unwrap();
2262 assert_eq!(model.feature_count(), 1);
2263
2264 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2269 let result = model.learn(&sf_mixed, 1.0);
2270 assert!(
2271 result.is_ok(),
2272 "Ignore must skip overflowing new feature, got {result:?}"
2273 );
2274 assert_eq!(model.feature_count(), 1);
2275 assert!(!model.params.contains_key(&1));
2276 assert_eq!(model.samples_seen(), 2);
2277 assert!(model.predict(&sf).is_ok());
2278 }
2279
2280 #[test]
2281 fn classifier_ignore_skips_overflowing_new_feature() {
2282 let mut model = FtrlClassifier::new(FtrlConfig {
2283 alpha: 0.5,
2284 beta: 1.0,
2285 l1: 0.0,
2286 l2: 0.0,
2287 max_features: Some(1),
2288 new_feature_policy: NewFeaturePolicy::Ignore,
2289 })
2290 .unwrap();
2291 let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2292 model.learn(&sf, true).unwrap();
2293 assert_eq!(model.feature_count(), 1);
2294
2295 let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2296 let result = model.learn(&sf_mixed, true);
2297 assert!(
2298 result.is_ok(),
2299 "Ignore must skip overflowing new feature, got {result:?}"
2300 );
2301 assert_eq!(model.feature_count(), 1);
2302 assert!(!model.params.contains_key(&1));
2303 assert_eq!(model.samples_seen(), 2);
2304 assert!(model.predict_proba(&sf).is_ok());
2305 }
2306
2307 #[test]
2312 #[cfg(feature = "serde")]
2313 fn regressor_serde_rejects_config_dependent_zero_denominator() {
2314 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}";
2319 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2320 assert!(
2321 result.is_err(),
2322 "config-dependent zero denominator must be rejected"
2323 );
2324 }
2325
2326 #[test]
2327 #[cfg(feature = "serde")]
2328 fn classifier_serde_rejects_config_dependent_zero_denominator() {
2329 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}";
2330 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2331 assert!(
2332 result.is_err(),
2333 "config-dependent zero denominator must be rejected"
2334 );
2335 }
2336
2337 #[test]
2338 #[cfg(feature = "serde")]
2339 fn regressor_serde_rejects_intercept_zero_denominator() {
2340 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}";
2343 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2344 assert!(
2345 result.is_err(),
2346 "intercept zero denominator must be rejected"
2347 );
2348 }
2349
2350 #[test]
2351 #[cfg(feature = "serde")]
2352 fn classifier_serde_rejects_intercept_zero_denominator() {
2353 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}";
2354 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2355 assert!(
2356 result.is_err(),
2357 "intercept zero denominator must be rejected"
2358 );
2359 }
2360
2361 #[test]
2362 #[cfg(feature = "serde")]
2363 fn regressor_valid_boundary_state_roundtrips() {
2364 let mut model = FtrlRegressor::new(FtrlConfig {
2368 alpha: 1.0,
2369 beta: 0.0,
2370 l1: 0.0,
2371 l2: 0.0,
2372 max_features: None,
2373 new_feature_policy: NewFeaturePolicy::default(),
2374 })
2375 .unwrap();
2376 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2377 model.learn(&sf, 3.0).unwrap();
2378 model.learn(&sf, 5.0).unwrap();
2379 let json = serde_json::to_string(&model).unwrap();
2380 let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2381 assert_eq!(restored.samples_seen(), model.samples_seen());
2382 assert_eq!(restored.feature_count(), model.feature_count());
2383 let p1 = model.predict(&sf).unwrap();
2384 let p2 = restored.predict(&sf).unwrap();
2385 assert!((p1 - p2).abs() < 1e-12);
2386 for (_, w) in restored.weights() {
2388 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2389 }
2390 assert!(restored.intercept().is_finite());
2391 }
2392
2393 #[test]
2394 #[cfg(feature = "serde")]
2395 fn classifier_valid_boundary_state_roundtrips() {
2396 let mut model = FtrlClassifier::new(FtrlConfig {
2397 alpha: 1.0,
2398 beta: 0.0,
2399 l1: 0.0,
2400 l2: 0.0,
2401 max_features: None,
2402 new_feature_policy: NewFeaturePolicy::default(),
2403 })
2404 .unwrap();
2405 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2406 model.learn(&sf, true).unwrap();
2407 model.learn(&sf, false).unwrap();
2408 let json = serde_json::to_string(&model).unwrap();
2409 let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2410 assert_eq!(restored.samples_seen(), model.samples_seen());
2411 assert_eq!(restored.feature_count(), model.feature_count());
2412 let p1 = model.predict_proba(&sf).unwrap();
2413 let p2 = restored.predict_proba(&sf).unwrap();
2414 assert!((p1 - p2).abs() < 1e-12);
2415 for (_, w) in restored.weights() {
2416 assert!(w.is_finite(), "restored weight must be finite, got {w}");
2417 }
2418 assert!(restored.intercept().is_finite());
2419 }
2420
2421 #[test]
2426 #[cfg(feature = "serde")]
2427 fn regressor_serde_rejects_params_above_max_features() {
2428 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}";
2431 let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2432 let err = match result {
2433 Ok(_) => panic!("expected serde error, got Ok"),
2434 Err(e) => e,
2435 };
2436 let msg = err.to_string();
2437 assert!(
2438 msg.contains("max_features") && msg.contains("feature count"),
2439 "error must mention feature count / max_features, got: {msg}"
2440 );
2441 }
2442
2443 #[test]
2444 #[cfg(feature = "serde")]
2445 fn classifier_serde_rejects_params_above_max_features() {
2446 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}";
2447 let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2448 let err = match result {
2449 Ok(_) => panic!("expected serde error, got Ok"),
2450 Err(e) => e,
2451 };
2452 let msg = err.to_string();
2453 assert!(
2454 msg.contains("max_features") && msg.contains("feature count"),
2455 "error must mention feature count / max_features, got: {msg}"
2456 );
2457 }
2458
2459 #[test]
2460 #[cfg(feature = "serde")]
2461 fn regressor_serde_accepts_params_equal_to_max_features() {
2462 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}";
2465 let model: FtrlRegressor =
2466 serde_json::from_str(json).expect("equal count must be accepted");
2467 assert_eq!(model.feature_count(), 2);
2468 assert_eq!(model.samples_seen(), 3);
2469 for (_, w) in model.weights() {
2470 assert!(w.is_finite(), "weight must be finite, got {w}");
2471 }
2472 assert!(model.intercept().is_finite());
2473 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2474 let pred = model.predict(&sf).expect("predict must succeed");
2475 assert!(pred.is_finite(), "prediction must be finite, got {pred}");
2476 }
2477
2478 #[test]
2479 #[cfg(feature = "serde")]
2480 fn classifier_serde_accepts_params_equal_to_max_features() {
2481 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}";
2482 let model: FtrlClassifier =
2483 serde_json::from_str(json).expect("equal count must be accepted");
2484 assert_eq!(model.feature_count(), 2);
2485 assert_eq!(model.samples_seen(), 3);
2486 for (_, w) in model.weights() {
2487 assert!(w.is_finite(), "weight must be finite, got {w}");
2488 }
2489 assert!(model.intercept().is_finite());
2490 let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2491 let p = model
2492 .predict_proba(&sf)
2493 .expect("predict_proba must succeed");
2494 assert!(p.is_finite(), "probability must be finite, got {p}");
2495 assert!(
2496 (0.0..=1.0).contains(&p),
2497 "probability must be in [0,1], got {p}"
2498 );
2499 }
2500
2501 #[test]
2502 #[cfg(feature = "serde")]
2503 fn regressor_serde_allows_unbounded_params_when_max_features_none() {
2504 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}";
2507 let model: FtrlRegressor =
2508 serde_json::from_str(json).expect("unbounded state must be accepted");
2509 assert_eq!(model.feature_count(), 3);
2510 for (_, w) in model.weights() {
2511 assert!(w.is_finite(), "weight must be finite, got {w}");
2512 }
2513 assert!(model.intercept().is_finite());
2514 }
2515
2516 #[test]
2517 #[cfg(feature = "serde")]
2518 fn classifier_serde_allows_unbounded_params_when_max_features_none() {
2519 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}";
2520 let model: FtrlClassifier =
2521 serde_json::from_str(json).expect("unbounded state must be accepted");
2522 assert_eq!(model.feature_count(), 3);
2523 for (_, w) in model.weights() {
2524 assert!(w.is_finite(), "weight must be finite, got {w}");
2525 }
2526 assert!(model.intercept().is_finite());
2527 }
2528}