Skip to main content

rill_ml/models/
ftrl.rs

1//! FTRL-Proximal online learning for sparse features.
2//!
3//! Implements the Follow-The-Regularized-Leader Proximal algorithm,
4//! which is well-suited for high-dimensional sparse data. L1
5//! regularization produces sparse weight vectors, and the per-coordinate
6//! learning rate adapts to feature frequency.
7//!
8//! See: McMahan et al., "Ad Click Prediction: a View from the Trenches"
9//! (KDD 2013).
10//!
11//! # Per-coordinate learning rate
12//!
13//! `eta_i = alpha / (beta + sqrt(n_i))`
14//!
15//! # Weight computation
16//!
17//! For feature `i`:
18//!
19//! ```text
20//! if |z_i| <= lambda1:
21//!     w_i = 0
22//! else:
23//!     w_i = -(z_i - sign(z_i) * lambda1) / (lambda2 + (beta + sqrt(n_i)) / alpha)
24//! ```
25//!
26//! The intercept uses `lambda1 = 0` (no L1 regularization).
27//!
28//! # Failure atomicity
29//!
30//! [`FtrlRegressor::learn`] and [`FtrlClassifier::learn`] compute the next
31//! state for every affected entry before mutating any field. If any
32//! intermediate value is non-finite or the sample counter overflows, the
33//! call returns `Err` and the model state is unchanged.
34//!
35//! # Dynamic feature growth
36//!
37//! By default the feature table grows without bound. Long-running
38//! deployments handling untrusted streams should set
39//! [`FtrlConfig::max_features`] to a finite value and pick a
40//! [`NewFeaturePolicy`]. `None` preserves backwards compatibility but
41//! allows memory growth proportional to the number of distinct
42//! `FeatureId`s ever observed.
43
44use 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/// Policy for handling new `FeatureId`s once [`FtrlConfig::max_features`]
51/// has been reached.
52///
53/// `max_features` only constrains feature insertion; features already in
54/// the model continue to train regardless of the policy.
55#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
56#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
57pub enum NewFeaturePolicy {
58    /// Reject the entire `learn` call with [`RillError::InvalidState`] when
59    /// adding the new features in the current sample would exceed
60    /// `max_features`. The call is failure-atomic: no state changes.
61    #[default]
62    Reject,
63    /// Keep training existing features but silently skip any new feature
64    /// that would exceed `max_features`. The sample counter still
65    /// advances and existing features still receive their gradient update.
66    Ignore,
67}
68
69/// Configuration for FTRL models.
70///
71/// Controls the per-coordinate learning rate and regularization strengths.
72/// All fields must be finite; `alpha` must be strictly positive and the
73/// regularization parameters must be non-negative.
74#[derive(Debug, Clone)]
75#[cfg_attr(feature = "serde", derive(serde::Serialize))]
76pub struct FtrlConfig {
77    /// Alpha: learning rate scaling. Must be `> 0`.
78    pub alpha: f64,
79    /// Beta: smoothing constant. Must be `>= 0`.
80    pub beta: f64,
81    /// L1 regularization strength. Must be `>= 0`.
82    pub l1: f64,
83    /// L2 regularization strength. Must be `>= 0`.
84    pub l2: f64,
85    /// Maximum number of distinct features the model will store.
86    ///
87    /// `None` (the default) allows unbounded growth, which is
88    /// backwards-compatible but unsafe for long-running services that
89    /// consume untrusted feature streams. Set to a finite value to
90    /// bound memory and trigger the [`NewFeaturePolicy`].
91    pub max_features: Option<usize>,
92    /// Policy applied when `max_features` is set and a new `FeatureId`
93    /// would exceed the cap.
94    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    /// Validate configuration parameters.
112    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/// Per-feature FTRL state.
186///
187/// Tracks the sum of (sigma-corrected) gradients `z` and the sum of squared
188/// gradients `n`. The per-coordinate adaptive learning rate is derived from
189/// `n`: features seen more frequently get smaller steps.
190#[derive(Debug, Clone, Default)]
191#[cfg_attr(feature = "serde", derive(serde::Serialize))]
192pub struct FtrlParam {
193    /// Sum of gradients (with sigma correction).
194    z: f64,
195    /// Sum of squared gradients. Must remain finite and non-negative.
196    n: f64,
197}
198
199impl FtrlParam {
200    /// Compute the FTRL weight given the config.
201    ///
202    /// Returns `0` when `|z| <= l1` (L1 soft-thresholding).
203    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    /// Compute the intercept weight (no L1 regularization).
215    ///
216    /// Returns `0` when no gradient has been observed yet (`n == 0`),
217    /// avoiding a potential `0 / 0` when `l2` and `beta` are both zero.
218    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    /// Compute the next `(z, n)` state after applying a gradient, without
229    /// mutating `self`.
230    ///
231    /// `sigma = (sqrt(n_new) - sqrt(n_old)) / alpha`
232    /// `z_new = z + g - sigma * w`
233    /// `n_new = n + g^2`
234    ///
235    /// Every intermediate result is checked for finiteness so a partial
236    /// overflow cannot poison state. The caller is responsible for
237    /// committing the returned values atomically.
238    fn next_updated(
239        &self,
240        gradient: f64,
241        weight: f64,
242        config: &FtrlConfig,
243    ) -> Result<(f64, f64), RillError> {
244        let gradient_sq = gradient * gradient;
245        ensure_finite("ftrl_gradient_squared", gradient_sq)?;
246        let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
247        let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
248        ensure_finite("ftrl_sigma", sigma)?;
249        let sigma_w = sigma * weight;
250        ensure_finite("ftrl_sigma_weight", sigma_w)?;
251        let z_delta = gradient - sigma_w;
252        ensure_finite("ftrl_z_delta", z_delta)?;
253        let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
254        Ok((z_new, n_new))
255    }
256
257    /// Validate that `z` is finite and `n` is finite and non-negative.
258    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
259    fn validate(&self) -> Result<(), RillError> {
260        ensure_finite("ftrl_z", self.z)?;
261        ensure_finite("ftrl_n", self.n)?;
262        if self.n < 0.0 {
263            return Err(RillError::InvalidState(format!(
264                "ftrl n must be non-negative, got {0}",
265                self.n
266            )));
267        }
268        Ok(())
269    }
270}
271
272#[cfg(feature = "serde")]
273impl<'de> serde::Deserialize<'de> for FtrlParam {
274    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
275    where
276        D: serde::Deserializer<'de>,
277    {
278        #[derive(serde::Deserialize)]
279        struct FtrlParamState {
280            z: f64,
281            n: f64,
282        }
283
284        let state = FtrlParamState::deserialize(deserializer)?;
285        let param = FtrlParam {
286            z: state.z,
287            n: state.n,
288        };
289        param.validate().map_err(serde::de::Error::custom)?;
290        Ok(param)
291    }
292}
293
294/// Compute the dot product `w · x` over sparse features.
295///
296/// Iterates only over the features present in `features` (not all stored
297/// params), looking up each feature's current FTRL weight. Each
298/// contribution and the running sum are checked for finiteness so a
299/// single overflowing term cannot poison the prediction.
300fn compute_dot(
301    params: &BTreeMap<FeatureId, FtrlParam>,
302    config: &FtrlConfig,
303    features: &SparseFeatures,
304) -> Result<f64, RillError> {
305    if features.is_empty() {
306        return Err(RillError::EmptyFeatures);
307    }
308    let mut dot = 0.0;
309    for &(id, value) in features.values() {
310        ensure_finite("sparse_value", value)?;
311        if let Some(param) = params.get(&id) {
312            let w = param.weight(config);
313            ensure_finite("ftrl_weight", w)?;
314            let contribution = w * value;
315            ensure_finite("ftrl_dot_contribution", contribution)?;
316            dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
317        }
318    }
319    Ok(dot)
320}
321
322/// FTRL regressor with squared loss.
323///
324/// Learns `y ≈ w · x + b` incrementally. The gradient of the squared loss
325/// w.r.t. the prediction is `prediction - target`, so each feature's
326/// gradient is `(prediction - target) * x_i`.
327///
328/// # Examples
329///
330/// ```
331/// use rill_ml::models::{FtrlConfig, FtrlRegressor};
332/// use rill_ml::sparse::SparseFeatures;
333/// use rill_ml::SparseRegressor;
334///
335/// let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
336/// let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
337/// let _pred = model.predict(&sf).unwrap();
338/// model.learn(&sf, 3.0).unwrap();
339/// ```
340#[derive(Debug, Clone)]
341#[cfg_attr(feature = "serde", derive(serde::Serialize))]
342pub struct FtrlRegressor {
343    config: FtrlConfig,
344    params: BTreeMap<FeatureId, FtrlParam>,
345    intercept: FtrlParam,
346    samples_seen: u64,
347}
348
349impl FtrlRegressor {
350    /// Create a new FTRL regressor.
351    ///
352    /// Returns an error if the configuration is invalid.
353    pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
354        config.validate()?;
355        Ok(Self {
356            config,
357            params: BTreeMap::new(),
358            intercept: FtrlParam::default(),
359            samples_seen: 0,
360        })
361    }
362
363    /// The model configuration.
364    pub const fn config(&self) -> &FtrlConfig {
365        &self.config
366    }
367
368    /// Return the current non-zero feature weights, sorted by `FeatureId`.
369    ///
370    /// Features whose FTRL weight is exactly zero (due to L1
371    /// soft-thresholding or never having been updated) are excluded.
372    pub fn weights(&self) -> Vec<(FeatureId, f64)> {
373        self.params
374            .iter()
375            .map(|(&id, param)| (id, param.weight(&self.config)))
376            .filter(|&(_, w)| w != 0.0)
377            .collect()
378    }
379
380    /// Compute the current intercept (bias) weight.
381    pub fn intercept(&self) -> f64 {
382        self.intercept.intercept_weight(&self.config)
383    }
384
385    /// Number of distinct features the model has seen.
386    pub fn feature_count(&self) -> usize {
387        self.params.len()
388    }
389
390    /// Compute the raw prediction `w · x + b` without updating state.
391    fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
392        let dot = compute_dot(&self.params, &self.config, features)?;
393        let intercept = self.intercept.intercept_weight(&self.config);
394        ensure_finite("ftrl_intercept", intercept)?;
395        Ok(dot + intercept)
396    }
397}
398
399#[cfg(feature = "serde")]
400impl<'de> serde::Deserialize<'de> for FtrlRegressor {
401    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
402    where
403        D: serde::Deserializer<'de>,
404    {
405        #[derive(serde::Deserialize)]
406        struct FtrlRegressorState {
407            config: FtrlConfig,
408            params: BTreeMap<FeatureId, FtrlParam>,
409            intercept: FtrlParam,
410            samples_seen: u64,
411        }
412
413        let state = FtrlRegressorState::deserialize(deserializer)?;
414        let model = FtrlRegressor {
415            config: state.config,
416            params: state.params,
417            intercept: state.intercept,
418            samples_seen: state.samples_seen,
419        };
420        // config and each FtrlParam are already validated by their own
421        // Deserialize impls; only the top-level invariants remain.
422        model
423            .validate_invariants()
424            .map_err(serde::de::Error::custom)?;
425        Ok(model)
426    }
427}
428
429impl FtrlRegressor {
430    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
431    fn validate_invariants(&self) -> Result<(), RillError> {
432        // config and individual params are validated at deserialization;
433        // nothing top-level to check beyond that.
434        self.config.validate()?;
435        self.intercept.validate()?;
436        for param in self.params.values() {
437            param.validate()?;
438        }
439        Ok(())
440    }
441}
442
443impl SparseRegressor for FtrlRegressor {
444    fn samples_seen(&self) -> u64 {
445        self.samples_seen
446    }
447
448    fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
449        self.predict_inner(features)
450    }
451
452    fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
453        if features.is_empty() {
454            return Err(RillError::EmptyFeatures);
455        }
456        ensure_finite("target", target)?;
457
458        // Reserve the sample counter first. If it overflows, no state changes.
459        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
460
461        let prediction = self.predict_inner(features)?;
462        ensure_finite("ftrl_prediction", prediction)?;
463        let grad = prediction - target;
464        ensure_finite("ftrl_gradient", grad)?;
465
466        // Pre-judge the max_features policy. We never insert a subset of
467        // new features then fail: either all new features are accepted or
468        // (under Ignore) all new features are skipped.
469        let new_ids_count = features
470            .values()
471            .iter()
472            .filter(|(id, _)| !self.params.contains_key(id))
473            .count();
474        let mut skip_new_features = false;
475        if let Some(max_features) = self.config.max_features {
476            let projected = self.params.len().saturating_add(new_ids_count);
477            if projected > max_features {
478                match self.config.new_feature_policy {
479                    NewFeaturePolicy::Reject => {
480                        return Err(RillError::InvalidState(format!(
481                            "FTRL feature count {projected} exceeds max_features {max_features}"
482                        )));
483                    }
484                    NewFeaturePolicy::Ignore => {
485                        skip_new_features = true;
486                    }
487                }
488            }
489        }
490
491        // Compute the next (z, n) for every feature without touching self.
492        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
493        for &(id, value) in features.values() {
494            let g = grad * value;
495            ensure_finite("ftrl_feature_gradient", g)?;
496
497            let is_new = !self.params.contains_key(&id);
498            if is_new && skip_new_features {
499                continue;
500            }
501
502            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
503                let w = param.weight(&self.config);
504                param.next_updated(g, w, &self.config)?
505            } else {
506                let param = FtrlParam::default();
507                let w = param.weight(&self.config);
508                param.next_updated(g, w, &self.config)?
509            };
510            updates.push((id, new_z, new_n));
511        }
512
513        // Compute the next intercept state.
514        let w_b = self.intercept.intercept_weight(&self.config);
515        let (new_intercept_z, new_intercept_n) =
516            self.intercept.next_updated(grad, w_b, &self.config)?;
517
518        // Commit atomically. No failure path beyond this point.
519        for (id, new_z, new_n) in updates {
520            let param = self.params.entry(id).or_default();
521            param.z = new_z;
522            param.n = new_n;
523        }
524        self.intercept.z = new_intercept_z;
525        self.intercept.n = new_intercept_n;
526        self.samples_seen = next_samples_seen;
527
528        Ok(())
529    }
530
531    fn reset(&mut self) {
532        self.params.clear();
533        self.intercept = FtrlParam::default();
534        self.samples_seen = 0;
535    }
536}
537
538/// FTRL binary classifier with log loss.
539///
540/// Predicts `P(y=1 | x) = sigmoid(w · x + b)`. The gradient of the log loss
541/// w.r.t. the logit simplifies to `probability - target`, so each feature's
542/// gradient is `(probability - target) * x_i`.
543///
544/// The returned probability lies in `[0, 1]`. Extreme logits can produce
545/// exactly `0.0` or `1.0` after `sigmoid`; downstream consumers such as
546/// [`crate::loss::log_loss::BinaryLogLoss`] clip internally.
547///
548/// # Examples
549///
550/// ```
551/// use rill_ml::models::{FtrlClassifier, FtrlConfig};
552/// use rill_ml::sparse::SparseFeatures;
553/// use rill_ml::SparseClassifier;
554///
555/// let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
556/// let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
557/// let _proba = model.predict_proba(&sf).unwrap();
558/// model.learn(&sf, true).unwrap();
559/// ```
560#[derive(Debug, Clone)]
561#[cfg_attr(feature = "serde", derive(serde::Serialize))]
562pub struct FtrlClassifier {
563    config: FtrlConfig,
564    params: BTreeMap<FeatureId, FtrlParam>,
565    intercept: FtrlParam,
566    samples_seen: u64,
567}
568
569impl FtrlClassifier {
570    /// Create a new FTRL classifier.
571    ///
572    /// Returns an error if the configuration is invalid.
573    pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
574        config.validate()?;
575        Ok(Self {
576            config,
577            params: BTreeMap::new(),
578            intercept: FtrlParam::default(),
579            samples_seen: 0,
580        })
581    }
582
583    /// The model configuration.
584    pub const fn config(&self) -> &FtrlConfig {
585        &self.config
586    }
587
588    /// Return the current non-zero feature weights, sorted by `FeatureId`.
589    ///
590    /// Features whose FTRL weight is exactly zero (due to L1
591    /// soft-thresholding or never having been updated) are excluded.
592    pub fn weights(&self) -> Vec<(FeatureId, f64)> {
593        self.params
594            .iter()
595            .map(|(&id, param)| (id, param.weight(&self.config)))
596            .filter(|&(_, w)| w != 0.0)
597            .collect()
598    }
599
600    /// Compute the current intercept (bias) weight.
601    pub fn intercept(&self) -> f64 {
602        self.intercept.intercept_weight(&self.config)
603    }
604
605    /// Number of distinct features the model has seen.
606    pub fn feature_count(&self) -> usize {
607        self.params.len()
608    }
609
610    /// Compute the probability `sigmoid(w · x + b)` without updating state.
611    fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
612        let dot = compute_dot(&self.params, &self.config, features)?;
613        let intercept = self.intercept.intercept_weight(&self.config);
614        ensure_finite("ftrl_intercept", intercept)?;
615        let logit = dot + intercept;
616        ensure_finite("ftrl_logit", logit)?;
617        Ok(sigmoid(logit))
618    }
619}
620
621#[cfg(feature = "serde")]
622impl<'de> serde::Deserialize<'de> for FtrlClassifier {
623    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
624    where
625        D: serde::Deserializer<'de>,
626    {
627        #[derive(serde::Deserialize)]
628        struct FtrlClassifierState {
629            config: FtrlConfig,
630            params: BTreeMap<FeatureId, FtrlParam>,
631            intercept: FtrlParam,
632            samples_seen: u64,
633        }
634
635        let state = FtrlClassifierState::deserialize(deserializer)?;
636        let model = FtrlClassifier {
637            config: state.config,
638            params: state.params,
639            intercept: state.intercept,
640            samples_seen: state.samples_seen,
641        };
642        model
643            .validate_invariants()
644            .map_err(serde::de::Error::custom)?;
645        Ok(model)
646    }
647}
648
649impl FtrlClassifier {
650    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
651    fn validate_invariants(&self) -> Result<(), RillError> {
652        self.config.validate()?;
653        self.intercept.validate()?;
654        for param in self.params.values() {
655            param.validate()?;
656        }
657        Ok(())
658    }
659}
660
661impl SparseClassifier for FtrlClassifier {
662    fn samples_seen(&self) -> u64 {
663        self.samples_seen
664    }
665
666    fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
667        self.predict_proba_inner(features)
668    }
669
670    fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
671        if features.is_empty() {
672            return Err(RillError::EmptyFeatures);
673        }
674
675        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
676
677        let probability = self.predict_proba_inner(features)?;
678        ensure_finite("ftrl_probability", probability)?;
679        let y = if target { 1.0 } else { 0.0 };
680        let grad = probability - y;
681        ensure_finite("ftrl_gradient", grad)?;
682
683        let new_ids_count = features
684            .values()
685            .iter()
686            .filter(|(id, _)| !self.params.contains_key(id))
687            .count();
688        let mut skip_new_features = false;
689        if let Some(max_features) = self.config.max_features {
690            let projected = self.params.len().saturating_add(new_ids_count);
691            if projected > max_features {
692                match self.config.new_feature_policy {
693                    NewFeaturePolicy::Reject => {
694                        return Err(RillError::InvalidState(format!(
695                            "FTRL feature count {projected} exceeds max_features {max_features}"
696                        )));
697                    }
698                    NewFeaturePolicy::Ignore => {
699                        skip_new_features = true;
700                    }
701                }
702            }
703        }
704
705        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
706        for &(id, value) in features.values() {
707            let g = grad * value;
708            ensure_finite("ftrl_feature_gradient", g)?;
709
710            let is_new = !self.params.contains_key(&id);
711            if is_new && skip_new_features {
712                continue;
713            }
714
715            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
716                let w = param.weight(&self.config);
717                param.next_updated(g, w, &self.config)?
718            } else {
719                let param = FtrlParam::default();
720                let w = param.weight(&self.config);
721                param.next_updated(g, w, &self.config)?
722            };
723            updates.push((id, new_z, new_n));
724        }
725
726        let w_b = self.intercept.intercept_weight(&self.config);
727        let (new_intercept_z, new_intercept_n) =
728            self.intercept.next_updated(grad, w_b, &self.config)?;
729
730        for (id, new_z, new_n) in updates {
731            let param = self.params.entry(id).or_default();
732            param.z = new_z;
733            param.n = new_n;
734        }
735        self.intercept.z = new_intercept_z;
736        self.intercept.n = new_intercept_n;
737        self.samples_seen = next_samples_seen;
738
739        Ok(())
740    }
741
742    fn reset(&mut self) {
743        self.params.clear();
744        self.intercept = FtrlParam::default();
745        self.samples_seen = 0;
746    }
747}
748
749#[cfg(test)]
750mod tests {
751    use super::*;
752    use rand::SeedableRng;
753
754    // -----------------------------------------------------------------
755    // FtrlRegressor tests
756    // -----------------------------------------------------------------
757
758    #[test]
759    fn cold_start_returns_zero() {
760        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
761        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
762        let pred = model.predict(&sf).unwrap();
763        assert!(pred.abs() < 1e-12);
764    }
765
766    #[test]
767    fn learn_linear_data_converges() {
768        // y = 2 * x, single feature
769        let mut model = FtrlRegressor::new(FtrlConfig {
770            alpha: 0.5,
771            beta: 1.0,
772            l1: 0.0,
773            l2: 0.0,
774            max_features: None,
775            new_feature_policy: NewFeaturePolicy::default(),
776        })
777        .unwrap();
778        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
779        let mut first_err = 0.0;
780        let mut last_err = 0.0;
781        for i in 0..500 {
782            let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
783            let y = 2.0 * x;
784            let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
785            let pred = model.predict(&sf).unwrap();
786            let err = (pred - y).abs();
787            if i < 10 {
788                first_err += err;
789            }
790            if i >= 490 {
791                last_err += err;
792            }
793            model.learn(&sf, y).unwrap();
794        }
795        assert!(last_err < first_err, "error should decrease");
796        let weights = model.weights();
797        assert_eq!(weights.len(), 1);
798        assert!(
799            (weights[0].1 - 2.0).abs() < 0.5,
800            "weight should approach 2.0"
801        );
802    }
803
804    #[test]
805    fn l1_produces_sparse_weights() {
806        // High L1 should drive most weights to zero.
807        let mut model = FtrlRegressor::new(FtrlConfig {
808            alpha: 0.1,
809            beta: 1.0,
810            l1: 100.0,
811            l2: 0.0,
812            max_features: None,
813            new_feature_policy: NewFeaturePolicy::default(),
814        })
815        .unwrap();
816        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
817        for _ in 0..200 {
818            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
819            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
820            let y = 0.5 * x1;
821            let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
822            model.learn(&sf, y).unwrap();
823        }
824        let weights = model.weights();
825        // With very high L1, all weights should be zero.
826        assert!(
827            weights.is_empty(),
828            "weights should all be zero, got {weights:?}"
829        );
830    }
831
832    #[test]
833    fn dynamic_features() {
834        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
835        assert_eq!(model.feature_count(), 0);
836        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
837        model.learn(&sf1, 1.0).unwrap();
838        assert_eq!(model.feature_count(), 1);
839        // A new feature id appears.
840        let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
841        model.learn(&sf2, 2.0).unwrap();
842        assert_eq!(model.feature_count(), 2);
843        // Feature 0 still present.
844        assert!(model.params.contains_key(&0));
845        assert!(model.params.contains_key(&5));
846    }
847
848    #[test]
849    fn predict_does_not_update_state() {
850        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
851        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
852        let _ = model.predict(&sf).unwrap();
853        assert_eq!(model.samples_seen(), 0);
854        assert_eq!(model.feature_count(), 0);
855        // Learn once, then predict again.
856        model.learn(&sf, 1.0).unwrap();
857        let count_after_learn = model.feature_count();
858        let _ = model.predict(&sf).unwrap();
859        assert_eq!(model.feature_count(), count_after_learn);
860        assert_eq!(model.samples_seen(), 1);
861    }
862
863    #[test]
864    fn non_finite_value_rejected() {
865        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
866        // SparseFeatures::from_sorted rejects non-finite values at construction.
867        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
868        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
869        assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
870        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
871        assert!(model.predict(&sf).is_ok());
872    }
873
874    #[test]
875    fn non_finite_target_rejected() {
876        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
877        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
878        assert!(model.learn(&sf, f64::NAN).is_err());
879        assert!(model.learn(&sf, f64::INFINITY).is_err());
880        assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
881        // State should not change on error.
882        assert_eq!(model.samples_seen(), 0);
883    }
884
885    #[test]
886    fn empty_features_rejected() {
887        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
888        let sf = SparseFeatures::new();
889        assert!(model.predict(&sf).is_err());
890        assert!(model.learn(&sf, 1.0).is_err());
891    }
892
893    #[test]
894    fn reset_clears_state() {
895        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
896        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
897        model.learn(&sf, 3.0).unwrap();
898        model.learn(&sf, 3.0).unwrap();
899        assert_eq!(model.samples_seen(), 2);
900        assert_eq!(model.feature_count(), 2);
901        model.reset();
902        assert_eq!(model.samples_seen(), 0);
903        assert_eq!(model.feature_count(), 0);
904        assert!(model.predict(&sf).unwrap().abs() < 1e-12);
905    }
906
907    #[test]
908    fn invalid_config_rejected() {
909        assert!(
910            FtrlRegressor::new(FtrlConfig {
911                alpha: 0.0,
912                ..FtrlConfig::default()
913            })
914            .is_err()
915        );
916        assert!(
917            FtrlRegressor::new(FtrlConfig {
918                alpha: -1.0,
919                ..FtrlConfig::default()
920            })
921            .is_err()
922        );
923        assert!(
924            FtrlRegressor::new(FtrlConfig {
925                beta: -1.0,
926                ..FtrlConfig::default()
927            })
928            .is_err()
929        );
930        assert!(
931            FtrlRegressor::new(FtrlConfig {
932                l1: -1.0,
933                ..FtrlConfig::default()
934            })
935            .is_err()
936        );
937        assert!(
938            FtrlRegressor::new(FtrlConfig {
939                l2: -1.0,
940                ..FtrlConfig::default()
941            })
942            .is_err()
943        );
944        assert!(
945            FtrlRegressor::new(FtrlConfig {
946                alpha: f64::NAN,
947                ..FtrlConfig::default()
948            })
949            .is_err()
950        );
951        assert!(
952            FtrlRegressor::new(FtrlConfig {
953                max_features: Some(0),
954                ..FtrlConfig::default()
955            })
956            .is_err()
957        );
958    }
959
960    #[test]
961    #[cfg(feature = "serde")]
962    fn serde_roundtrip() {
963        let mut model = FtrlRegressor::new(FtrlConfig {
964            alpha: 0.2,
965            beta: 0.5,
966            l1: 0.5,
967            l2: 0.5,
968            max_features: Some(100),
969            new_feature_policy: NewFeaturePolicy::Reject,
970        })
971        .unwrap();
972        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
973        model.learn(&sf, 5.0).unwrap();
974        let json = serde_json::to_string(&model).unwrap();
975        let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
976        assert_eq!(restored.samples_seen(), model.samples_seen());
977        assert_eq!(restored.feature_count(), model.feature_count());
978        let pred_orig = model.predict(&sf).unwrap();
979        let pred_restored = restored.predict(&sf).unwrap();
980        assert!((pred_orig - pred_restored).abs() < 1e-12);
981    }
982
983    #[test]
984    fn weights_returns_nonzero_only() {
985        let mut model = FtrlRegressor::new(FtrlConfig {
986            alpha: 0.5,
987            beta: 1.0,
988            l1: 0.0,
989            l2: 0.0,
990            max_features: None,
991            new_feature_policy: NewFeaturePolicy::default(),
992        })
993        .unwrap();
994        // Learn feature 0 strongly, feature 1 barely.
995        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
996        for _ in 0..50 {
997            model.learn(&sf, 1.0).unwrap();
998        }
999        let weights = model.weights();
1000        // All returned weights should be non-zero.
1001        for &(_, w) in &weights {
1002            assert!(w != 0.0);
1003        }
1004        // Feature 0 should be in the list.
1005        assert!(weights.iter().any(|&(id, _)| id == 0));
1006    }
1007
1008    #[test]
1009    fn multiple_features() {
1010        // y = 1.0 * x0 + (-1.0) * x1 + 0.5
1011        let mut model = FtrlRegressor::new(FtrlConfig {
1012            alpha: 0.5,
1013            beta: 1.0,
1014            l1: 0.0,
1015            l2: 0.0,
1016            max_features: None,
1017            new_feature_policy: NewFeaturePolicy::default(),
1018        })
1019        .unwrap();
1020        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1021        for _ in 0..500 {
1022            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1023            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1024            let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1025            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1026            model.learn(&sf, y).unwrap();
1027        }
1028        let weights = model.weights();
1029        assert_eq!(weights.len(), 2);
1030        let w0 = weights
1031            .iter()
1032            .find(|&&(id, _)| id == 0)
1033            .map(|&(_, w)| w)
1034            .unwrap();
1035        let w1 = weights
1036            .iter()
1037            .find(|&&(id, _)| id == 1)
1038            .map(|&(_, w)| w)
1039            .unwrap();
1040        assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1041        assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1042        assert!(
1043            (model.intercept() - 0.5).abs() < 0.5,
1044            "intercept should approach 0.5"
1045        );
1046    }
1047
1048    #[test]
1049    fn intercept_learned() {
1050        // y = 3.0 (constant), single feature with value 0.0 so that only
1051        // the intercept can learn (feature gradient is always 0).
1052        let mut model = FtrlRegressor::new(FtrlConfig {
1053            alpha: 0.5,
1054            beta: 1.0,
1055            l1: 0.0,
1056            l2: 0.0,
1057            max_features: None,
1058            new_feature_policy: NewFeaturePolicy::default(),
1059        })
1060        .unwrap();
1061        let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1062        for _ in 0..300 {
1063            model.learn(&sf, 3.0).unwrap();
1064        }
1065        let pred = model.predict(&sf).unwrap();
1066        assert!(
1067            (pred - 3.0).abs() < 0.5,
1068            "prediction should approach 3.0, got {pred}"
1069        );
1070        assert!(
1071            (model.intercept() - 3.0).abs() < 0.5,
1072            "intercept should approach 3.0"
1073        );
1074        // Feature weight should be 0 (never updated since x=0).
1075        assert!(model.weights().is_empty());
1076    }
1077
1078    #[test]
1079    fn high_dim_sparse() {
1080        // 1000 possible features, only 5 active per sample.
1081        // Target is a linear combination of the active features.
1082        let mut model = FtrlRegressor::new(FtrlConfig {
1083            alpha: 0.3,
1084            beta: 1.0,
1085            l1: 0.0,
1086            l2: 0.0,
1087            max_features: None,
1088            new_feature_policy: NewFeaturePolicy::default(),
1089        })
1090        .unwrap();
1091        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1092        // True weights for features 0..5.
1093        let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1094        let mut first_err = 0.0;
1095        let mut last_err = 0.0;
1096        for i in 0..2000 {
1097            let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1098            for (j, &w) in true_w.iter().enumerate() {
1099                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1100                active.push((j as u64, x * w));
1101            }
1102            // Add some noise features with zero contribution.
1103            for k in 5..10 {
1104                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1105                active.push((k as u64 + 100, x));
1106            }
1107            active.sort_by_key(|(id, _)| *id);
1108            let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1109            let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1110            let pred = model.predict(&sf).unwrap();
1111            let err = (pred - y).abs();
1112            if i < 20 {
1113                first_err += err;
1114            }
1115            if i >= 1980 {
1116                last_err += err;
1117            }
1118            model.learn(&sf, y).unwrap();
1119        }
1120        assert!(
1121            last_err < first_err,
1122            "error should decrease in high-dim sparse"
1123        );
1124    }
1125
1126    // -----------------------------------------------------------------
1127    // FtrlRegressor: failure atomicity and overflow (ML-001/002)
1128    // -----------------------------------------------------------------
1129
1130    #[test]
1131    fn regressor_overflow_does_not_mutate_state() {
1132        // Finite inputs but intermediate `gradient * value` or
1133        // `gradient^2` overflows. The whole learn call must fail and
1134        // leave the model untouched.
1135        let mut model = FtrlRegressor::new(FtrlConfig {
1136            alpha: 0.1,
1137            beta: 1.0,
1138            l1: 0.0,
1139            l2: 0.0,
1140            max_features: None,
1141            new_feature_policy: NewFeaturePolicy::default(),
1142        })
1143        .unwrap();
1144        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1145        let result = model.learn(&sf, 1e100);
1146        assert!(result.is_err(), "expected overflow error, got {result:?}");
1147        assert_eq!(model.samples_seen(), 0);
1148        assert_eq!(model.feature_count(), 0);
1149        assert!(model.params.is_empty());
1150        assert_eq!(model.intercept.z, 0.0);
1151        assert_eq!(model.intercept.n, 0.0);
1152    }
1153
1154    #[test]
1155    fn regressor_partial_update_is_atomic() {
1156        // Two features in one sample. Feature 0 would succeed on its own,
1157        // feature 1 overflows. Neither may be committed.
1158        let mut model = FtrlRegressor::new(FtrlConfig {
1159            alpha: 0.1,
1160            beta: 1.0,
1161            l1: 0.0,
1162            l2: 0.0,
1163            max_features: None,
1164            new_feature_policy: NewFeaturePolicy::default(),
1165        })
1166        .unwrap();
1167        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1168        assert!(model.learn(&sf, 1e100).is_err());
1169        // Feature 0 must NOT be inserted.
1170        assert!(!model.params.contains_key(&0));
1171        assert!(!model.params.contains_key(&1));
1172        assert_eq!(model.samples_seen(), 0);
1173    }
1174
1175    #[test]
1176    #[cfg(feature = "serde")]
1177    fn regressor_samples_seen_overflow_is_atomic() {
1178        let json = format!(
1179            "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1180            u64::MAX
1181        );
1182        let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1183        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1184        let result = model.learn(&sf, 1.0);
1185        assert!(result.is_err(), "expected counter overflow");
1186        assert_eq!(model.samples_seen(), u64::MAX);
1187        assert_eq!(model.feature_count(), 0);
1188        assert_eq!(model.intercept.z, 0.0);
1189        assert_eq!(model.intercept.n, 0.0);
1190    }
1191
1192    // -----------------------------------------------------------------
1193    // FtrlRegressor: max_features boundary (ML-003)
1194    // -----------------------------------------------------------------
1195
1196    #[test]
1197    fn regressor_max_features_reject_at_limit() {
1198        let mut model = FtrlRegressor::new(FtrlConfig {
1199            alpha: 0.5,
1200            beta: 1.0,
1201            l1: 0.0,
1202            l2: 0.0,
1203            max_features: Some(2),
1204            new_feature_policy: NewFeaturePolicy::Reject,
1205        })
1206        .unwrap();
1207        // Reach exactly the limit.
1208        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1209        model.learn(&sf, 1.0).unwrap();
1210        assert_eq!(model.feature_count(), 2);
1211        // One more new feature: Reject must fail atomically.
1212        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1213        assert!(model.learn(&sf_new, 1.0).is_err());
1214        assert_eq!(model.feature_count(), 2);
1215        assert_eq!(model.samples_seen(), 1);
1216        // Existing features still train.
1217        let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1218        model.learn(&sf_existing, 1.0).unwrap();
1219        assert_eq!(model.feature_count(), 2);
1220        assert_eq!(model.samples_seen(), 2);
1221    }
1222
1223    #[test]
1224    fn regressor_max_features_ignore_skips_new() {
1225        let mut model = FtrlRegressor::new(FtrlConfig {
1226            alpha: 0.5,
1227            beta: 1.0,
1228            l1: 0.0,
1229            l2: 0.0,
1230            max_features: Some(2),
1231            new_feature_policy: NewFeaturePolicy::Ignore,
1232        })
1233        .unwrap();
1234        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1235        model.learn(&sf, 1.0).unwrap();
1236        // Sample with one existing and one new feature. Under Ignore the
1237        // new feature is skipped, the existing one is updated, and the
1238        // counter still advances.
1239        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1240        model.learn(&sf_mixed, 1.0).unwrap();
1241        assert_eq!(model.feature_count(), 2);
1242        assert!(!model.params.contains_key(&2));
1243        assert_eq!(model.samples_seen(), 2);
1244    }
1245
1246    #[test]
1247    fn regressor_max_features_multi_new_prejudge() {
1248        let mut model = FtrlRegressor::new(FtrlConfig {
1249            alpha: 0.5,
1250            beta: 1.0,
1251            l1: 0.0,
1252            l2: 0.0,
1253            max_features: Some(2),
1254            new_feature_policy: NewFeaturePolicy::Reject,
1255        })
1256        .unwrap();
1257        // A single sample with three new features. Reject fails atomically
1258        // without inserting any subset.
1259        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1260        assert!(model.learn(&sf, 1.0).is_err());
1261        assert_eq!(model.feature_count(), 0);
1262        assert_eq!(model.samples_seen(), 0);
1263    }
1264
1265    // -----------------------------------------------------------------
1266    // FtrlRegressor: serde validation (ML-004)
1267    // -----------------------------------------------------------------
1268
1269    #[test]
1270    #[cfg(feature = "serde")]
1271    fn regressor_serde_rejects_negative_n() {
1272        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1273        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1274        assert!(result.is_err(), "negative n must be rejected");
1275    }
1276
1277    #[test]
1278    #[cfg(feature = "serde")]
1279    fn regressor_serde_rejects_invalid_config() {
1280        let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1281        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1282        assert!(result.is_err(), "invalid alpha must be rejected");
1283    }
1284
1285    #[test]
1286    #[cfg(feature = "serde")]
1287    fn regressor_serde_accepts_missing_optional_fields() {
1288        // Old state without max_features/new_feature_policy must still load.
1289        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1290        let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1291        assert!(model.config().max_features.is_none());
1292        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1293    }
1294
1295    // -----------------------------------------------------------------
1296    // FtrlClassifier tests
1297    // -----------------------------------------------------------------
1298
1299    #[test]
1300    fn cold_start_returns_0_5() {
1301        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1302        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1303        let p = model.predict_proba(&sf).unwrap();
1304        assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1305    }
1306
1307    #[test]
1308    fn learn_separable_data() {
1309        // Linearly separable: class 1 when x0 > 0.
1310        let mut model = FtrlClassifier::new(FtrlConfig {
1311            alpha: 0.5,
1312            beta: 1.0,
1313            l1: 0.0,
1314            l2: 0.0,
1315            max_features: None,
1316            new_feature_policy: NewFeaturePolicy::default(),
1317        })
1318        .unwrap();
1319        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1320        for _ in 0..1000 {
1321            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1322            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1323            let y = x0 > 0.0;
1324            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1325            model.learn(&sf, y).unwrap();
1326        }
1327        let p_pos = model
1328            .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1329            .unwrap();
1330        let p_neg = model
1331            .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1332            .unwrap();
1333        assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1334        assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1335    }
1336
1337    #[test]
1338    fn classifier_l1_produces_sparse_weights() {
1339        let mut model = FtrlClassifier::new(FtrlConfig {
1340            alpha: 0.1,
1341            beta: 1.0,
1342            l1: 100.0,
1343            l2: 0.0,
1344            max_features: None,
1345            new_feature_policy: NewFeaturePolicy::default(),
1346        })
1347        .unwrap();
1348        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1349        for _ in 0..200 {
1350            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1351            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1352            let y = x0 > 0.0;
1353            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1354            model.learn(&sf, y).unwrap();
1355        }
1356        let weights = model.weights();
1357        assert!(
1358            weights.is_empty(),
1359            "weights should all be zero with high L1, got {weights:?}"
1360        );
1361    }
1362
1363    #[test]
1364    fn classifier_dynamic_features() {
1365        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1366        assert_eq!(model.feature_count(), 0);
1367        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1368        model.learn(&sf1, true).unwrap();
1369        assert_eq!(model.feature_count(), 1);
1370        let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1371        model.learn(&sf2, false).unwrap();
1372        assert_eq!(model.feature_count(), 2);
1373    }
1374
1375    #[test]
1376    fn classifier_predict_does_not_update_state() {
1377        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1378        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1379        let _ = model.predict_proba(&sf).unwrap();
1380        assert_eq!(model.samples_seen(), 0);
1381        assert_eq!(model.feature_count(), 0);
1382        model.learn(&sf, true).unwrap();
1383        let count = model.feature_count();
1384        let _ = model.predict_proba(&sf).unwrap();
1385        assert_eq!(model.feature_count(), count);
1386        assert_eq!(model.samples_seen(), 1);
1387    }
1388
1389    #[test]
1390    fn classifier_non_finite_value_rejected() {
1391        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1392        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1393        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1394        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1395        assert!(model.predict_proba(&sf).is_ok());
1396    }
1397
1398    #[test]
1399    fn classifier_empty_features_rejected() {
1400        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1401        let sf = SparseFeatures::new();
1402        assert!(model.predict_proba(&sf).is_err());
1403        assert!(model.learn(&sf, true).is_err());
1404    }
1405
1406    #[test]
1407    fn classifier_reset_clears_state() {
1408        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1409        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1410        model.learn(&sf, true).unwrap();
1411        model.learn(&sf, false).unwrap();
1412        assert_eq!(model.samples_seen(), 2);
1413        assert!(model.feature_count() > 0);
1414        model.reset();
1415        assert_eq!(model.samples_seen(), 0);
1416        assert_eq!(model.feature_count(), 0);
1417        let p = model.predict_proba(&sf).unwrap();
1418        assert!((p - 0.5).abs() < 1e-12);
1419    }
1420
1421    #[test]
1422    fn classifier_invalid_config_rejected() {
1423        assert!(
1424            FtrlClassifier::new(FtrlConfig {
1425                alpha: 0.0,
1426                ..FtrlConfig::default()
1427            })
1428            .is_err()
1429        );
1430        assert!(
1431            FtrlClassifier::new(FtrlConfig {
1432                beta: -0.1,
1433                ..FtrlConfig::default()
1434            })
1435            .is_err()
1436        );
1437        assert!(
1438            FtrlClassifier::new(FtrlConfig {
1439                l1: -1.0,
1440                ..FtrlConfig::default()
1441            })
1442            .is_err()
1443        );
1444        assert!(
1445            FtrlClassifier::new(FtrlConfig {
1446                l2: -1.0,
1447                ..FtrlConfig::default()
1448            })
1449            .is_err()
1450        );
1451        assert!(
1452            FtrlClassifier::new(FtrlConfig {
1453                alpha: f64::INFINITY,
1454                ..FtrlConfig::default()
1455            })
1456            .is_err()
1457        );
1458    }
1459
1460    #[test]
1461    #[cfg(feature = "serde")]
1462    fn classifier_serde_roundtrip() {
1463        let mut model = FtrlClassifier::new(FtrlConfig {
1464            alpha: 0.3,
1465            beta: 0.5,
1466            l1: 0.1,
1467            l2: 0.2,
1468            max_features: Some(100),
1469            new_feature_policy: NewFeaturePolicy::Reject,
1470        })
1471        .unwrap();
1472        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1473        model.learn(&sf, true).unwrap();
1474        model.learn(&sf, false).unwrap();
1475        let json = serde_json::to_string(&model).unwrap();
1476        let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1477        assert_eq!(restored.samples_seen(), model.samples_seen());
1478        assert_eq!(restored.feature_count(), model.feature_count());
1479        let p1 = model.predict_proba(&sf).unwrap();
1480        let p2 = restored.predict_proba(&sf).unwrap();
1481        assert!((p1 - p2).abs() < 1e-12);
1482    }
1483
1484    #[test]
1485    fn predict_proba_in_range() {
1486        let mut model = FtrlClassifier::new(FtrlConfig {
1487            alpha: 0.5,
1488            beta: 1.0,
1489            l1: 0.0,
1490            l2: 0.0,
1491            max_features: None,
1492            new_feature_policy: NewFeaturePolicy::default(),
1493        })
1494        .unwrap();
1495        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1496        for _ in 0..200 {
1497            let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1498            let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1499            let y = x0 > 0.0;
1500            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1501            model.learn(&sf, y).unwrap();
1502            let p = model.predict_proba(&sf).unwrap();
1503            assert!(
1504                (0.0..=1.0).contains(&p),
1505                "probability must be in [0,1], got {p}"
1506            );
1507        }
1508    }
1509
1510    #[test]
1511    fn learn_improves_accuracy() {
1512        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1513        // Generate a fixed test set.
1514        let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1515            .map(|_| {
1516                let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1517                let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1518                let y = x0 + x1 > 0.0;
1519                (
1520                    SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1521                    y,
1522                )
1523            })
1524            .collect();
1525
1526        let mut model = FtrlClassifier::new(FtrlConfig {
1527            alpha: 0.5,
1528            beta: 1.0,
1529            l1: 0.0,
1530            l2: 0.0,
1531            max_features: None,
1532            new_feature_policy: NewFeaturePolicy::default(),
1533        })
1534        .unwrap();
1535
1536        // Accuracy before learning (always predicts 0.5 -> threshold 0.5 -> true).
1537        let acc_before: f64 = test_set
1538            .iter()
1539            .map(|(sf, y)| {
1540                let pred = model.predict(sf).unwrap();
1541                if pred == *y { 1.0 } else { 0.0 }
1542            })
1543            .sum::<f64>()
1544            / test_set.len() as f64;
1545
1546        // Train on fresh data.
1547        for _ in 0..1000 {
1548            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1549            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1550            let y = x0 + x1 > 0.0;
1551            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1552            model.learn(&sf, y).unwrap();
1553        }
1554
1555        let acc_after: f64 = test_set
1556            .iter()
1557            .map(|(sf, y)| {
1558                let pred = model.predict(sf).unwrap();
1559                if pred == *y { 1.0 } else { 0.0 }
1560            })
1561            .sum::<f64>()
1562            / test_set.len() as f64;
1563
1564        assert!(
1565            acc_after > acc_before,
1566            "accuracy should improve: {acc_before} -> {acc_after}"
1567        );
1568    }
1569
1570    #[test]
1571    fn classifier_weights_returns_nonzero_only() {
1572        let mut model = FtrlClassifier::new(FtrlConfig {
1573            alpha: 0.5,
1574            beta: 1.0,
1575            l1: 0.0,
1576            l2: 0.0,
1577            max_features: None,
1578            new_feature_policy: NewFeaturePolicy::default(),
1579        })
1580        .unwrap();
1581        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1582        for _ in 0..50 {
1583            model.learn(&sf, true).unwrap();
1584        }
1585        let weights = model.weights();
1586        for &(_, w) in &weights {
1587            assert!(w != 0.0);
1588        }
1589    }
1590
1591    #[test]
1592    fn classifier_multiple_features() {
1593        let mut model = FtrlClassifier::new(FtrlConfig {
1594            alpha: 0.5,
1595            beta: 1.0,
1596            l1: 0.0,
1597            l2: 0.0,
1598            max_features: None,
1599            new_feature_policy: NewFeaturePolicy::default(),
1600        })
1601        .unwrap();
1602        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1603        for _ in 0..1000 {
1604            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1605            let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1606            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1607            // y = 1 if x0 + x1 > 0
1608            let y = x0 + x1 > 0.0;
1609            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1610            model.learn(&sf, y).unwrap();
1611        }
1612        let weights = model.weights();
1613        // Features 0 and 1 should have non-zero weights; feature 2 may or may not.
1614        assert!(weights.iter().any(|&(id, _)| id == 0));
1615        assert!(weights.iter().any(|&(id, _)| id == 1));
1616        // Verify prediction quality.
1617        let p_pos = model
1618            .predict_proba(
1619                &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1620            )
1621            .unwrap();
1622        let p_neg = model
1623            .predict_proba(
1624                &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1625            )
1626            .unwrap();
1627        assert!(p_pos > 0.8);
1628        assert!(p_neg < 0.2);
1629    }
1630
1631    #[test]
1632    fn log_loss_converges() {
1633        // Average log loss should decrease over training.
1634        let mut model = FtrlClassifier::new(FtrlConfig {
1635            alpha: 0.5,
1636            beta: 1.0,
1637            l1: 0.0,
1638            l2: 0.0,
1639            max_features: None,
1640            new_feature_policy: NewFeaturePolicy::default(),
1641        })
1642        .unwrap();
1643        let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1644        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1645        let mut first_loss = 0.0;
1646        let mut last_loss = 0.0;
1647        for i in 0..1000 {
1648            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1649            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1650            let y = x0 > 0.0;
1651            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1652            let p = model.predict_proba(&sf).unwrap();
1653            let loss = loss_fn.loss(p, y);
1654            if i < 20 {
1655                first_loss += loss;
1656            }
1657            if i >= 980 {
1658                last_loss += loss;
1659            }
1660            model.learn(&sf, y).unwrap();
1661        }
1662        assert!(last_loss < first_loss, "log loss should decrease");
1663    }
1664
1665    // -----------------------------------------------------------------
1666    // FtrlClassifier: failure atomicity and overflow (ML-001/002)
1667    // -----------------------------------------------------------------
1668
1669    #[test]
1670    fn classifier_overflow_does_not_mutate_state() {
1671        let mut model = FtrlClassifier::new(FtrlConfig {
1672            alpha: 0.1,
1673            beta: 1.0,
1674            l1: 0.0,
1675            l2: 0.0,
1676            max_features: None,
1677            new_feature_policy: NewFeaturePolicy::default(),
1678        })
1679        .unwrap();
1680        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1681        // Cold-start probability is 0.5, grad = 0.5 - 0.0 = 0.5.
1682        // g for feature 1 = 0.5 * 1e300 = 5e299 (finite), g^2 = 2.5e599 = inf.
1683        let result = model.learn(&sf, false);
1684        assert!(result.is_err(), "expected overflow error, got {result:?}");
1685        assert_eq!(model.samples_seen(), 0);
1686        assert_eq!(model.feature_count(), 0);
1687        assert!(model.params.is_empty());
1688        assert_eq!(model.intercept.z, 0.0);
1689        assert_eq!(model.intercept.n, 0.0);
1690    }
1691
1692    #[test]
1693    fn classifier_partial_update_is_atomic() {
1694        let mut model = FtrlClassifier::new(FtrlConfig {
1695            alpha: 0.1,
1696            beta: 1.0,
1697            l1: 0.0,
1698            l2: 0.0,
1699            max_features: None,
1700            new_feature_policy: NewFeaturePolicy::default(),
1701        })
1702        .unwrap();
1703        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1704        assert!(model.learn(&sf, false).is_err());
1705        assert!(!model.params.contains_key(&0));
1706        assert!(!model.params.contains_key(&1));
1707        assert_eq!(model.samples_seen(), 0);
1708    }
1709
1710    #[test]
1711    #[cfg(feature = "serde")]
1712    fn classifier_samples_seen_overflow_is_atomic() {
1713        let json = format!(
1714            "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1715            u64::MAX
1716        );
1717        let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1718        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1719        let result = model.learn(&sf, true);
1720        assert!(result.is_err(), "expected counter overflow");
1721        assert_eq!(model.samples_seen(), u64::MAX);
1722        assert_eq!(model.feature_count(), 0);
1723        assert_eq!(model.intercept.z, 0.0);
1724        assert_eq!(model.intercept.n, 0.0);
1725    }
1726
1727    // -----------------------------------------------------------------
1728    // FtrlClassifier: max_features boundary (ML-003)
1729    // -----------------------------------------------------------------
1730
1731    #[test]
1732    fn classifier_max_features_reject_at_limit() {
1733        let mut model = FtrlClassifier::new(FtrlConfig {
1734            alpha: 0.5,
1735            beta: 1.0,
1736            l1: 0.0,
1737            l2: 0.0,
1738            max_features: Some(2),
1739            new_feature_policy: NewFeaturePolicy::Reject,
1740        })
1741        .unwrap();
1742        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1743        model.learn(&sf, true).unwrap();
1744        assert_eq!(model.feature_count(), 2);
1745        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1746        assert!(model.learn(&sf_new, true).is_err());
1747        assert_eq!(model.feature_count(), 2);
1748        assert_eq!(model.samples_seen(), 1);
1749    }
1750
1751    #[test]
1752    fn classifier_max_features_ignore_skips_new() {
1753        let mut model = FtrlClassifier::new(FtrlConfig {
1754            alpha: 0.5,
1755            beta: 1.0,
1756            l1: 0.0,
1757            l2: 0.0,
1758            max_features: Some(2),
1759            new_feature_policy: NewFeaturePolicy::Ignore,
1760        })
1761        .unwrap();
1762        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1763        model.learn(&sf, true).unwrap();
1764        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1765        model.learn(&sf_mixed, false).unwrap();
1766        assert_eq!(model.feature_count(), 2);
1767        assert!(!model.params.contains_key(&2));
1768        assert_eq!(model.samples_seen(), 2);
1769    }
1770
1771    #[test]
1772    fn classifier_max_features_multi_new_prejudge() {
1773        let mut model = FtrlClassifier::new(FtrlConfig {
1774            alpha: 0.5,
1775            beta: 1.0,
1776            l1: 0.0,
1777            l2: 0.0,
1778            max_features: Some(2),
1779            new_feature_policy: NewFeaturePolicy::Reject,
1780        })
1781        .unwrap();
1782        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1783        assert!(model.learn(&sf, true).is_err());
1784        assert_eq!(model.feature_count(), 0);
1785        assert_eq!(model.samples_seen(), 0);
1786    }
1787
1788    // -----------------------------------------------------------------
1789    // FtrlClassifier: serde validation (ML-004)
1790    // -----------------------------------------------------------------
1791
1792    #[test]
1793    #[cfg(feature = "serde")]
1794    fn classifier_serde_rejects_negative_n() {
1795        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1796        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1797        assert!(result.is_err(), "negative n must be rejected");
1798    }
1799
1800    #[test]
1801    #[cfg(feature = "serde")]
1802    fn classifier_serde_rejects_invalid_config() {
1803        let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1804        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1805        assert!(result.is_err(), "invalid alpha must be rejected");
1806    }
1807
1808    #[test]
1809    #[cfg(feature = "serde")]
1810    fn classifier_serde_accepts_missing_optional_fields() {
1811        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1812        let model: FtrlClassifier = serde_json::from_str(json).unwrap();
1813        assert!(model.config().max_features.is_none());
1814        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1815    }
1816}