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 FTRL feature weight with explicit config-aware safety
229    /// checks.
230    ///
231    /// Unlike [`weight`](Self::weight), this method verifies every
232    /// intermediate quantity (numerator, denominator, quotient) and rejects
233    /// states where the config + param combination would produce a
234    /// non-finite or non-computable weight. It is intended for the serde
235    /// trust boundary and may also be used by the predict path.
236    ///
237    /// # Contract
238    ///
239    /// - `z` finite, `n` finite and `>= 0`.
240    /// - L1 soft-thresholding path (`|z| <= l1`) returns `Ok(0.0)`.
241    /// - Denominator must be finite and non-zero; a zero denominator
242    ///   (e.g. `l2 = 0`, `beta = 0`, `sqrt(n) / alpha` underflows to 0)
243    ///   is rejected explicitly rather than discovered via the final
244    ///   quotient.
245    /// - The final quotient must be finite.
246    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
247    fn weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
248        ensure_finite("ftrl_z", self.z)?;
249        ensure_finite("ftrl_n", self.n)?;
250        if self.n < 0.0 {
251            return Err(RillError::InvalidState(format!(
252                "ftrl n must be non-negative, got {}",
253                self.n
254            )));
255        }
256        // L1 soft-thresholding: |z| <= l1 → weight is 0.
257        if self.z.abs() <= config.l1 {
258            return Ok(0.0);
259        }
260        let sign = self.z.signum();
261        let numerator = -(self.z - sign * config.l1);
262        ensure_finite("ftrl_weight_numerator", numerator)?;
263        let sqrt_n = self.n.sqrt();
264        ensure_finite("ftrl_weight_sqrt_n", sqrt_n)?;
265        let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
266        ensure_finite("ftrl_weight_denominator", denominator)?;
267        if denominator == 0.0 {
268            return Err(RillError::InvalidState(format!(
269                "ftrl weight denominator is zero (z={}, n={}, alpha={}, beta={}, l1={}, l2={})",
270                self.z, self.n, config.alpha, config.beta, config.l1, config.l2
271            )));
272        }
273        let weight = numerator / denominator;
274        ensure_finite("ftrl_weight", weight)?;
275        Ok(weight)
276    }
277
278    /// Compute the intercept weight with explicit config-aware safety
279    /// checks.
280    ///
281    /// See [`weight_checked`](Self::weight_checked) for the rationale. The
282    /// intercept uses `l1 = 0` (no L1 regularization) and short-circuits
283    /// the cold-start case `n == 0, z == 0` to `Ok(0.0)`. A state with
284    /// `n == 0` but `z != 0` is rejected because it cannot produce a
285    /// finite intercept weight.
286    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
287    fn intercept_weight_checked(&self, config: &FtrlConfig) -> Result<f64, RillError> {
288        ensure_finite("ftrl_z", self.z)?;
289        ensure_finite("ftrl_n", self.n)?;
290        if self.n < 0.0 {
291            return Err(RillError::InvalidState(format!(
292                "ftrl n must be non-negative, got {}",
293                self.n
294            )));
295        }
296        // Cold start: n == 0. intercept_weight short-circuits to 0 only
297        // when z == 0 as well; a non-zero z with n == 0 indicates a
298        // corrupted or malicious state.
299        if self.n == 0.0 {
300            if self.z != 0.0 {
301                return Err(RillError::InvalidState(format!(
302                    "ftrl intercept has n=0 but z={} (non-zero); cannot produce a finite weight",
303                    self.z
304                )));
305            }
306            return Ok(0.0);
307        }
308        let numerator = -self.z;
309        ensure_finite("ftrl_intercept_numerator", numerator)?;
310        let sqrt_n = self.n.sqrt();
311        ensure_finite("ftrl_intercept_sqrt_n", sqrt_n)?;
312        let denominator = config.l2 + (config.beta + sqrt_n) / config.alpha;
313        ensure_finite("ftrl_intercept_denominator", denominator)?;
314        if denominator == 0.0 {
315            return Err(RillError::InvalidState(format!(
316                "ftrl intercept denominator is zero (z={}, n={}, alpha={}, beta={}, l2={})",
317                self.z, self.n, config.alpha, config.beta, config.l2
318            )));
319        }
320        let weight = numerator / denominator;
321        ensure_finite("ftrl_intercept_weight", weight)?;
322        Ok(weight)
323    }
324
325    /// Compute the next `(z, n)` state after applying a gradient, without
326    /// mutating `self`.
327    ///
328    /// `sigma = (sqrt(n_new) - sqrt(n_old)) / alpha`
329    /// `z_new = z + g - sigma * w`
330    /// `n_new = n + g^2`
331    ///
332    /// Every intermediate result is checked for finiteness so a partial
333    /// overflow cannot poison state. The caller is responsible for
334    /// committing the returned values atomically.
335    fn next_updated(
336        &self,
337        gradient: f64,
338        weight: f64,
339        config: &FtrlConfig,
340    ) -> Result<(f64, f64), RillError> {
341        let gradient_sq = gradient * gradient;
342        ensure_finite("ftrl_gradient_squared", gradient_sq)?;
343        // Detect underflow: a non-zero gradient whose square underflows to
344        // zero. In that case `n_new == n_old` while `z` still advances by
345        // `gradient`, which can produce an unusable state (`n == 0, z != 0`)
346        // on cold-start features and a zero denominator in `weight` on the
347        // next prediction. Reject explicitly rather than silently committing
348        // a state that would make the next `predict()` fail.
349        if gradient != 0.0 && gradient_sq == 0.0 {
350            return Err(RillError::NonFiniteValue {
351                field: "ftrl_gradient_squared",
352                value: gradient_sq,
353            });
354        }
355        let n_new = checked_finite_add(self.n, gradient_sq, "ftrl_n_new")?;
356        let sigma = (n_new.sqrt() - self.n.sqrt()) / config.alpha;
357        ensure_finite("ftrl_sigma", sigma)?;
358        let sigma_w = sigma * weight;
359        ensure_finite("ftrl_sigma_weight", sigma_w)?;
360        let z_delta = gradient - sigma_w;
361        ensure_finite("ftrl_z_delta", z_delta)?;
362        let z_new = checked_finite_add(self.z, z_delta, "ftrl_z_new")?;
363        Ok((z_new, n_new))
364    }
365
366    /// Validate that `z` is finite, `n` is finite and non-negative, and the
367    /// state can produce a finite weight on the next `predict()` call.
368    ///
369    /// The FTRL weight formula divides by `l2 + (beta + sqrt(n)) / alpha`.
370    /// When `n == 0` and `l2 == 0` and `beta == 0` the denominator is zero.
371    /// `intercept_weight` already short-circuits `n == 0` to `0.0`, but
372    /// `weight` does not. A state with `n == 0` and `z != 0` (where
373    /// `|z| > l1`) would therefore produce `±inf` on the next prediction.
374    /// Such a state can only arise from a non-zero gradient whose square
375    /// underflows to zero (handled in [`FtrlParam::next_updated`]) or from
376    /// a maliciously crafted serde payload; both must be rejected.
377    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
378    fn validate(&self) -> Result<(), RillError> {
379        ensure_finite("ftrl_z", self.z)?;
380        ensure_finite("ftrl_n", self.n)?;
381        if self.n < 0.0 {
382            return Err(RillError::InvalidState(format!(
383                "ftrl n must be non-negative, got {0}",
384                self.n
385            )));
386        }
387        if self.n == 0.0 && self.z != 0.0 {
388            return Err(RillError::InvalidState(format!(
389                "ftrl param has n=0 but z={0} (non-zero); this state cannot \
390                 produce a finite weight",
391                self.z
392            )));
393        }
394        Ok(())
395    }
396}
397
398#[cfg(feature = "serde")]
399impl<'de> serde::Deserialize<'de> for FtrlParam {
400    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
401    where
402        D: serde::Deserializer<'de>,
403    {
404        #[derive(serde::Deserialize)]
405        struct FtrlParamState {
406            z: f64,
407            n: f64,
408        }
409
410        let state = FtrlParamState::deserialize(deserializer)?;
411        let param = FtrlParam {
412            z: state.z,
413            n: state.n,
414        };
415        param.validate().map_err(serde::de::Error::custom)?;
416        Ok(param)
417    }
418}
419
420/// Compute the dot product `w · x` over sparse features.
421///
422/// Iterates only over the features present in `features` (not all stored
423/// params), looking up each feature's current FTRL weight. Each
424/// contribution and the running sum are checked for finiteness so a
425/// single overflowing term cannot poison the prediction.
426fn compute_dot(
427    params: &BTreeMap<FeatureId, FtrlParam>,
428    config: &FtrlConfig,
429    features: &SparseFeatures,
430) -> Result<f64, RillError> {
431    if features.is_empty() {
432        return Err(RillError::EmptyFeatures);
433    }
434    let mut dot = 0.0;
435    for &(id, value) in features.values() {
436        ensure_finite("sparse_value", value)?;
437        if let Some(param) = params.get(&id) {
438            let w = param.weight(config);
439            ensure_finite("ftrl_weight", w)?;
440            let contribution = w * value;
441            ensure_finite("ftrl_dot_contribution", contribution)?;
442            dot = checked_finite_add(dot, contribution, "ftrl_dot")?;
443        }
444    }
445    Ok(dot)
446}
447
448/// FTRL regressor with squared loss.
449///
450/// Learns `y ≈ w · x + b` incrementally. The gradient of the squared loss
451/// w.r.t. the prediction is `prediction - target`, so each feature's
452/// gradient is `(prediction - target) * x_i`.
453///
454/// # Examples
455///
456/// ```
457/// use rill_ml::models::{FtrlConfig, FtrlRegressor};
458/// use rill_ml::sparse::SparseFeatures;
459/// use rill_ml::SparseRegressor;
460///
461/// let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
462/// let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
463/// let _pred = model.predict(&sf).unwrap();
464/// model.learn(&sf, 3.0).unwrap();
465/// ```
466#[derive(Debug, Clone)]
467#[cfg_attr(feature = "serde", derive(serde::Serialize))]
468pub struct FtrlRegressor {
469    config: FtrlConfig,
470    params: BTreeMap<FeatureId, FtrlParam>,
471    intercept: FtrlParam,
472    samples_seen: u64,
473}
474
475impl FtrlRegressor {
476    /// Create a new FTRL regressor.
477    ///
478    /// Returns an error if the configuration is invalid.
479    pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
480        config.validate()?;
481        Ok(Self {
482            config,
483            params: BTreeMap::new(),
484            intercept: FtrlParam::default(),
485            samples_seen: 0,
486        })
487    }
488
489    /// The model configuration.
490    pub const fn config(&self) -> &FtrlConfig {
491        &self.config
492    }
493
494    /// Return the current non-zero feature weights, sorted by `FeatureId`.
495    ///
496    /// Features whose FTRL weight is exactly zero (due to L1
497    /// soft-thresholding or never having been updated) are excluded.
498    pub fn weights(&self) -> Vec<(FeatureId, f64)> {
499        self.params
500            .iter()
501            .map(|(&id, param)| (id, param.weight(&self.config)))
502            .filter(|&(_, w)| w != 0.0)
503            .collect()
504    }
505
506    /// Compute the current intercept (bias) weight.
507    pub fn intercept(&self) -> f64 {
508        self.intercept.intercept_weight(&self.config)
509    }
510
511    /// Number of distinct features the model has seen.
512    pub fn feature_count(&self) -> usize {
513        self.params.len()
514    }
515
516    /// Compute the raw prediction `w · x + b` without updating state.
517    fn predict_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
518        let dot = compute_dot(&self.params, &self.config, features)?;
519        let intercept = self.intercept.intercept_weight(&self.config);
520        ensure_finite("ftrl_intercept", intercept)?;
521        checked_finite_add(dot, intercept, "ftrl_prediction")
522    }
523}
524
525#[cfg(feature = "serde")]
526impl<'de> serde::Deserialize<'de> for FtrlRegressor {
527    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
528    where
529        D: serde::Deserializer<'de>,
530    {
531        #[derive(serde::Deserialize)]
532        struct FtrlRegressorState {
533            config: FtrlConfig,
534            params: BTreeMap<FeatureId, FtrlParam>,
535            intercept: FtrlParam,
536            samples_seen: u64,
537        }
538
539        let state = FtrlRegressorState::deserialize(deserializer)?;
540        let model = FtrlRegressor {
541            config: state.config,
542            params: state.params,
543            intercept: state.intercept,
544            samples_seen: state.samples_seen,
545        };
546        // config and each FtrlParam are already validated by their own
547        // Deserialize impls; only the top-level invariants remain.
548        model
549            .validate_invariants()
550            .map_err(serde::de::Error::custom)?;
551        Ok(model)
552    }
553}
554
555impl FtrlRegressor {
556    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
557    fn validate_invariants(&self) -> Result<(), RillError> {
558        // config and individual params are validated at deserialization;
559        // here we additionally verify that the config + param combination
560        // produces a finite, computable weight for every feature and the
561        // intercept. This catches states that pass the basic z/n checks
562        // but would still yield a zero-denominator or infinite weight
563        // (e.g. very large alpha with tiny n, beta=0, l2=0).
564        self.config.validate()?;
565        for (id, param) in &self.params {
566            param.validate()?;
567            param
568                .weight_checked(&self.config)
569                .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
570        }
571        self.intercept.validate()?;
572        self.intercept
573            .intercept_weight_checked(&self.config)
574            .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
575        Ok(())
576    }
577}
578
579impl SparseRegressor for FtrlRegressor {
580    fn samples_seen(&self) -> u64 {
581        self.samples_seen
582    }
583
584    fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
585        self.predict_inner(features)
586    }
587
588    fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
589        if features.is_empty() {
590            return Err(RillError::EmptyFeatures);
591        }
592        ensure_finite("target", target)?;
593
594        // Reserve the sample counter first. If it overflows, no state changes.
595        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
596
597        let prediction = self.predict_inner(features)?;
598        ensure_finite("ftrl_prediction", prediction)?;
599        let grad = prediction - target;
600        ensure_finite("ftrl_gradient", grad)?;
601
602        // Pre-judge the max_features policy. We never insert a subset of
603        // new features then fail: either all new features are accepted or
604        // (under Ignore) all new features are skipped.
605        let new_ids_count = features
606            .values()
607            .iter()
608            .filter(|(id, _)| !self.params.contains_key(id))
609            .count();
610        let mut skip_new_features = false;
611        if let Some(max_features) = self.config.max_features {
612            let projected = self.params.len().saturating_add(new_ids_count);
613            if projected > max_features {
614                match self.config.new_feature_policy {
615                    NewFeaturePolicy::Reject => {
616                        return Err(RillError::InvalidState(format!(
617                            "FTRL feature count {projected} exceeds max_features {max_features}"
618                        )));
619                    }
620                    NewFeaturePolicy::Ignore => {
621                        skip_new_features = true;
622                    }
623                }
624            }
625        }
626
627        // Compute the next (z, n) for every feature without touching self.
628        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
629        for &(id, value) in features.values() {
630            // Ignore policy must be evaluated before any arithmetic on the
631            // new feature's value. Otherwise an oversized new feature whose
632            // `grad * value` would overflow could fail the whole `learn()`
633            // call even though the feature is supposed to be skipped.
634            let is_new = !self.params.contains_key(&id);
635            if is_new && skip_new_features {
636                continue;
637            }
638
639            let g = grad * value;
640            ensure_finite("ftrl_feature_gradient", g)?;
641
642            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
643                let w = param.weight(&self.config);
644                param.next_updated(g, w, &self.config)?
645            } else {
646                let param = FtrlParam::default();
647                let w = param.weight(&self.config);
648                param.next_updated(g, w, &self.config)?
649            };
650            // Verify the next state produces a finite weight before
651            // committing. This catches any path that would leave the model
652            // in a state where the next `predict()` fails due to internal
653            // state (e.g. a zero denominator in the weight formula).
654            let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
655            ensure_finite("ftrl_next_weight", next_w)?;
656            updates.push((id, new_z, new_n));
657        }
658
659        // Compute the next intercept state.
660        let w_b = self.intercept.intercept_weight(&self.config);
661        let (new_intercept_z, new_intercept_n) =
662            self.intercept.next_updated(grad, w_b, &self.config)?;
663        // Verify the next intercept produces a finite weight too.
664        let next_intercept_w = FtrlParam {
665            z: new_intercept_z,
666            n: new_intercept_n,
667        }
668        .intercept_weight(&self.config);
669        ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
670
671        // Commit atomically. No failure path beyond this point.
672        for (id, new_z, new_n) in updates {
673            let param = self.params.entry(id).or_default();
674            param.z = new_z;
675            param.n = new_n;
676        }
677        self.intercept.z = new_intercept_z;
678        self.intercept.n = new_intercept_n;
679        self.samples_seen = next_samples_seen;
680
681        Ok(())
682    }
683
684    fn reset(&mut self) {
685        self.params.clear();
686        self.intercept = FtrlParam::default();
687        self.samples_seen = 0;
688    }
689}
690
691/// FTRL binary classifier with log loss.
692///
693/// Predicts `P(y=1 | x) = sigmoid(w · x + b)`. The gradient of the log loss
694/// w.r.t. the logit simplifies to `probability - target`, so each feature's
695/// gradient is `(probability - target) * x_i`.
696///
697/// The returned probability lies in `[0, 1]`. Extreme logits can produce
698/// exactly `0.0` or `1.0` after `sigmoid`; downstream consumers such as
699/// [`crate::loss::log_loss::BinaryLogLoss`] clip internally.
700///
701/// # Examples
702///
703/// ```
704/// use rill_ml::models::{FtrlClassifier, FtrlConfig};
705/// use rill_ml::sparse::SparseFeatures;
706/// use rill_ml::SparseClassifier;
707///
708/// let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
709/// let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
710/// let _proba = model.predict_proba(&sf).unwrap();
711/// model.learn(&sf, true).unwrap();
712/// ```
713#[derive(Debug, Clone)]
714#[cfg_attr(feature = "serde", derive(serde::Serialize))]
715pub struct FtrlClassifier {
716    config: FtrlConfig,
717    params: BTreeMap<FeatureId, FtrlParam>,
718    intercept: FtrlParam,
719    samples_seen: u64,
720}
721
722impl FtrlClassifier {
723    /// Create a new FTRL classifier.
724    ///
725    /// Returns an error if the configuration is invalid.
726    pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
727        config.validate()?;
728        Ok(Self {
729            config,
730            params: BTreeMap::new(),
731            intercept: FtrlParam::default(),
732            samples_seen: 0,
733        })
734    }
735
736    /// The model configuration.
737    pub const fn config(&self) -> &FtrlConfig {
738        &self.config
739    }
740
741    /// Return the current non-zero feature weights, sorted by `FeatureId`.
742    ///
743    /// Features whose FTRL weight is exactly zero (due to L1
744    /// soft-thresholding or never having been updated) are excluded.
745    pub fn weights(&self) -> Vec<(FeatureId, f64)> {
746        self.params
747            .iter()
748            .map(|(&id, param)| (id, param.weight(&self.config)))
749            .filter(|&(_, w)| w != 0.0)
750            .collect()
751    }
752
753    /// Compute the current intercept (bias) weight.
754    pub fn intercept(&self) -> f64 {
755        self.intercept.intercept_weight(&self.config)
756    }
757
758    /// Number of distinct features the model has seen.
759    pub fn feature_count(&self) -> usize {
760        self.params.len()
761    }
762
763    /// Compute the probability `sigmoid(w · x + b)` without updating state.
764    fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
765        let dot = compute_dot(&self.params, &self.config, features)?;
766        let intercept = self.intercept.intercept_weight(&self.config);
767        ensure_finite("ftrl_intercept", intercept)?;
768        let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
769        Ok(sigmoid(logit))
770    }
771}
772
773#[cfg(feature = "serde")]
774impl<'de> serde::Deserialize<'de> for FtrlClassifier {
775    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
776    where
777        D: serde::Deserializer<'de>,
778    {
779        #[derive(serde::Deserialize)]
780        struct FtrlClassifierState {
781            config: FtrlConfig,
782            params: BTreeMap<FeatureId, FtrlParam>,
783            intercept: FtrlParam,
784            samples_seen: u64,
785        }
786
787        let state = FtrlClassifierState::deserialize(deserializer)?;
788        let model = FtrlClassifier {
789            config: state.config,
790            params: state.params,
791            intercept: state.intercept,
792            samples_seen: state.samples_seen,
793        };
794        model
795            .validate_invariants()
796            .map_err(serde::de::Error::custom)?;
797        Ok(model)
798    }
799}
800
801impl FtrlClassifier {
802    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
803    fn validate_invariants(&self) -> Result<(), RillError> {
804        // See FtrlRegressor::validate_invariants for rationale.
805        self.config.validate()?;
806        for (id, param) in &self.params {
807            param.validate()?;
808            param
809                .weight_checked(&self.config)
810                .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
811        }
812        self.intercept.validate()?;
813        self.intercept
814            .intercept_weight_checked(&self.config)
815            .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
816        Ok(())
817    }
818}
819
820impl SparseClassifier for FtrlClassifier {
821    fn samples_seen(&self) -> u64 {
822        self.samples_seen
823    }
824
825    fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
826        self.predict_proba_inner(features)
827    }
828
829    fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
830        if features.is_empty() {
831            return Err(RillError::EmptyFeatures);
832        }
833
834        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
835
836        let probability = self.predict_proba_inner(features)?;
837        ensure_finite("ftrl_probability", probability)?;
838        let y = if target { 1.0 } else { 0.0 };
839        let grad = probability - y;
840        ensure_finite("ftrl_gradient", grad)?;
841
842        let new_ids_count = features
843            .values()
844            .iter()
845            .filter(|(id, _)| !self.params.contains_key(id))
846            .count();
847        let mut skip_new_features = false;
848        if let Some(max_features) = self.config.max_features {
849            let projected = self.params.len().saturating_add(new_ids_count);
850            if projected > max_features {
851                match self.config.new_feature_policy {
852                    NewFeaturePolicy::Reject => {
853                        return Err(RillError::InvalidState(format!(
854                            "FTRL feature count {projected} exceeds max_features {max_features}"
855                        )));
856                    }
857                    NewFeaturePolicy::Ignore => {
858                        skip_new_features = true;
859                    }
860                }
861            }
862        }
863
864        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
865        for &(id, value) in features.values() {
866            // Ignore policy must be evaluated before any arithmetic on the
867            // new feature's value. Otherwise an oversized new feature whose
868            // `grad * value` would overflow could fail the whole `learn()`
869            // call even though the feature is supposed to be skipped.
870            let is_new = !self.params.contains_key(&id);
871            if is_new && skip_new_features {
872                continue;
873            }
874
875            let g = grad * value;
876            ensure_finite("ftrl_feature_gradient", g)?;
877
878            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
879                let w = param.weight(&self.config);
880                param.next_updated(g, w, &self.config)?
881            } else {
882                let param = FtrlParam::default();
883                let w = param.weight(&self.config);
884                param.next_updated(g, w, &self.config)?
885            };
886            // Verify the next state produces a finite weight before committing.
887            let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
888            ensure_finite("ftrl_next_weight", next_w)?;
889            updates.push((id, new_z, new_n));
890        }
891
892        let w_b = self.intercept.intercept_weight(&self.config);
893        let (new_intercept_z, new_intercept_n) =
894            self.intercept.next_updated(grad, w_b, &self.config)?;
895        // Verify the next intercept produces a finite weight too.
896        let next_intercept_w = FtrlParam {
897            z: new_intercept_z,
898            n: new_intercept_n,
899        }
900        .intercept_weight(&self.config);
901        ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
902
903        for (id, new_z, new_n) in updates {
904            let param = self.params.entry(id).or_default();
905            param.z = new_z;
906            param.n = new_n;
907        }
908        self.intercept.z = new_intercept_z;
909        self.intercept.n = new_intercept_n;
910        self.samples_seen = next_samples_seen;
911
912        Ok(())
913    }
914
915    fn reset(&mut self) {
916        self.params.clear();
917        self.intercept = FtrlParam::default();
918        self.samples_seen = 0;
919    }
920}
921
922#[cfg(test)]
923mod tests {
924    use super::*;
925    use rand::SeedableRng;
926
927    // -----------------------------------------------------------------
928    // FtrlRegressor tests
929    // -----------------------------------------------------------------
930
931    #[test]
932    fn cold_start_returns_zero() {
933        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
934        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
935        let pred = model.predict(&sf).unwrap();
936        assert!(pred.abs() < 1e-12);
937    }
938
939    #[test]
940    fn learn_linear_data_converges() {
941        // y = 2 * x, single feature
942        let mut model = FtrlRegressor::new(FtrlConfig {
943            alpha: 0.5,
944            beta: 1.0,
945            l1: 0.0,
946            l2: 0.0,
947            max_features: None,
948            new_feature_policy: NewFeaturePolicy::default(),
949        })
950        .unwrap();
951        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
952        let mut first_err = 0.0;
953        let mut last_err = 0.0;
954        for i in 0..500 {
955            let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
956            let y = 2.0 * x;
957            let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
958            let pred = model.predict(&sf).unwrap();
959            let err = (pred - y).abs();
960            if i < 10 {
961                first_err += err;
962            }
963            if i >= 490 {
964                last_err += err;
965            }
966            model.learn(&sf, y).unwrap();
967        }
968        assert!(last_err < first_err, "error should decrease");
969        let weights = model.weights();
970        assert_eq!(weights.len(), 1);
971        assert!(
972            (weights[0].1 - 2.0).abs() < 0.5,
973            "weight should approach 2.0"
974        );
975    }
976
977    #[test]
978    fn l1_produces_sparse_weights() {
979        // High L1 should drive most weights to zero.
980        let mut model = FtrlRegressor::new(FtrlConfig {
981            alpha: 0.1,
982            beta: 1.0,
983            l1: 100.0,
984            l2: 0.0,
985            max_features: None,
986            new_feature_policy: NewFeaturePolicy::default(),
987        })
988        .unwrap();
989        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
990        for _ in 0..200 {
991            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
992            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
993            let y = 0.5 * x1;
994            let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
995            model.learn(&sf, y).unwrap();
996        }
997        let weights = model.weights();
998        // With very high L1, all weights should be zero.
999        assert!(
1000            weights.is_empty(),
1001            "weights should all be zero, got {weights:?}"
1002        );
1003    }
1004
1005    #[test]
1006    fn dynamic_features() {
1007        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1008        assert_eq!(model.feature_count(), 0);
1009        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1010        model.learn(&sf1, 1.0).unwrap();
1011        assert_eq!(model.feature_count(), 1);
1012        // A new feature id appears.
1013        let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1014        model.learn(&sf2, 2.0).unwrap();
1015        assert_eq!(model.feature_count(), 2);
1016        // Feature 0 still present.
1017        assert!(model.params.contains_key(&0));
1018        assert!(model.params.contains_key(&5));
1019    }
1020
1021    #[test]
1022    fn predict_does_not_update_state() {
1023        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1024        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1025        let _ = model.predict(&sf).unwrap();
1026        assert_eq!(model.samples_seen(), 0);
1027        assert_eq!(model.feature_count(), 0);
1028        // Learn once, then predict again.
1029        model.learn(&sf, 1.0).unwrap();
1030        let count_after_learn = model.feature_count();
1031        let _ = model.predict(&sf).unwrap();
1032        assert_eq!(model.feature_count(), count_after_learn);
1033        assert_eq!(model.samples_seen(), 1);
1034    }
1035
1036    #[test]
1037    fn non_finite_value_rejected() {
1038        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1039        // SparseFeatures::from_sorted rejects non-finite values at construction.
1040        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1041        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1042        assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1043        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1044        assert!(model.predict(&sf).is_ok());
1045    }
1046
1047    #[test]
1048    fn non_finite_target_rejected() {
1049        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1050        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1051        assert!(model.learn(&sf, f64::NAN).is_err());
1052        assert!(model.learn(&sf, f64::INFINITY).is_err());
1053        assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1054        // State should not change on error.
1055        assert_eq!(model.samples_seen(), 0);
1056    }
1057
1058    #[test]
1059    fn empty_features_rejected() {
1060        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1061        let sf = SparseFeatures::new();
1062        assert!(model.predict(&sf).is_err());
1063        assert!(model.learn(&sf, 1.0).is_err());
1064    }
1065
1066    #[test]
1067    fn reset_clears_state() {
1068        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1069        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1070        model.learn(&sf, 3.0).unwrap();
1071        model.learn(&sf, 3.0).unwrap();
1072        assert_eq!(model.samples_seen(), 2);
1073        assert_eq!(model.feature_count(), 2);
1074        model.reset();
1075        assert_eq!(model.samples_seen(), 0);
1076        assert_eq!(model.feature_count(), 0);
1077        assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1078    }
1079
1080    #[test]
1081    fn invalid_config_rejected() {
1082        assert!(
1083            FtrlRegressor::new(FtrlConfig {
1084                alpha: 0.0,
1085                ..FtrlConfig::default()
1086            })
1087            .is_err()
1088        );
1089        assert!(
1090            FtrlRegressor::new(FtrlConfig {
1091                alpha: -1.0,
1092                ..FtrlConfig::default()
1093            })
1094            .is_err()
1095        );
1096        assert!(
1097            FtrlRegressor::new(FtrlConfig {
1098                beta: -1.0,
1099                ..FtrlConfig::default()
1100            })
1101            .is_err()
1102        );
1103        assert!(
1104            FtrlRegressor::new(FtrlConfig {
1105                l1: -1.0,
1106                ..FtrlConfig::default()
1107            })
1108            .is_err()
1109        );
1110        assert!(
1111            FtrlRegressor::new(FtrlConfig {
1112                l2: -1.0,
1113                ..FtrlConfig::default()
1114            })
1115            .is_err()
1116        );
1117        assert!(
1118            FtrlRegressor::new(FtrlConfig {
1119                alpha: f64::NAN,
1120                ..FtrlConfig::default()
1121            })
1122            .is_err()
1123        );
1124        assert!(
1125            FtrlRegressor::new(FtrlConfig {
1126                max_features: Some(0),
1127                ..FtrlConfig::default()
1128            })
1129            .is_err()
1130        );
1131    }
1132
1133    #[test]
1134    #[cfg(feature = "serde")]
1135    fn serde_roundtrip() {
1136        let mut model = FtrlRegressor::new(FtrlConfig {
1137            alpha: 0.2,
1138            beta: 0.5,
1139            l1: 0.5,
1140            l2: 0.5,
1141            max_features: Some(100),
1142            new_feature_policy: NewFeaturePolicy::Reject,
1143        })
1144        .unwrap();
1145        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1146        model.learn(&sf, 5.0).unwrap();
1147        let json = serde_json::to_string(&model).unwrap();
1148        let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1149        assert_eq!(restored.samples_seen(), model.samples_seen());
1150        assert_eq!(restored.feature_count(), model.feature_count());
1151        let pred_orig = model.predict(&sf).unwrap();
1152        let pred_restored = restored.predict(&sf).unwrap();
1153        assert!((pred_orig - pred_restored).abs() < 1e-12);
1154    }
1155
1156    #[test]
1157    fn weights_returns_nonzero_only() {
1158        let mut model = FtrlRegressor::new(FtrlConfig {
1159            alpha: 0.5,
1160            beta: 1.0,
1161            l1: 0.0,
1162            l2: 0.0,
1163            max_features: None,
1164            new_feature_policy: NewFeaturePolicy::default(),
1165        })
1166        .unwrap();
1167        // Learn feature 0 strongly, feature 1 barely.
1168        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1169        for _ in 0..50 {
1170            model.learn(&sf, 1.0).unwrap();
1171        }
1172        let weights = model.weights();
1173        // All returned weights should be non-zero.
1174        for &(_, w) in &weights {
1175            assert!(w != 0.0);
1176        }
1177        // Feature 0 should be in the list.
1178        assert!(weights.iter().any(|&(id, _)| id == 0));
1179    }
1180
1181    #[test]
1182    fn multiple_features() {
1183        // y = 1.0 * x0 + (-1.0) * x1 + 0.5
1184        let mut model = FtrlRegressor::new(FtrlConfig {
1185            alpha: 0.5,
1186            beta: 1.0,
1187            l1: 0.0,
1188            l2: 0.0,
1189            max_features: None,
1190            new_feature_policy: NewFeaturePolicy::default(),
1191        })
1192        .unwrap();
1193        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1194        for _ in 0..500 {
1195            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1196            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1197            let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1198            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1199            model.learn(&sf, y).unwrap();
1200        }
1201        let weights = model.weights();
1202        assert_eq!(weights.len(), 2);
1203        let w0 = weights
1204            .iter()
1205            .find(|&&(id, _)| id == 0)
1206            .map(|&(_, w)| w)
1207            .unwrap();
1208        let w1 = weights
1209            .iter()
1210            .find(|&&(id, _)| id == 1)
1211            .map(|&(_, w)| w)
1212            .unwrap();
1213        assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1214        assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1215        assert!(
1216            (model.intercept() - 0.5).abs() < 0.5,
1217            "intercept should approach 0.5"
1218        );
1219    }
1220
1221    #[test]
1222    fn intercept_learned() {
1223        // y = 3.0 (constant), single feature with value 0.0 so that only
1224        // the intercept can learn (feature gradient is always 0).
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: None,
1231            new_feature_policy: NewFeaturePolicy::default(),
1232        })
1233        .unwrap();
1234        let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1235        for _ in 0..300 {
1236            model.learn(&sf, 3.0).unwrap();
1237        }
1238        let pred = model.predict(&sf).unwrap();
1239        assert!(
1240            (pred - 3.0).abs() < 0.5,
1241            "prediction should approach 3.0, got {pred}"
1242        );
1243        assert!(
1244            (model.intercept() - 3.0).abs() < 0.5,
1245            "intercept should approach 3.0"
1246        );
1247        // Feature weight should be 0 (never updated since x=0).
1248        assert!(model.weights().is_empty());
1249    }
1250
1251    #[test]
1252    fn high_dim_sparse() {
1253        // 1000 possible features, only 5 active per sample.
1254        // Target is a linear combination of the active features.
1255        let mut model = FtrlRegressor::new(FtrlConfig {
1256            alpha: 0.3,
1257            beta: 1.0,
1258            l1: 0.0,
1259            l2: 0.0,
1260            max_features: None,
1261            new_feature_policy: NewFeaturePolicy::default(),
1262        })
1263        .unwrap();
1264        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1265        // True weights for features 0..5.
1266        let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1267        let mut first_err = 0.0;
1268        let mut last_err = 0.0;
1269        for i in 0..2000 {
1270            let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1271            for (j, &w) in true_w.iter().enumerate() {
1272                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1273                active.push((j as u64, x * w));
1274            }
1275            // Add some noise features with zero contribution.
1276            for k in 5..10 {
1277                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1278                active.push((k as u64 + 100, x));
1279            }
1280            active.sort_by_key(|(id, _)| *id);
1281            let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1282            let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1283            let pred = model.predict(&sf).unwrap();
1284            let err = (pred - y).abs();
1285            if i < 20 {
1286                first_err += err;
1287            }
1288            if i >= 1980 {
1289                last_err += err;
1290            }
1291            model.learn(&sf, y).unwrap();
1292        }
1293        assert!(
1294            last_err < first_err,
1295            "error should decrease in high-dim sparse"
1296        );
1297    }
1298
1299    // -----------------------------------------------------------------
1300    // FtrlRegressor: failure atomicity and overflow (ML-001/002)
1301    // -----------------------------------------------------------------
1302
1303    #[test]
1304    fn regressor_overflow_does_not_mutate_state() {
1305        // Finite inputs but intermediate `gradient * value` or
1306        // `gradient^2` overflows. The whole learn call must fail and
1307        // leave the model untouched.
1308        let mut model = FtrlRegressor::new(FtrlConfig {
1309            alpha: 0.1,
1310            beta: 1.0,
1311            l1: 0.0,
1312            l2: 0.0,
1313            max_features: None,
1314            new_feature_policy: NewFeaturePolicy::default(),
1315        })
1316        .unwrap();
1317        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1318        let result = model.learn(&sf, 1e100);
1319        assert!(result.is_err(), "expected overflow error, got {result:?}");
1320        assert_eq!(model.samples_seen(), 0);
1321        assert_eq!(model.feature_count(), 0);
1322        assert!(model.params.is_empty());
1323        assert_eq!(model.intercept.z, 0.0);
1324        assert_eq!(model.intercept.n, 0.0);
1325    }
1326
1327    #[test]
1328    fn regressor_partial_update_is_atomic() {
1329        // Two features in one sample. Feature 0 would succeed on its own,
1330        // feature 1 overflows. Neither may be committed.
1331        let mut model = FtrlRegressor::new(FtrlConfig {
1332            alpha: 0.1,
1333            beta: 1.0,
1334            l1: 0.0,
1335            l2: 0.0,
1336            max_features: None,
1337            new_feature_policy: NewFeaturePolicy::default(),
1338        })
1339        .unwrap();
1340        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1341        assert!(model.learn(&sf, 1e100).is_err());
1342        // Feature 0 must NOT be inserted.
1343        assert!(!model.params.contains_key(&0));
1344        assert!(!model.params.contains_key(&1));
1345        assert_eq!(model.samples_seen(), 0);
1346    }
1347
1348    #[test]
1349    #[cfg(feature = "serde")]
1350    fn regressor_samples_seen_overflow_is_atomic() {
1351        let json = format!(
1352            "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1353            u64::MAX
1354        );
1355        let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1356        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1357        let result = model.learn(&sf, 1.0);
1358        assert!(result.is_err(), "expected counter overflow");
1359        assert_eq!(model.samples_seen(), u64::MAX);
1360        assert_eq!(model.feature_count(), 0);
1361        assert_eq!(model.intercept.z, 0.0);
1362        assert_eq!(model.intercept.n, 0.0);
1363    }
1364
1365    // -----------------------------------------------------------------
1366    // FtrlRegressor: max_features boundary (ML-003)
1367    // -----------------------------------------------------------------
1368
1369    #[test]
1370    fn regressor_max_features_reject_at_limit() {
1371        let mut model = FtrlRegressor::new(FtrlConfig {
1372            alpha: 0.5,
1373            beta: 1.0,
1374            l1: 0.0,
1375            l2: 0.0,
1376            max_features: Some(2),
1377            new_feature_policy: NewFeaturePolicy::Reject,
1378        })
1379        .unwrap();
1380        // Reach exactly the limit.
1381        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1382        model.learn(&sf, 1.0).unwrap();
1383        assert_eq!(model.feature_count(), 2);
1384        // One more new feature: Reject must fail atomically.
1385        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1386        assert!(model.learn(&sf_new, 1.0).is_err());
1387        assert_eq!(model.feature_count(), 2);
1388        assert_eq!(model.samples_seen(), 1);
1389        // Existing features still train.
1390        let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1391        model.learn(&sf_existing, 1.0).unwrap();
1392        assert_eq!(model.feature_count(), 2);
1393        assert_eq!(model.samples_seen(), 2);
1394    }
1395
1396    #[test]
1397    fn regressor_max_features_ignore_skips_new() {
1398        let mut model = FtrlRegressor::new(FtrlConfig {
1399            alpha: 0.5,
1400            beta: 1.0,
1401            l1: 0.0,
1402            l2: 0.0,
1403            max_features: Some(2),
1404            new_feature_policy: NewFeaturePolicy::Ignore,
1405        })
1406        .unwrap();
1407        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1408        model.learn(&sf, 1.0).unwrap();
1409        // Sample with one existing and one new feature. Under Ignore the
1410        // new feature is skipped, the existing one is updated, and the
1411        // counter still advances.
1412        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1413        model.learn(&sf_mixed, 1.0).unwrap();
1414        assert_eq!(model.feature_count(), 2);
1415        assert!(!model.params.contains_key(&2));
1416        assert_eq!(model.samples_seen(), 2);
1417    }
1418
1419    #[test]
1420    fn regressor_max_features_multi_new_prejudge() {
1421        let mut model = FtrlRegressor::new(FtrlConfig {
1422            alpha: 0.5,
1423            beta: 1.0,
1424            l1: 0.0,
1425            l2: 0.0,
1426            max_features: Some(2),
1427            new_feature_policy: NewFeaturePolicy::Reject,
1428        })
1429        .unwrap();
1430        // A single sample with three new features. Reject fails atomically
1431        // without inserting any subset.
1432        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1433        assert!(model.learn(&sf, 1.0).is_err());
1434        assert_eq!(model.feature_count(), 0);
1435        assert_eq!(model.samples_seen(), 0);
1436    }
1437
1438    // -----------------------------------------------------------------
1439    // FtrlRegressor: serde validation (ML-004)
1440    // -----------------------------------------------------------------
1441
1442    #[test]
1443    #[cfg(feature = "serde")]
1444    fn regressor_serde_rejects_negative_n() {
1445        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1446        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1447        assert!(result.is_err(), "negative n must be rejected");
1448    }
1449
1450    #[test]
1451    #[cfg(feature = "serde")]
1452    fn regressor_serde_rejects_invalid_config() {
1453        let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1454        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1455        assert!(result.is_err(), "invalid alpha must be rejected");
1456    }
1457
1458    #[test]
1459    #[cfg(feature = "serde")]
1460    fn regressor_serde_accepts_missing_optional_fields() {
1461        // Old state without max_features/new_feature_policy must still load.
1462        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1463        let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1464        assert!(model.config().max_features.is_none());
1465        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1466    }
1467
1468    // -----------------------------------------------------------------
1469    // FtrlClassifier tests
1470    // -----------------------------------------------------------------
1471
1472    #[test]
1473    fn cold_start_returns_0_5() {
1474        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1475        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1476        let p = model.predict_proba(&sf).unwrap();
1477        assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1478    }
1479
1480    #[test]
1481    fn learn_separable_data() {
1482        // Linearly separable: class 1 when x0 > 0.
1483        let mut model = FtrlClassifier::new(FtrlConfig {
1484            alpha: 0.5,
1485            beta: 1.0,
1486            l1: 0.0,
1487            l2: 0.0,
1488            max_features: None,
1489            new_feature_policy: NewFeaturePolicy::default(),
1490        })
1491        .unwrap();
1492        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1493        for _ in 0..1000 {
1494            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1495            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1496            let y = x0 > 0.0;
1497            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1498            model.learn(&sf, y).unwrap();
1499        }
1500        let p_pos = model
1501            .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1502            .unwrap();
1503        let p_neg = model
1504            .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1505            .unwrap();
1506        assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1507        assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1508    }
1509
1510    #[test]
1511    fn classifier_l1_produces_sparse_weights() {
1512        let mut model = FtrlClassifier::new(FtrlConfig {
1513            alpha: 0.1,
1514            beta: 1.0,
1515            l1: 100.0,
1516            l2: 0.0,
1517            max_features: None,
1518            new_feature_policy: NewFeaturePolicy::default(),
1519        })
1520        .unwrap();
1521        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1522        for _ in 0..200 {
1523            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1524            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1525            let y = x0 > 0.0;
1526            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1527            model.learn(&sf, y).unwrap();
1528        }
1529        let weights = model.weights();
1530        assert!(
1531            weights.is_empty(),
1532            "weights should all be zero with high L1, got {weights:?}"
1533        );
1534    }
1535
1536    #[test]
1537    fn classifier_dynamic_features() {
1538        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1539        assert_eq!(model.feature_count(), 0);
1540        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1541        model.learn(&sf1, true).unwrap();
1542        assert_eq!(model.feature_count(), 1);
1543        let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1544        model.learn(&sf2, false).unwrap();
1545        assert_eq!(model.feature_count(), 2);
1546    }
1547
1548    #[test]
1549    fn classifier_predict_does_not_update_state() {
1550        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1551        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1552        let _ = model.predict_proba(&sf).unwrap();
1553        assert_eq!(model.samples_seen(), 0);
1554        assert_eq!(model.feature_count(), 0);
1555        model.learn(&sf, true).unwrap();
1556        let count = model.feature_count();
1557        let _ = model.predict_proba(&sf).unwrap();
1558        assert_eq!(model.feature_count(), count);
1559        assert_eq!(model.samples_seen(), 1);
1560    }
1561
1562    #[test]
1563    fn classifier_non_finite_value_rejected() {
1564        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1565        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1566        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1567        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1568        assert!(model.predict_proba(&sf).is_ok());
1569    }
1570
1571    #[test]
1572    fn classifier_empty_features_rejected() {
1573        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1574        let sf = SparseFeatures::new();
1575        assert!(model.predict_proba(&sf).is_err());
1576        assert!(model.learn(&sf, true).is_err());
1577    }
1578
1579    #[test]
1580    fn classifier_reset_clears_state() {
1581        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1582        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1583        model.learn(&sf, true).unwrap();
1584        model.learn(&sf, false).unwrap();
1585        assert_eq!(model.samples_seen(), 2);
1586        assert!(model.feature_count() > 0);
1587        model.reset();
1588        assert_eq!(model.samples_seen(), 0);
1589        assert_eq!(model.feature_count(), 0);
1590        let p = model.predict_proba(&sf).unwrap();
1591        assert!((p - 0.5).abs() < 1e-12);
1592    }
1593
1594    #[test]
1595    fn classifier_invalid_config_rejected() {
1596        assert!(
1597            FtrlClassifier::new(FtrlConfig {
1598                alpha: 0.0,
1599                ..FtrlConfig::default()
1600            })
1601            .is_err()
1602        );
1603        assert!(
1604            FtrlClassifier::new(FtrlConfig {
1605                beta: -0.1,
1606                ..FtrlConfig::default()
1607            })
1608            .is_err()
1609        );
1610        assert!(
1611            FtrlClassifier::new(FtrlConfig {
1612                l1: -1.0,
1613                ..FtrlConfig::default()
1614            })
1615            .is_err()
1616        );
1617        assert!(
1618            FtrlClassifier::new(FtrlConfig {
1619                l2: -1.0,
1620                ..FtrlConfig::default()
1621            })
1622            .is_err()
1623        );
1624        assert!(
1625            FtrlClassifier::new(FtrlConfig {
1626                alpha: f64::INFINITY,
1627                ..FtrlConfig::default()
1628            })
1629            .is_err()
1630        );
1631    }
1632
1633    #[test]
1634    #[cfg(feature = "serde")]
1635    fn classifier_serde_roundtrip() {
1636        let mut model = FtrlClassifier::new(FtrlConfig {
1637            alpha: 0.3,
1638            beta: 0.5,
1639            l1: 0.1,
1640            l2: 0.2,
1641            max_features: Some(100),
1642            new_feature_policy: NewFeaturePolicy::Reject,
1643        })
1644        .unwrap();
1645        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1646        model.learn(&sf, true).unwrap();
1647        model.learn(&sf, false).unwrap();
1648        let json = serde_json::to_string(&model).unwrap();
1649        let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1650        assert_eq!(restored.samples_seen(), model.samples_seen());
1651        assert_eq!(restored.feature_count(), model.feature_count());
1652        let p1 = model.predict_proba(&sf).unwrap();
1653        let p2 = restored.predict_proba(&sf).unwrap();
1654        assert!((p1 - p2).abs() < 1e-12);
1655    }
1656
1657    #[test]
1658    fn predict_proba_in_range() {
1659        let mut model = FtrlClassifier::new(FtrlConfig {
1660            alpha: 0.5,
1661            beta: 1.0,
1662            l1: 0.0,
1663            l2: 0.0,
1664            max_features: None,
1665            new_feature_policy: NewFeaturePolicy::default(),
1666        })
1667        .unwrap();
1668        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1669        for _ in 0..200 {
1670            let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1671            let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1672            let y = x0 > 0.0;
1673            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1674            model.learn(&sf, y).unwrap();
1675            let p = model.predict_proba(&sf).unwrap();
1676            assert!(
1677                (0.0..=1.0).contains(&p),
1678                "probability must be in [0,1], got {p}"
1679            );
1680        }
1681    }
1682
1683    #[test]
1684    fn learn_improves_accuracy() {
1685        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1686        // Generate a fixed test set.
1687        let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1688            .map(|_| {
1689                let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1690                let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1691                let y = x0 + x1 > 0.0;
1692                (
1693                    SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1694                    y,
1695                )
1696            })
1697            .collect();
1698
1699        let mut model = FtrlClassifier::new(FtrlConfig {
1700            alpha: 0.5,
1701            beta: 1.0,
1702            l1: 0.0,
1703            l2: 0.0,
1704            max_features: None,
1705            new_feature_policy: NewFeaturePolicy::default(),
1706        })
1707        .unwrap();
1708
1709        // Accuracy before learning (always predicts 0.5 -> threshold 0.5 -> true).
1710        let acc_before: f64 = test_set
1711            .iter()
1712            .map(|(sf, y)| {
1713                let pred = model.predict(sf).unwrap();
1714                if pred == *y { 1.0 } else { 0.0 }
1715            })
1716            .sum::<f64>()
1717            / test_set.len() as f64;
1718
1719        // Train on fresh data.
1720        for _ in 0..1000 {
1721            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1722            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1723            let y = x0 + x1 > 0.0;
1724            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1725            model.learn(&sf, y).unwrap();
1726        }
1727
1728        let acc_after: f64 = test_set
1729            .iter()
1730            .map(|(sf, y)| {
1731                let pred = model.predict(sf).unwrap();
1732                if pred == *y { 1.0 } else { 0.0 }
1733            })
1734            .sum::<f64>()
1735            / test_set.len() as f64;
1736
1737        assert!(
1738            acc_after > acc_before,
1739            "accuracy should improve: {acc_before} -> {acc_after}"
1740        );
1741    }
1742
1743    #[test]
1744    fn classifier_weights_returns_nonzero_only() {
1745        let mut model = FtrlClassifier::new(FtrlConfig {
1746            alpha: 0.5,
1747            beta: 1.0,
1748            l1: 0.0,
1749            l2: 0.0,
1750            max_features: None,
1751            new_feature_policy: NewFeaturePolicy::default(),
1752        })
1753        .unwrap();
1754        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1755        for _ in 0..50 {
1756            model.learn(&sf, true).unwrap();
1757        }
1758        let weights = model.weights();
1759        for &(_, w) in &weights {
1760            assert!(w != 0.0);
1761        }
1762    }
1763
1764    #[test]
1765    fn classifier_multiple_features() {
1766        let mut model = FtrlClassifier::new(FtrlConfig {
1767            alpha: 0.5,
1768            beta: 1.0,
1769            l1: 0.0,
1770            l2: 0.0,
1771            max_features: None,
1772            new_feature_policy: NewFeaturePolicy::default(),
1773        })
1774        .unwrap();
1775        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1776        for _ in 0..1000 {
1777            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1778            let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1779            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1780            // y = 1 if x0 + x1 > 0
1781            let y = x0 + x1 > 0.0;
1782            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1783            model.learn(&sf, y).unwrap();
1784        }
1785        let weights = model.weights();
1786        // Features 0 and 1 should have non-zero weights; feature 2 may or may not.
1787        assert!(weights.iter().any(|&(id, _)| id == 0));
1788        assert!(weights.iter().any(|&(id, _)| id == 1));
1789        // Verify prediction quality.
1790        let p_pos = model
1791            .predict_proba(
1792                &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1793            )
1794            .unwrap();
1795        let p_neg = model
1796            .predict_proba(
1797                &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1798            )
1799            .unwrap();
1800        assert!(p_pos > 0.8);
1801        assert!(p_neg < 0.2);
1802    }
1803
1804    #[test]
1805    fn log_loss_converges() {
1806        // Average log loss should decrease over training.
1807        let mut model = FtrlClassifier::new(FtrlConfig {
1808            alpha: 0.5,
1809            beta: 1.0,
1810            l1: 0.0,
1811            l2: 0.0,
1812            max_features: None,
1813            new_feature_policy: NewFeaturePolicy::default(),
1814        })
1815        .unwrap();
1816        let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1817        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1818        let mut first_loss = 0.0;
1819        let mut last_loss = 0.0;
1820        for i in 0..1000 {
1821            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1822            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1823            let y = x0 > 0.0;
1824            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1825            let p = model.predict_proba(&sf).unwrap();
1826            let loss = loss_fn.loss(p, y);
1827            if i < 20 {
1828                first_loss += loss;
1829            }
1830            if i >= 980 {
1831                last_loss += loss;
1832            }
1833            model.learn(&sf, y).unwrap();
1834        }
1835        assert!(last_loss < first_loss, "log loss should decrease");
1836    }
1837
1838    // -----------------------------------------------------------------
1839    // FtrlClassifier: failure atomicity and overflow (ML-001/002)
1840    // -----------------------------------------------------------------
1841
1842    #[test]
1843    fn classifier_overflow_does_not_mutate_state() {
1844        let mut model = FtrlClassifier::new(FtrlConfig {
1845            alpha: 0.1,
1846            beta: 1.0,
1847            l1: 0.0,
1848            l2: 0.0,
1849            max_features: None,
1850            new_feature_policy: NewFeaturePolicy::default(),
1851        })
1852        .unwrap();
1853        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1854        // Cold-start probability is 0.5, grad = 0.5 - 0.0 = 0.5.
1855        // g for feature 1 = 0.5 * 1e300 = 5e299 (finite), g^2 = 2.5e599 = inf.
1856        let result = model.learn(&sf, false);
1857        assert!(result.is_err(), "expected overflow error, got {result:?}");
1858        assert_eq!(model.samples_seen(), 0);
1859        assert_eq!(model.feature_count(), 0);
1860        assert!(model.params.is_empty());
1861        assert_eq!(model.intercept.z, 0.0);
1862        assert_eq!(model.intercept.n, 0.0);
1863    }
1864
1865    #[test]
1866    fn classifier_partial_update_is_atomic() {
1867        let mut model = FtrlClassifier::new(FtrlConfig {
1868            alpha: 0.1,
1869            beta: 1.0,
1870            l1: 0.0,
1871            l2: 0.0,
1872            max_features: None,
1873            new_feature_policy: NewFeaturePolicy::default(),
1874        })
1875        .unwrap();
1876        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1877        assert!(model.learn(&sf, false).is_err());
1878        assert!(!model.params.contains_key(&0));
1879        assert!(!model.params.contains_key(&1));
1880        assert_eq!(model.samples_seen(), 0);
1881    }
1882
1883    #[test]
1884    #[cfg(feature = "serde")]
1885    fn classifier_samples_seen_overflow_is_atomic() {
1886        let json = format!(
1887            "{{\"config\":{{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{}},\"intercept\":{{\"z\":0.0,\"n\":0.0}},\"samples_seen\":{}}}",
1888            u64::MAX
1889        );
1890        let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1891        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1892        let result = model.learn(&sf, true);
1893        assert!(result.is_err(), "expected counter overflow");
1894        assert_eq!(model.samples_seen(), u64::MAX);
1895        assert_eq!(model.feature_count(), 0);
1896        assert_eq!(model.intercept.z, 0.0);
1897        assert_eq!(model.intercept.n, 0.0);
1898    }
1899
1900    // -----------------------------------------------------------------
1901    // FtrlClassifier: max_features boundary (ML-003)
1902    // -----------------------------------------------------------------
1903
1904    #[test]
1905    fn classifier_max_features_reject_at_limit() {
1906        let mut model = FtrlClassifier::new(FtrlConfig {
1907            alpha: 0.5,
1908            beta: 1.0,
1909            l1: 0.0,
1910            l2: 0.0,
1911            max_features: Some(2),
1912            new_feature_policy: NewFeaturePolicy::Reject,
1913        })
1914        .unwrap();
1915        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1916        model.learn(&sf, true).unwrap();
1917        assert_eq!(model.feature_count(), 2);
1918        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1919        assert!(model.learn(&sf_new, true).is_err());
1920        assert_eq!(model.feature_count(), 2);
1921        assert_eq!(model.samples_seen(), 1);
1922    }
1923
1924    #[test]
1925    fn classifier_max_features_ignore_skips_new() {
1926        let mut model = FtrlClassifier::new(FtrlConfig {
1927            alpha: 0.5,
1928            beta: 1.0,
1929            l1: 0.0,
1930            l2: 0.0,
1931            max_features: Some(2),
1932            new_feature_policy: NewFeaturePolicy::Ignore,
1933        })
1934        .unwrap();
1935        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1936        model.learn(&sf, true).unwrap();
1937        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1938        model.learn(&sf_mixed, false).unwrap();
1939        assert_eq!(model.feature_count(), 2);
1940        assert!(!model.params.contains_key(&2));
1941        assert_eq!(model.samples_seen(), 2);
1942    }
1943
1944    #[test]
1945    fn classifier_max_features_multi_new_prejudge() {
1946        let mut model = FtrlClassifier::new(FtrlConfig {
1947            alpha: 0.5,
1948            beta: 1.0,
1949            l1: 0.0,
1950            l2: 0.0,
1951            max_features: Some(2),
1952            new_feature_policy: NewFeaturePolicy::Reject,
1953        })
1954        .unwrap();
1955        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1956        assert!(model.learn(&sf, true).is_err());
1957        assert_eq!(model.feature_count(), 0);
1958        assert_eq!(model.samples_seen(), 0);
1959    }
1960
1961    // -----------------------------------------------------------------
1962    // FtrlClassifier: serde validation (ML-004)
1963    // -----------------------------------------------------------------
1964
1965    #[test]
1966    #[cfg(feature = "serde")]
1967    fn classifier_serde_rejects_negative_n() {
1968        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":-1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1969        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1970        assert!(result.is_err(), "negative n must be rejected");
1971    }
1972
1973    #[test]
1974    #[cfg(feature = "serde")]
1975    fn classifier_serde_rejects_invalid_config() {
1976        let json = "{\"config\":{\"alpha\":-1.0,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1977        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1978        assert!(result.is_err(), "invalid alpha must be rejected");
1979    }
1980
1981    #[test]
1982    #[cfg(feature = "serde")]
1983    fn classifier_serde_accepts_missing_optional_fields() {
1984        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0},\"params\":{},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
1985        let model: FtrlClassifier = serde_json::from_str(json).unwrap();
1986        assert!(model.config().max_features.is_none());
1987        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1988    }
1989
1990    // -----------------------------------------------------------------
1991    // Second-audit: FTRL underflow / zero-denominator / Ignore order
1992    // -----------------------------------------------------------------
1993
1994    #[test]
1995    fn regressor_gradient_squared_underflow_is_atomic() {
1996        // Cold start: prediction = 0, target = -1e-200 → grad = 1e-200.
1997        // For feature 0 (value=1.0): g = 1e-200 (non-zero, finite).
1998        // g^2 = 1e-400 underflows to 0.0. Without the explicit underflow
1999        // check, n_new would stay at 0 while z advances, producing a state
2000        // whose next predict() divides by zero.
2001        let mut model = FtrlRegressor::new(FtrlConfig {
2002            alpha: 1.0,
2003            beta: 0.0,
2004            l1: 0.0,
2005            l2: 0.0,
2006            max_features: None,
2007            new_feature_policy: NewFeaturePolicy::default(),
2008        })
2009        .unwrap();
2010        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2011        let result = model.learn(&sf, -1e-200);
2012        assert!(result.is_err(), "expected underflow error, got {result:?}");
2013        assert_eq!(model.samples_seen(), 0);
2014        assert_eq!(model.feature_count(), 0);
2015        assert!(model.params.is_empty());
2016        assert_eq!(model.intercept.z, 0.0);
2017        assert_eq!(model.intercept.n, 0.0);
2018    }
2019
2020    #[test]
2021    fn classifier_gradient_squared_underflow_is_atomic() {
2022        // Cold start: probability = 0.5, target = false (y=0), grad = 0.5.
2023        // Feature 0 value = 1e-200: g = 0.5 * 1e-200 = 5e-201 (non-zero).
2024        // g^2 = 2.5e-401 underflows to 0.0.
2025        let mut model = FtrlClassifier::new(FtrlConfig {
2026            alpha: 1.0,
2027            beta: 0.0,
2028            l1: 0.0,
2029            l2: 0.0,
2030            max_features: None,
2031            new_feature_policy: NewFeaturePolicy::default(),
2032        })
2033        .unwrap();
2034        let sf = SparseFeatures::from_sorted(vec![(0, 1e-200)]).unwrap();
2035        let result = model.learn(&sf, false);
2036        assert!(result.is_err(), "expected underflow error, got {result:?}");
2037        assert_eq!(model.samples_seen(), 0);
2038        assert_eq!(model.feature_count(), 0);
2039    }
2040
2041    #[test]
2042    fn regressor_boundary_config_predict_after_learn_always_finite() {
2043        // beta=0, l2=0, l1=0 is the zero-denominator boundary. Every
2044        // successful learn must keep the model in a state where predict
2045        // returns a finite value.
2046        let mut model = FtrlRegressor::new(FtrlConfig {
2047            alpha: 1.0,
2048            beta: 0.0,
2049            l1: 0.0,
2050            l2: 0.0,
2051            max_features: None,
2052            new_feature_policy: NewFeaturePolicy::default(),
2053        })
2054        .unwrap();
2055        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2056        for _ in 0..100 {
2057            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2058            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2059            let y = 2.0 * x0 - x1;
2060            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2061            model.learn(&sf, y).unwrap();
2062            let pred = model.predict(&sf);
2063            assert!(
2064                pred.is_ok(),
2065                "predict failed after successful learn: {pred:?}"
2066            );
2067            assert!(
2068                pred.unwrap().is_finite(),
2069                "predict must return finite value after successful learn"
2070            );
2071        }
2072    }
2073
2074    #[test]
2075    fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2076        let mut model = FtrlClassifier::new(FtrlConfig {
2077            alpha: 1.0,
2078            beta: 0.0,
2079            l1: 0.0,
2080            l2: 0.0,
2081            max_features: None,
2082            new_feature_policy: NewFeaturePolicy::default(),
2083        })
2084        .unwrap();
2085        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2086        for _ in 0..100 {
2087            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2088            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2089            let y = x0 > 0.0;
2090            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2091            model.learn(&sf, y).unwrap();
2092            let proba = model.predict_proba(&sf);
2093            assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2094            let p = proba.unwrap();
2095            assert!(p.is_finite(), "probability must be finite, got {p}");
2096            assert!(
2097                (0.0..=1.0).contains(&p),
2098                "probability must be in [0,1], got {p}"
2099            );
2100        }
2101    }
2102
2103    #[test]
2104    #[cfg(feature = "serde")]
2105    fn regressor_serde_rejects_n_zero_z_nonzero() {
2106        // n=0, z=1.0: weight formula denominator = l2 + (beta + sqrt(0))/alpha.
2107        // With beta=0, l2=0, alpha=1: denominator = 0 → weight = inf.
2108        // validate() must reject this state.
2109        let json = "{\"config\":{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":0.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2110        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2111        assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2112    }
2113
2114    #[test]
2115    #[cfg(feature = "serde")]
2116    fn classifier_serde_rejects_n_zero_z_nonzero() {
2117        let json = "{\"config\":{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":0.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2118        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2119        assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2120    }
2121
2122    #[test]
2123    #[cfg(feature = "serde")]
2124    fn regressor_predict_dot_plus_intercept_overflow() {
2125        // Craft a model where dot + intercept overflows f64.
2126        // config: alpha=1, beta=0, l1=0, l2=0 → weight = -z / sqrt(n)
2127        // z = -f64::MAX * 0.75, n = 1.0 → weight = f64::MAX * 0.75
2128        // dot = weight * 1.0 = f64::MAX * 0.75
2129        // intercept_weight = f64::MAX * 0.75
2130        // dot + intercept = f64::MAX * 1.5 → overflow to inf.
2131        let z = -f64::MAX * 0.75;
2132        let json = format!(
2133            "{{\"config\":{{\"alpha\":1.0,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"}},\"params\":{{\"0\":{{\"z\":{0},\"n\":1.0}}}},\"intercept\":{{\"z\":{0},\"n\":1.0}},\"samples_seen\":1}}",
2134            z
2135        );
2136        let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2137        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2138        let result = model.predict(&sf);
2139        assert!(
2140            result.is_err(),
2141            "expected dot+intercept overflow error, got {result:?}"
2142        );
2143    }
2144
2145    #[test]
2146    fn regressor_ignore_skips_overflowing_new_feature() {
2147        let mut model = FtrlRegressor::new(FtrlConfig {
2148            alpha: 0.5,
2149            beta: 1.0,
2150            l1: 0.0,
2151            l2: 0.0,
2152            max_features: Some(1),
2153            new_feature_policy: NewFeaturePolicy::Ignore,
2154        })
2155        .unwrap();
2156        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2157        model.learn(&sf, 1.0).unwrap();
2158        assert_eq!(model.feature_count(), 1);
2159
2160        // New feature 1 with huge value. grad is O(1), so
2161        // g = grad * f64::MAX is finite, but g^2 overflows to inf.
2162        // Under Ignore, the new feature must be skipped BEFORE the
2163        // multiplication; otherwise the whole learn() fails.
2164        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2165        let result = model.learn(&sf_mixed, 1.0);
2166        assert!(
2167            result.is_ok(),
2168            "Ignore must skip overflowing new feature, got {result:?}"
2169        );
2170        assert_eq!(model.feature_count(), 1);
2171        assert!(!model.params.contains_key(&1));
2172        assert_eq!(model.samples_seen(), 2);
2173        assert!(model.predict(&sf).is_ok());
2174    }
2175
2176    #[test]
2177    fn classifier_ignore_skips_overflowing_new_feature() {
2178        let mut model = FtrlClassifier::new(FtrlConfig {
2179            alpha: 0.5,
2180            beta: 1.0,
2181            l1: 0.0,
2182            l2: 0.0,
2183            max_features: Some(1),
2184            new_feature_policy: NewFeaturePolicy::Ignore,
2185        })
2186        .unwrap();
2187        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2188        model.learn(&sf, true).unwrap();
2189        assert_eq!(model.feature_count(), 1);
2190
2191        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2192        let result = model.learn(&sf_mixed, true);
2193        assert!(
2194            result.is_ok(),
2195            "Ignore must skip overflowing new feature, got {result:?}"
2196        );
2197        assert_eq!(model.feature_count(), 1);
2198        assert!(!model.params.contains_key(&1));
2199        assert_eq!(model.samples_seen(), 2);
2200        assert!(model.predict_proba(&sf).is_ok());
2201    }
2202
2203    // -----------------------------------------------------------------
2204    // §6.2: FTRL config-aware serde validation
2205    // -----------------------------------------------------------------
2206
2207    #[test]
2208    #[cfg(feature = "serde")]
2209    fn regressor_serde_rejects_config_dependent_zero_denominator() {
2210        // alpha = f64::MAX, beta = 0, l2 = 0, l1 = 0, n = 1e-300, z = 1.0.
2211        // sqrt(n) = 1e-150; denominator = (0 + 1e-150) / f64::MAX underflows
2212        // to 0.0. The basic validate() passes (n > 0, z finite) but
2213        // weight_checked must reject the zero denominator explicitly.
2214        let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1e-300}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2215        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2216        assert!(
2217            result.is_err(),
2218            "config-dependent zero denominator must be rejected"
2219        );
2220    }
2221
2222    #[test]
2223    #[cfg(feature = "serde")]
2224    fn classifier_serde_rejects_config_dependent_zero_denominator() {
2225        let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1e-300}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2226        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2227        assert!(
2228            result.is_err(),
2229            "config-dependent zero denominator must be rejected"
2230        );
2231    }
2232
2233    #[test]
2234    #[cfg(feature = "serde")]
2235    fn regressor_serde_rejects_intercept_zero_denominator() {
2236        // Same dangerous config, but the malicious state is in the
2237        // intercept rather than a feature param.
2238        let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":1.0,\"n\":1e-300},\"samples_seen\":0}";
2239        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2240        assert!(
2241            result.is_err(),
2242            "intercept zero denominator must be rejected"
2243        );
2244    }
2245
2246    #[test]
2247    #[cfg(feature = "serde")]
2248    fn classifier_serde_rejects_intercept_zero_denominator() {
2249        let json = "{\"config\":{\"alpha\":1.7976931348623157e308,\"beta\":0.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{},\"intercept\":{\"z\":1.0,\"n\":1e-300},\"samples_seen\":0}";
2250        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2251        assert!(
2252            result.is_err(),
2253            "intercept zero denominator must be rejected"
2254        );
2255    }
2256
2257    #[test]
2258    #[cfg(feature = "serde")]
2259    fn regressor_valid_boundary_state_roundtrips() {
2260        // Boundary config (beta=0, l2=0) with legitimately trained state
2261        // must round-trip through serde without being rejected by the new
2262        // config-aware checks.
2263        let mut model = FtrlRegressor::new(FtrlConfig {
2264            alpha: 1.0,
2265            beta: 0.0,
2266            l1: 0.0,
2267            l2: 0.0,
2268            max_features: None,
2269            new_feature_policy: NewFeaturePolicy::default(),
2270        })
2271        .unwrap();
2272        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2273        model.learn(&sf, 3.0).unwrap();
2274        model.learn(&sf, 5.0).unwrap();
2275        let json = serde_json::to_string(&model).unwrap();
2276        let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2277        assert_eq!(restored.samples_seen(), model.samples_seen());
2278        assert_eq!(restored.feature_count(), model.feature_count());
2279        let p1 = model.predict(&sf).unwrap();
2280        let p2 = restored.predict(&sf).unwrap();
2281        assert!((p1 - p2).abs() < 1e-12);
2282        // weights() must not produce Infinity on the restored model.
2283        for (_, w) in restored.weights() {
2284            assert!(w.is_finite(), "restored weight must be finite, got {w}");
2285        }
2286        assert!(restored.intercept().is_finite());
2287    }
2288
2289    #[test]
2290    #[cfg(feature = "serde")]
2291    fn classifier_valid_boundary_state_roundtrips() {
2292        let mut model = FtrlClassifier::new(FtrlConfig {
2293            alpha: 1.0,
2294            beta: 0.0,
2295            l1: 0.0,
2296            l2: 0.0,
2297            max_features: None,
2298            new_feature_policy: NewFeaturePolicy::default(),
2299        })
2300        .unwrap();
2301        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2302        model.learn(&sf, true).unwrap();
2303        model.learn(&sf, false).unwrap();
2304        let json = serde_json::to_string(&model).unwrap();
2305        let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2306        assert_eq!(restored.samples_seen(), model.samples_seen());
2307        assert_eq!(restored.feature_count(), model.feature_count());
2308        let p1 = model.predict_proba(&sf).unwrap();
2309        let p2 = restored.predict_proba(&sf).unwrap();
2310        assert!((p1 - p2).abs() < 1e-12);
2311        for (_, w) in restored.weights() {
2312            assert!(w.is_finite(), "restored weight must be finite, got {w}");
2313        }
2314        assert!(restored.intercept().is_finite());
2315    }
2316}