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        // Top-level invariant: the stored feature count must not exceed
566        // `max_features`. A malicious payload with `max_features = 1` but
567        // two stored params violates the model contract and is rejected
568        // here rather than silently accepted.
569        if let Some(max_features) = self.config.max_features
570            && self.params.len() > max_features
571        {
572            return Err(RillError::InvalidState(format!(
573                "FTRL stored feature count {} exceeds max_features {}",
574                self.params.len(),
575                max_features
576            )));
577        }
578        for (id, param) in &self.params {
579            param.validate()?;
580            param
581                .weight_checked(&self.config)
582                .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
583        }
584        self.intercept.validate()?;
585        self.intercept
586            .intercept_weight_checked(&self.config)
587            .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
588        Ok(())
589    }
590}
591
592impl SparseRegressor for FtrlRegressor {
593    fn samples_seen(&self) -> u64 {
594        self.samples_seen
595    }
596
597    fn predict(&self, features: &SparseFeatures) -> Result<f64, RillError> {
598        self.predict_inner(features)
599    }
600
601    fn learn(&mut self, features: &SparseFeatures, target: f64) -> Result<(), RillError> {
602        if features.is_empty() {
603            return Err(RillError::EmptyFeatures);
604        }
605        ensure_finite("target", target)?;
606
607        // Reserve the sample counter first. If it overflows, no state changes.
608        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
609
610        let prediction = self.predict_inner(features)?;
611        ensure_finite("ftrl_prediction", prediction)?;
612        let grad = prediction - target;
613        ensure_finite("ftrl_gradient", grad)?;
614
615        // Pre-judge the max_features policy. We never insert a subset of
616        // new features then fail: either all new features are accepted or
617        // (under Ignore) all new features are skipped.
618        let new_ids_count = features
619            .values()
620            .iter()
621            .filter(|(id, _)| !self.params.contains_key(id))
622            .count();
623        let mut skip_new_features = false;
624        if let Some(max_features) = self.config.max_features {
625            let projected = self.params.len().saturating_add(new_ids_count);
626            if projected > max_features {
627                match self.config.new_feature_policy {
628                    NewFeaturePolicy::Reject => {
629                        return Err(RillError::InvalidState(format!(
630                            "FTRL feature count {projected} exceeds max_features {max_features}"
631                        )));
632                    }
633                    NewFeaturePolicy::Ignore => {
634                        skip_new_features = true;
635                    }
636                }
637            }
638        }
639
640        // Compute the next (z, n) for every feature without touching self.
641        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
642        for &(id, value) in features.values() {
643            // Ignore policy must be evaluated before any arithmetic on the
644            // new feature's value. Otherwise an oversized new feature whose
645            // `grad * value` would overflow could fail the whole `learn()`
646            // call even though the feature is supposed to be skipped.
647            let is_new = !self.params.contains_key(&id);
648            if is_new && skip_new_features {
649                continue;
650            }
651
652            let g = grad * value;
653            ensure_finite("ftrl_feature_gradient", g)?;
654
655            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
656                let w = param.weight(&self.config);
657                param.next_updated(g, w, &self.config)?
658            } else {
659                let param = FtrlParam::default();
660                let w = param.weight(&self.config);
661                param.next_updated(g, w, &self.config)?
662            };
663            // Verify the next state produces a finite weight before
664            // committing. This catches any path that would leave the model
665            // in a state where the next `predict()` fails due to internal
666            // state (e.g. a zero denominator in the weight formula).
667            let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
668            ensure_finite("ftrl_next_weight", next_w)?;
669            updates.push((id, new_z, new_n));
670        }
671
672        // Compute the next intercept state.
673        let w_b = self.intercept.intercept_weight(&self.config);
674        let (new_intercept_z, new_intercept_n) =
675            self.intercept.next_updated(grad, w_b, &self.config)?;
676        // Verify the next intercept produces a finite weight too.
677        let next_intercept_w = FtrlParam {
678            z: new_intercept_z,
679            n: new_intercept_n,
680        }
681        .intercept_weight(&self.config);
682        ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
683
684        // Commit atomically. No failure path beyond this point.
685        for (id, new_z, new_n) in updates {
686            let param = self.params.entry(id).or_default();
687            param.z = new_z;
688            param.n = new_n;
689        }
690        self.intercept.z = new_intercept_z;
691        self.intercept.n = new_intercept_n;
692        self.samples_seen = next_samples_seen;
693
694        Ok(())
695    }
696
697    fn reset(&mut self) {
698        self.params.clear();
699        self.intercept = FtrlParam::default();
700        self.samples_seen = 0;
701    }
702}
703
704/// FTRL binary classifier with log loss.
705///
706/// Predicts `P(y=1 | x) = sigmoid(w · x + b)`. The gradient of the log loss
707/// w.r.t. the logit simplifies to `probability - target`, so each feature's
708/// gradient is `(probability - target) * x_i`.
709///
710/// The returned probability lies in `[0, 1]`. Extreme logits can produce
711/// exactly `0.0` or `1.0` after `sigmoid`; downstream consumers such as
712/// [`crate::loss::log_loss::BinaryLogLoss`] clip internally.
713///
714/// # Examples
715///
716/// ```
717/// use rill_ml::models::{FtrlClassifier, FtrlConfig};
718/// use rill_ml::sparse::SparseFeatures;
719/// use rill_ml::SparseClassifier;
720///
721/// let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
722/// let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
723/// let _proba = model.predict_proba(&sf).unwrap();
724/// model.learn(&sf, true).unwrap();
725/// ```
726#[derive(Debug, Clone)]
727#[cfg_attr(feature = "serde", derive(serde::Serialize))]
728pub struct FtrlClassifier {
729    config: FtrlConfig,
730    params: BTreeMap<FeatureId, FtrlParam>,
731    intercept: FtrlParam,
732    samples_seen: u64,
733}
734
735impl FtrlClassifier {
736    /// Create a new FTRL classifier.
737    ///
738    /// Returns an error if the configuration is invalid.
739    pub fn new(config: FtrlConfig) -> Result<Self, RillError> {
740        config.validate()?;
741        Ok(Self {
742            config,
743            params: BTreeMap::new(),
744            intercept: FtrlParam::default(),
745            samples_seen: 0,
746        })
747    }
748
749    /// The model configuration.
750    pub const fn config(&self) -> &FtrlConfig {
751        &self.config
752    }
753
754    /// Return the current non-zero feature weights, sorted by `FeatureId`.
755    ///
756    /// Features whose FTRL weight is exactly zero (due to L1
757    /// soft-thresholding or never having been updated) are excluded.
758    pub fn weights(&self) -> Vec<(FeatureId, f64)> {
759        self.params
760            .iter()
761            .map(|(&id, param)| (id, param.weight(&self.config)))
762            .filter(|&(_, w)| w != 0.0)
763            .collect()
764    }
765
766    /// Compute the current intercept (bias) weight.
767    pub fn intercept(&self) -> f64 {
768        self.intercept.intercept_weight(&self.config)
769    }
770
771    /// Number of distinct features the model has seen.
772    pub fn feature_count(&self) -> usize {
773        self.params.len()
774    }
775
776    /// Compute the probability `sigmoid(w · x + b)` without updating state.
777    fn predict_proba_inner(&self, features: &SparseFeatures) -> Result<f64, RillError> {
778        let dot = compute_dot(&self.params, &self.config, features)?;
779        let intercept = self.intercept.intercept_weight(&self.config);
780        ensure_finite("ftrl_intercept", intercept)?;
781        let logit = checked_finite_add(dot, intercept, "ftrl_logit")?;
782        Ok(sigmoid(logit))
783    }
784}
785
786#[cfg(feature = "serde")]
787impl<'de> serde::Deserialize<'de> for FtrlClassifier {
788    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
789    where
790        D: serde::Deserializer<'de>,
791    {
792        #[derive(serde::Deserialize)]
793        struct FtrlClassifierState {
794            config: FtrlConfig,
795            params: BTreeMap<FeatureId, FtrlParam>,
796            intercept: FtrlParam,
797            samples_seen: u64,
798        }
799
800        let state = FtrlClassifierState::deserialize(deserializer)?;
801        let model = FtrlClassifier {
802            config: state.config,
803            params: state.params,
804            intercept: state.intercept,
805            samples_seen: state.samples_seen,
806        };
807        model
808            .validate_invariants()
809            .map_err(serde::de::Error::custom)?;
810        Ok(model)
811    }
812}
813
814impl FtrlClassifier {
815    #[cfg_attr(not(feature = "serde"), allow(dead_code))]
816    fn validate_invariants(&self) -> Result<(), RillError> {
817        // See FtrlRegressor::validate_invariants for rationale.
818        self.config.validate()?;
819        // Top-level invariant: the stored feature count must not exceed
820        // `max_features`. See FtrlRegressor::validate_invariants.
821        if let Some(max_features) = self.config.max_features
822            && self.params.len() > max_features
823        {
824            return Err(RillError::InvalidState(format!(
825                "FTRL stored feature count {} exceeds max_features {}",
826                self.params.len(),
827                max_features
828            )));
829        }
830        for (id, param) in &self.params {
831            param.validate()?;
832            param
833                .weight_checked(&self.config)
834                .map_err(|e| RillError::InvalidState(format!("ftrl feature {id} weight: {e}")))?;
835        }
836        self.intercept.validate()?;
837        self.intercept
838            .intercept_weight_checked(&self.config)
839            .map_err(|e| RillError::InvalidState(format!("ftrl intercept weight: {e}")))?;
840        Ok(())
841    }
842}
843
844impl SparseClassifier for FtrlClassifier {
845    fn samples_seen(&self) -> u64 {
846        self.samples_seen
847    }
848
849    fn predict_proba(&self, features: &SparseFeatures) -> Result<f64, RillError> {
850        self.predict_proba_inner(features)
851    }
852
853    fn learn(&mut self, features: &SparseFeatures, target: bool) -> Result<(), RillError> {
854        if features.is_empty() {
855            return Err(RillError::EmptyFeatures);
856        }
857
858        let next_samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
859
860        let probability = self.predict_proba_inner(features)?;
861        ensure_finite("ftrl_probability", probability)?;
862        let y = if target { 1.0 } else { 0.0 };
863        let grad = probability - y;
864        ensure_finite("ftrl_gradient", grad)?;
865
866        let new_ids_count = features
867            .values()
868            .iter()
869            .filter(|(id, _)| !self.params.contains_key(id))
870            .count();
871        let mut skip_new_features = false;
872        if let Some(max_features) = self.config.max_features {
873            let projected = self.params.len().saturating_add(new_ids_count);
874            if projected > max_features {
875                match self.config.new_feature_policy {
876                    NewFeaturePolicy::Reject => {
877                        return Err(RillError::InvalidState(format!(
878                            "FTRL feature count {projected} exceeds max_features {max_features}"
879                        )));
880                    }
881                    NewFeaturePolicy::Ignore => {
882                        skip_new_features = true;
883                    }
884                }
885            }
886        }
887
888        let mut updates: Vec<(FeatureId, f64, f64)> = Vec::with_capacity(features.len());
889        for &(id, value) in features.values() {
890            // Ignore policy must be evaluated before any arithmetic on the
891            // new feature's value. Otherwise an oversized new feature whose
892            // `grad * value` would overflow could fail the whole `learn()`
893            // call even though the feature is supposed to be skipped.
894            let is_new = !self.params.contains_key(&id);
895            if is_new && skip_new_features {
896                continue;
897            }
898
899            let g = grad * value;
900            ensure_finite("ftrl_feature_gradient", g)?;
901
902            let (new_z, new_n) = if let Some(param) = self.params.get(&id) {
903                let w = param.weight(&self.config);
904                param.next_updated(g, w, &self.config)?
905            } else {
906                let param = FtrlParam::default();
907                let w = param.weight(&self.config);
908                param.next_updated(g, w, &self.config)?
909            };
910            // Verify the next state produces a finite weight before committing.
911            let next_w = FtrlParam { z: new_z, n: new_n }.weight(&self.config);
912            ensure_finite("ftrl_next_weight", next_w)?;
913            updates.push((id, new_z, new_n));
914        }
915
916        let w_b = self.intercept.intercept_weight(&self.config);
917        let (new_intercept_z, new_intercept_n) =
918            self.intercept.next_updated(grad, w_b, &self.config)?;
919        // Verify the next intercept produces a finite weight too.
920        let next_intercept_w = FtrlParam {
921            z: new_intercept_z,
922            n: new_intercept_n,
923        }
924        .intercept_weight(&self.config);
925        ensure_finite("ftrl_next_intercept_weight", next_intercept_w)?;
926
927        for (id, new_z, new_n) in updates {
928            let param = self.params.entry(id).or_default();
929            param.z = new_z;
930            param.n = new_n;
931        }
932        self.intercept.z = new_intercept_z;
933        self.intercept.n = new_intercept_n;
934        self.samples_seen = next_samples_seen;
935
936        Ok(())
937    }
938
939    fn reset(&mut self) {
940        self.params.clear();
941        self.intercept = FtrlParam::default();
942        self.samples_seen = 0;
943    }
944}
945
946#[cfg(test)]
947mod tests {
948    use super::*;
949    use rand::SeedableRng;
950
951    // -----------------------------------------------------------------
952    // FtrlRegressor tests
953    // -----------------------------------------------------------------
954
955    #[test]
956    fn cold_start_returns_zero() {
957        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
958        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
959        let pred = model.predict(&sf).unwrap();
960        assert!(pred.abs() < 1e-12);
961    }
962
963    #[test]
964    fn learn_linear_data_converges() {
965        // y = 2 * x, single feature
966        let mut model = FtrlRegressor::new(FtrlConfig {
967            alpha: 0.5,
968            beta: 1.0,
969            l1: 0.0,
970            l2: 0.0,
971            max_features: None,
972            new_feature_policy: NewFeaturePolicy::default(),
973        })
974        .unwrap();
975        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(42);
976        let mut first_err = 0.0;
977        let mut last_err = 0.0;
978        for i in 0..500 {
979            let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
980            let y = 2.0 * x;
981            let sf = SparseFeatures::from_sorted(vec![(0, x)]).unwrap();
982            let pred = model.predict(&sf).unwrap();
983            let err = (pred - y).abs();
984            if i < 10 {
985                first_err += err;
986            }
987            if i >= 490 {
988                last_err += err;
989            }
990            model.learn(&sf, y).unwrap();
991        }
992        assert!(last_err < first_err, "error should decrease");
993        let weights = model.weights();
994        assert_eq!(weights.len(), 1);
995        assert!(
996            (weights[0].1 - 2.0).abs() < 0.5,
997            "weight should approach 2.0"
998        );
999    }
1000
1001    #[test]
1002    fn l1_produces_sparse_weights() {
1003        // High L1 should drive most weights to zero.
1004        let mut model = FtrlRegressor::new(FtrlConfig {
1005            alpha: 0.1,
1006            beta: 1.0,
1007            l1: 100.0,
1008            l2: 0.0,
1009            max_features: None,
1010            new_feature_policy: NewFeaturePolicy::default(),
1011        })
1012        .unwrap();
1013        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(1);
1014        for _ in 0..200 {
1015            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1016            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1017            let y = 0.5 * x1;
1018            let sf = SparseFeatures::from_sorted(vec![(0, x1), (1, x2)]).unwrap();
1019            model.learn(&sf, y).unwrap();
1020        }
1021        let weights = model.weights();
1022        // With very high L1, all weights should be zero.
1023        assert!(
1024            weights.is_empty(),
1025            "weights should all be zero, got {weights:?}"
1026        );
1027    }
1028
1029    #[test]
1030    fn dynamic_features() {
1031        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1032        assert_eq!(model.feature_count(), 0);
1033        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1034        model.learn(&sf1, 1.0).unwrap();
1035        assert_eq!(model.feature_count(), 1);
1036        // A new feature id appears.
1037        let sf2 = SparseFeatures::from_sorted(vec![(5, 2.0)]).unwrap();
1038        model.learn(&sf2, 2.0).unwrap();
1039        assert_eq!(model.feature_count(), 2);
1040        // Feature 0 still present.
1041        assert!(model.params.contains_key(&0));
1042        assert!(model.params.contains_key(&5));
1043    }
1044
1045    #[test]
1046    fn predict_does_not_update_state() {
1047        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1048        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1049        let _ = model.predict(&sf).unwrap();
1050        assert_eq!(model.samples_seen(), 0);
1051        assert_eq!(model.feature_count(), 0);
1052        // Learn once, then predict again.
1053        model.learn(&sf, 1.0).unwrap();
1054        let count_after_learn = model.feature_count();
1055        let _ = model.predict(&sf).unwrap();
1056        assert_eq!(model.feature_count(), count_after_learn);
1057        assert_eq!(model.samples_seen(), 1);
1058    }
1059
1060    #[test]
1061    fn non_finite_value_rejected() {
1062        let model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1063        // SparseFeatures::from_sorted rejects non-finite values at construction.
1064        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1065        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1066        assert!(SparseFeatures::from_sorted(vec![(0, f64::NEG_INFINITY)]).is_err());
1067        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1068        assert!(model.predict(&sf).is_ok());
1069    }
1070
1071    #[test]
1072    fn non_finite_target_rejected() {
1073        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1074        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1075        assert!(model.learn(&sf, f64::NAN).is_err());
1076        assert!(model.learn(&sf, f64::INFINITY).is_err());
1077        assert!(model.learn(&sf, f64::NEG_INFINITY).is_err());
1078        // State should not change on error.
1079        assert_eq!(model.samples_seen(), 0);
1080    }
1081
1082    #[test]
1083    fn empty_features_rejected() {
1084        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1085        let sf = SparseFeatures::new();
1086        assert!(model.predict(&sf).is_err());
1087        assert!(model.learn(&sf, 1.0).is_err());
1088    }
1089
1090    #[test]
1091    fn reset_clears_state() {
1092        let mut model = FtrlRegressor::new(FtrlConfig::default()).unwrap();
1093        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1094        model.learn(&sf, 3.0).unwrap();
1095        model.learn(&sf, 3.0).unwrap();
1096        assert_eq!(model.samples_seen(), 2);
1097        assert_eq!(model.feature_count(), 2);
1098        model.reset();
1099        assert_eq!(model.samples_seen(), 0);
1100        assert_eq!(model.feature_count(), 0);
1101        assert!(model.predict(&sf).unwrap().abs() < 1e-12);
1102    }
1103
1104    #[test]
1105    fn invalid_config_rejected() {
1106        assert!(
1107            FtrlRegressor::new(FtrlConfig {
1108                alpha: 0.0,
1109                ..FtrlConfig::default()
1110            })
1111            .is_err()
1112        );
1113        assert!(
1114            FtrlRegressor::new(FtrlConfig {
1115                alpha: -1.0,
1116                ..FtrlConfig::default()
1117            })
1118            .is_err()
1119        );
1120        assert!(
1121            FtrlRegressor::new(FtrlConfig {
1122                beta: -1.0,
1123                ..FtrlConfig::default()
1124            })
1125            .is_err()
1126        );
1127        assert!(
1128            FtrlRegressor::new(FtrlConfig {
1129                l1: -1.0,
1130                ..FtrlConfig::default()
1131            })
1132            .is_err()
1133        );
1134        assert!(
1135            FtrlRegressor::new(FtrlConfig {
1136                l2: -1.0,
1137                ..FtrlConfig::default()
1138            })
1139            .is_err()
1140        );
1141        assert!(
1142            FtrlRegressor::new(FtrlConfig {
1143                alpha: f64::NAN,
1144                ..FtrlConfig::default()
1145            })
1146            .is_err()
1147        );
1148        assert!(
1149            FtrlRegressor::new(FtrlConfig {
1150                max_features: Some(0),
1151                ..FtrlConfig::default()
1152            })
1153            .is_err()
1154        );
1155    }
1156
1157    #[test]
1158    #[cfg(feature = "serde")]
1159    fn serde_roundtrip() {
1160        let mut model = FtrlRegressor::new(FtrlConfig {
1161            alpha: 0.2,
1162            beta: 0.5,
1163            l1: 0.5,
1164            l2: 0.5,
1165            max_features: Some(100),
1166            new_feature_policy: NewFeaturePolicy::Reject,
1167        })
1168        .unwrap();
1169        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (3, 2.0)]).unwrap();
1170        model.learn(&sf, 5.0).unwrap();
1171        let json = serde_json::to_string(&model).unwrap();
1172        let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
1173        assert_eq!(restored.samples_seen(), model.samples_seen());
1174        assert_eq!(restored.feature_count(), model.feature_count());
1175        let pred_orig = model.predict(&sf).unwrap();
1176        let pred_restored = restored.predict(&sf).unwrap();
1177        assert!((pred_orig - pred_restored).abs() < 1e-12);
1178    }
1179
1180    #[test]
1181    fn weights_returns_nonzero_only() {
1182        let mut model = FtrlRegressor::new(FtrlConfig {
1183            alpha: 0.5,
1184            beta: 1.0,
1185            l1: 0.0,
1186            l2: 0.0,
1187            max_features: None,
1188            new_feature_policy: NewFeaturePolicy::default(),
1189        })
1190        .unwrap();
1191        // Learn feature 0 strongly, feature 1 barely.
1192        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1193        for _ in 0..50 {
1194            model.learn(&sf, 1.0).unwrap();
1195        }
1196        let weights = model.weights();
1197        // All returned weights should be non-zero.
1198        for &(_, w) in &weights {
1199            assert!(w != 0.0);
1200        }
1201        // Feature 0 should be in the list.
1202        assert!(weights.iter().any(|&(id, _)| id == 0));
1203    }
1204
1205    #[test]
1206    fn multiple_features() {
1207        // y = 1.0 * x0 + (-1.0) * x1 + 0.5
1208        let mut model = FtrlRegressor::new(FtrlConfig {
1209            alpha: 0.5,
1210            beta: 1.0,
1211            l1: 0.0,
1212            l2: 0.0,
1213            max_features: None,
1214            new_feature_policy: NewFeaturePolicy::default(),
1215        })
1216        .unwrap();
1217        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(99);
1218        for _ in 0..500 {
1219            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1220            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1221            let y = 1.0 * x0 - 1.0 * x1 + 0.5;
1222            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1223            model.learn(&sf, y).unwrap();
1224        }
1225        let weights = model.weights();
1226        assert_eq!(weights.len(), 2);
1227        let w0 = weights
1228            .iter()
1229            .find(|&&(id, _)| id == 0)
1230            .map(|&(_, w)| w)
1231            .unwrap();
1232        let w1 = weights
1233            .iter()
1234            .find(|&&(id, _)| id == 1)
1235            .map(|&(_, w)| w)
1236            .unwrap();
1237        assert!((w0 - 1.0).abs() < 0.5, "w0 should approach 1.0, got {w0}");
1238        assert!((w1 + 1.0).abs() < 0.5, "w1 should approach -1.0, got {w1}");
1239        assert!(
1240            (model.intercept() - 0.5).abs() < 0.5,
1241            "intercept should approach 0.5"
1242        );
1243    }
1244
1245    #[test]
1246    fn intercept_learned() {
1247        // y = 3.0 (constant), single feature with value 0.0 so that only
1248        // the intercept can learn (feature gradient is always 0).
1249        let mut model = FtrlRegressor::new(FtrlConfig {
1250            alpha: 0.5,
1251            beta: 1.0,
1252            l1: 0.0,
1253            l2: 0.0,
1254            max_features: None,
1255            new_feature_policy: NewFeaturePolicy::default(),
1256        })
1257        .unwrap();
1258        let sf = SparseFeatures::from_sorted(vec![(0, 0.0)]).unwrap();
1259        for _ in 0..300 {
1260            model.learn(&sf, 3.0).unwrap();
1261        }
1262        let pred = model.predict(&sf).unwrap();
1263        assert!(
1264            (pred - 3.0).abs() < 0.5,
1265            "prediction should approach 3.0, got {pred}"
1266        );
1267        assert!(
1268            (model.intercept() - 3.0).abs() < 0.5,
1269            "intercept should approach 3.0"
1270        );
1271        // Feature weight should be 0 (never updated since x=0).
1272        assert!(model.weights().is_empty());
1273    }
1274
1275    #[test]
1276    fn high_dim_sparse() {
1277        // 1000 possible features, only 5 active per sample.
1278        // Target is a linear combination of the active features.
1279        let mut model = FtrlRegressor::new(FtrlConfig {
1280            alpha: 0.3,
1281            beta: 1.0,
1282            l1: 0.0,
1283            l2: 0.0,
1284            max_features: None,
1285            new_feature_policy: NewFeaturePolicy::default(),
1286        })
1287        .unwrap();
1288        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
1289        // True weights for features 0..5.
1290        let true_w = [1.0, -0.5, 2.0, 0.3, -1.5];
1291        let mut first_err = 0.0;
1292        let mut last_err = 0.0;
1293        for i in 0..2000 {
1294            let mut active: Vec<(FeatureId, f64)> = Vec::with_capacity(5);
1295            for (j, &w) in true_w.iter().enumerate() {
1296                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1297                active.push((j as u64, x * w));
1298            }
1299            // Add some noise features with zero contribution.
1300            for k in 5..10 {
1301                let x = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1302                active.push((k as u64 + 100, x));
1303            }
1304            active.sort_by_key(|(id, _)| *id);
1305            let sf = SparseFeatures::from_sorted(active.clone()).unwrap();
1306            let y: f64 = active.iter().take(5).map(|(_, v)| v).sum();
1307            let pred = model.predict(&sf).unwrap();
1308            let err = (pred - y).abs();
1309            if i < 20 {
1310                first_err += err;
1311            }
1312            if i >= 1980 {
1313                last_err += err;
1314            }
1315            model.learn(&sf, y).unwrap();
1316        }
1317        assert!(
1318            last_err < first_err,
1319            "error should decrease in high-dim sparse"
1320        );
1321    }
1322
1323    // -----------------------------------------------------------------
1324    // FtrlRegressor: failure atomicity and overflow (ML-001/002)
1325    // -----------------------------------------------------------------
1326
1327    #[test]
1328    fn regressor_overflow_does_not_mutate_state() {
1329        // Finite inputs but intermediate `gradient * value` or
1330        // `gradient^2` overflows. The whole learn call must fail and
1331        // leave the model untouched.
1332        let mut model = FtrlRegressor::new(FtrlConfig {
1333            alpha: 0.1,
1334            beta: 1.0,
1335            l1: 0.0,
1336            l2: 0.0,
1337            max_features: None,
1338            new_feature_policy: NewFeaturePolicy::default(),
1339        })
1340        .unwrap();
1341        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1342        let result = model.learn(&sf, 1e100);
1343        assert!(result.is_err(), "expected overflow error, got {result:?}");
1344        assert_eq!(model.samples_seen(), 0);
1345        assert_eq!(model.feature_count(), 0);
1346        assert!(model.params.is_empty());
1347        assert_eq!(model.intercept.z, 0.0);
1348        assert_eq!(model.intercept.n, 0.0);
1349    }
1350
1351    #[test]
1352    fn regressor_partial_update_is_atomic() {
1353        // Two features in one sample. Feature 0 would succeed on its own,
1354        // feature 1 overflows. Neither may be committed.
1355        let mut model = FtrlRegressor::new(FtrlConfig {
1356            alpha: 0.1,
1357            beta: 1.0,
1358            l1: 0.0,
1359            l2: 0.0,
1360            max_features: None,
1361            new_feature_policy: NewFeaturePolicy::default(),
1362        })
1363        .unwrap();
1364        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e200)]).unwrap();
1365        assert!(model.learn(&sf, 1e100).is_err());
1366        // Feature 0 must NOT be inserted.
1367        assert!(!model.params.contains_key(&0));
1368        assert!(!model.params.contains_key(&1));
1369        assert_eq!(model.samples_seen(), 0);
1370    }
1371
1372    #[test]
1373    #[cfg(feature = "serde")]
1374    fn regressor_samples_seen_overflow_is_atomic() {
1375        let json = format!(
1376            "{{\"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\":{}}}",
1377            u64::MAX
1378        );
1379        let mut model: FtrlRegressor = serde_json::from_str(&json).unwrap();
1380        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1381        let result = model.learn(&sf, 1.0);
1382        assert!(result.is_err(), "expected counter overflow");
1383        assert_eq!(model.samples_seen(), u64::MAX);
1384        assert_eq!(model.feature_count(), 0);
1385        assert_eq!(model.intercept.z, 0.0);
1386        assert_eq!(model.intercept.n, 0.0);
1387    }
1388
1389    // -----------------------------------------------------------------
1390    // FtrlRegressor: max_features boundary (ML-003)
1391    // -----------------------------------------------------------------
1392
1393    #[test]
1394    fn regressor_max_features_reject_at_limit() {
1395        let mut model = FtrlRegressor::new(FtrlConfig {
1396            alpha: 0.5,
1397            beta: 1.0,
1398            l1: 0.0,
1399            l2: 0.0,
1400            max_features: Some(2),
1401            new_feature_policy: NewFeaturePolicy::Reject,
1402        })
1403        .unwrap();
1404        // Reach exactly the limit.
1405        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1406        model.learn(&sf, 1.0).unwrap();
1407        assert_eq!(model.feature_count(), 2);
1408        // One more new feature: Reject must fail atomically.
1409        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1410        assert!(model.learn(&sf_new, 1.0).is_err());
1411        assert_eq!(model.feature_count(), 2);
1412        assert_eq!(model.samples_seen(), 1);
1413        // Existing features still train.
1414        let sf_existing = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1415        model.learn(&sf_existing, 1.0).unwrap();
1416        assert_eq!(model.feature_count(), 2);
1417        assert_eq!(model.samples_seen(), 2);
1418    }
1419
1420    #[test]
1421    fn regressor_max_features_ignore_skips_new() {
1422        let mut model = FtrlRegressor::new(FtrlConfig {
1423            alpha: 0.5,
1424            beta: 1.0,
1425            l1: 0.0,
1426            l2: 0.0,
1427            max_features: Some(2),
1428            new_feature_policy: NewFeaturePolicy::Ignore,
1429        })
1430        .unwrap();
1431        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1432        model.learn(&sf, 1.0).unwrap();
1433        // Sample with one existing and one new feature. Under Ignore the
1434        // new feature is skipped, the existing one is updated, and the
1435        // counter still advances.
1436        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1437        model.learn(&sf_mixed, 1.0).unwrap();
1438        assert_eq!(model.feature_count(), 2);
1439        assert!(!model.params.contains_key(&2));
1440        assert_eq!(model.samples_seen(), 2);
1441    }
1442
1443    #[test]
1444    fn regressor_max_features_multi_new_prejudge() {
1445        let mut model = FtrlRegressor::new(FtrlConfig {
1446            alpha: 0.5,
1447            beta: 1.0,
1448            l1: 0.0,
1449            l2: 0.0,
1450            max_features: Some(2),
1451            new_feature_policy: NewFeaturePolicy::Reject,
1452        })
1453        .unwrap();
1454        // A single sample with three new features. Reject fails atomically
1455        // without inserting any subset.
1456        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1457        assert!(model.learn(&sf, 1.0).is_err());
1458        assert_eq!(model.feature_count(), 0);
1459        assert_eq!(model.samples_seen(), 0);
1460    }
1461
1462    // -----------------------------------------------------------------
1463    // FtrlRegressor: serde validation (ML-004)
1464    // -----------------------------------------------------------------
1465
1466    #[test]
1467    #[cfg(feature = "serde")]
1468    fn regressor_serde_rejects_negative_n() {
1469        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}";
1470        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1471        assert!(result.is_err(), "negative n must be rejected");
1472    }
1473
1474    #[test]
1475    #[cfg(feature = "serde")]
1476    fn regressor_serde_rejects_invalid_config() {
1477        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}";
1478        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
1479        assert!(result.is_err(), "invalid alpha must be rejected");
1480    }
1481
1482    #[test]
1483    #[cfg(feature = "serde")]
1484    fn regressor_serde_accepts_missing_optional_fields() {
1485        // Old state without max_features/new_feature_policy must still load.
1486        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}";
1487        let model: FtrlRegressor = serde_json::from_str(json).unwrap();
1488        assert!(model.config().max_features.is_none());
1489        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
1490    }
1491
1492    // -----------------------------------------------------------------
1493    // FtrlClassifier tests
1494    // -----------------------------------------------------------------
1495
1496    #[test]
1497    fn cold_start_returns_0_5() {
1498        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1499        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1500        let p = model.predict_proba(&sf).unwrap();
1501        assert!((p - 0.5).abs() < 1e-12, "cold start should predict 0.5");
1502    }
1503
1504    #[test]
1505    fn learn_separable_data() {
1506        // Linearly separable: class 1 when x0 > 0.
1507        let mut model = FtrlClassifier::new(FtrlConfig {
1508            alpha: 0.5,
1509            beta: 1.0,
1510            l1: 0.0,
1511            l2: 0.0,
1512            max_features: None,
1513            new_feature_policy: NewFeaturePolicy::default(),
1514        })
1515        .unwrap();
1516        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
1517        for _ in 0..1000 {
1518            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1519            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1520            let y = x0 > 0.0;
1521            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1522            model.learn(&sf, y).unwrap();
1523        }
1524        let p_pos = model
1525            .predict_proba(&SparseFeatures::from_sorted(vec![(0, 2.0), (1, 0.0)]).unwrap())
1526            .unwrap();
1527        let p_neg = model
1528            .predict_proba(&SparseFeatures::from_sorted(vec![(0, -2.0), (1, 0.0)]).unwrap())
1529            .unwrap();
1530        assert!(p_pos > 0.7, "p_pos should be high, got {p_pos}");
1531        assert!(p_neg < 0.3, "p_neg should be low, got {p_neg}");
1532    }
1533
1534    #[test]
1535    fn classifier_l1_produces_sparse_weights() {
1536        let mut model = FtrlClassifier::new(FtrlConfig {
1537            alpha: 0.1,
1538            beta: 1.0,
1539            l1: 100.0,
1540            l2: 0.0,
1541            max_features: None,
1542            new_feature_policy: NewFeaturePolicy::default(),
1543        })
1544        .unwrap();
1545        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(5);
1546        for _ in 0..200 {
1547            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1548            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1549            let y = x0 > 0.0;
1550            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1551            model.learn(&sf, y).unwrap();
1552        }
1553        let weights = model.weights();
1554        assert!(
1555            weights.is_empty(),
1556            "weights should all be zero with high L1, got {weights:?}"
1557        );
1558    }
1559
1560    #[test]
1561    fn classifier_dynamic_features() {
1562        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1563        assert_eq!(model.feature_count(), 0);
1564        let sf1 = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1565        model.learn(&sf1, true).unwrap();
1566        assert_eq!(model.feature_count(), 1);
1567        let sf2 = SparseFeatures::from_sorted(vec![(10, 1.0)]).unwrap();
1568        model.learn(&sf2, false).unwrap();
1569        assert_eq!(model.feature_count(), 2);
1570    }
1571
1572    #[test]
1573    fn classifier_predict_does_not_update_state() {
1574        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1575        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1576        let _ = model.predict_proba(&sf).unwrap();
1577        assert_eq!(model.samples_seen(), 0);
1578        assert_eq!(model.feature_count(), 0);
1579        model.learn(&sf, true).unwrap();
1580        let count = model.feature_count();
1581        let _ = model.predict_proba(&sf).unwrap();
1582        assert_eq!(model.feature_count(), count);
1583        assert_eq!(model.samples_seen(), 1);
1584    }
1585
1586    #[test]
1587    fn classifier_non_finite_value_rejected() {
1588        let model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1589        assert!(SparseFeatures::from_sorted(vec![(0, f64::NAN)]).is_err());
1590        assert!(SparseFeatures::from_sorted(vec![(0, f64::INFINITY)]).is_err());
1591        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1592        assert!(model.predict_proba(&sf).is_ok());
1593    }
1594
1595    #[test]
1596    fn classifier_empty_features_rejected() {
1597        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1598        let sf = SparseFeatures::new();
1599        assert!(model.predict_proba(&sf).is_err());
1600        assert!(model.learn(&sf, true).is_err());
1601    }
1602
1603    #[test]
1604    fn classifier_reset_clears_state() {
1605        let mut model = FtrlClassifier::new(FtrlConfig::default()).unwrap();
1606        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1607        model.learn(&sf, true).unwrap();
1608        model.learn(&sf, false).unwrap();
1609        assert_eq!(model.samples_seen(), 2);
1610        assert!(model.feature_count() > 0);
1611        model.reset();
1612        assert_eq!(model.samples_seen(), 0);
1613        assert_eq!(model.feature_count(), 0);
1614        let p = model.predict_proba(&sf).unwrap();
1615        assert!((p - 0.5).abs() < 1e-12);
1616    }
1617
1618    #[test]
1619    fn classifier_invalid_config_rejected() {
1620        assert!(
1621            FtrlClassifier::new(FtrlConfig {
1622                alpha: 0.0,
1623                ..FtrlConfig::default()
1624            })
1625            .is_err()
1626        );
1627        assert!(
1628            FtrlClassifier::new(FtrlConfig {
1629                beta: -0.1,
1630                ..FtrlConfig::default()
1631            })
1632            .is_err()
1633        );
1634        assert!(
1635            FtrlClassifier::new(FtrlConfig {
1636                l1: -1.0,
1637                ..FtrlConfig::default()
1638            })
1639            .is_err()
1640        );
1641        assert!(
1642            FtrlClassifier::new(FtrlConfig {
1643                l2: -1.0,
1644                ..FtrlConfig::default()
1645            })
1646            .is_err()
1647        );
1648        assert!(
1649            FtrlClassifier::new(FtrlConfig {
1650                alpha: f64::INFINITY,
1651                ..FtrlConfig::default()
1652            })
1653            .is_err()
1654        );
1655    }
1656
1657    #[test]
1658    #[cfg(feature = "serde")]
1659    fn classifier_serde_roundtrip() {
1660        let mut model = FtrlClassifier::new(FtrlConfig {
1661            alpha: 0.3,
1662            beta: 0.5,
1663            l1: 0.1,
1664            l2: 0.2,
1665            max_features: Some(100),
1666            new_feature_policy: NewFeaturePolicy::Reject,
1667        })
1668        .unwrap();
1669        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (2, -1.0)]).unwrap();
1670        model.learn(&sf, true).unwrap();
1671        model.learn(&sf, false).unwrap();
1672        let json = serde_json::to_string(&model).unwrap();
1673        let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
1674        assert_eq!(restored.samples_seen(), model.samples_seen());
1675        assert_eq!(restored.feature_count(), model.feature_count());
1676        let p1 = model.predict_proba(&sf).unwrap();
1677        let p2 = restored.predict_proba(&sf).unwrap();
1678        assert!((p1 - p2).abs() < 1e-12);
1679    }
1680
1681    #[test]
1682    fn predict_proba_in_range() {
1683        let mut model = FtrlClassifier::new(FtrlConfig {
1684            alpha: 0.5,
1685            beta: 1.0,
1686            l1: 0.0,
1687            l2: 0.0,
1688            max_features: None,
1689            new_feature_policy: NewFeaturePolicy::default(),
1690        })
1691        .unwrap();
1692        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(17);
1693        for _ in 0..200 {
1694            let x0 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1695            let x1 = rand::Rng::gen_range(&mut rng, -5.0..5.0);
1696            let y = x0 > 0.0;
1697            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1698            model.learn(&sf, y).unwrap();
1699            let p = model.predict_proba(&sf).unwrap();
1700            assert!(
1701                (0.0..=1.0).contains(&p),
1702                "probability must be in [0,1], got {p}"
1703            );
1704        }
1705    }
1706
1707    #[test]
1708    fn learn_improves_accuracy() {
1709        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(21);
1710        // Generate a fixed test set.
1711        let test_set: Vec<(SparseFeatures, bool)> = (0..100)
1712            .map(|_| {
1713                let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1714                let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1715                let y = x0 + x1 > 0.0;
1716                (
1717                    SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap(),
1718                    y,
1719                )
1720            })
1721            .collect();
1722
1723        let mut model = FtrlClassifier::new(FtrlConfig {
1724            alpha: 0.5,
1725            beta: 1.0,
1726            l1: 0.0,
1727            l2: 0.0,
1728            max_features: None,
1729            new_feature_policy: NewFeaturePolicy::default(),
1730        })
1731        .unwrap();
1732
1733        // Accuracy before learning (always predicts 0.5 -> threshold 0.5 -> true).
1734        let acc_before: f64 = test_set
1735            .iter()
1736            .map(|(sf, y)| {
1737                let pred = model.predict(sf).unwrap();
1738                if pred == *y { 1.0 } else { 0.0 }
1739            })
1740            .sum::<f64>()
1741            / test_set.len() as f64;
1742
1743        // Train on fresh data.
1744        for _ in 0..1000 {
1745            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1746            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1747            let y = x0 + x1 > 0.0;
1748            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1749            model.learn(&sf, y).unwrap();
1750        }
1751
1752        let acc_after: f64 = test_set
1753            .iter()
1754            .map(|(sf, y)| {
1755                let pred = model.predict(sf).unwrap();
1756                if pred == *y { 1.0 } else { 0.0 }
1757            })
1758            .sum::<f64>()
1759            / test_set.len() as f64;
1760
1761        assert!(
1762            acc_after > acc_before,
1763            "accuracy should improve: {acc_before} -> {acc_after}"
1764        );
1765    }
1766
1767    #[test]
1768    fn classifier_weights_returns_nonzero_only() {
1769        let mut model = FtrlClassifier::new(FtrlConfig {
1770            alpha: 0.5,
1771            beta: 1.0,
1772            l1: 0.0,
1773            l2: 0.0,
1774            max_features: None,
1775            new_feature_policy: NewFeaturePolicy::default(),
1776        })
1777        .unwrap();
1778        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 0.0001)]).unwrap();
1779        for _ in 0..50 {
1780            model.learn(&sf, true).unwrap();
1781        }
1782        let weights = model.weights();
1783        for &(_, w) in &weights {
1784            assert!(w != 0.0);
1785        }
1786    }
1787
1788    #[test]
1789    fn classifier_multiple_features() {
1790        let mut model = FtrlClassifier::new(FtrlConfig {
1791            alpha: 0.5,
1792            beta: 1.0,
1793            l1: 0.0,
1794            l2: 0.0,
1795            max_features: None,
1796            new_feature_policy: NewFeaturePolicy::default(),
1797        })
1798        .unwrap();
1799        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(33);
1800        for _ in 0..1000 {
1801            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1802            let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1803            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1804            // y = 1 if x0 + x1 > 0
1805            let y = x0 + x1 > 0.0;
1806            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1), (2, x2)]).unwrap();
1807            model.learn(&sf, y).unwrap();
1808        }
1809        let weights = model.weights();
1810        // Features 0 and 1 should have non-zero weights; feature 2 may or may not.
1811        assert!(weights.iter().any(|&(id, _)| id == 0));
1812        assert!(weights.iter().any(|&(id, _)| id == 1));
1813        // Verify prediction quality.
1814        let p_pos = model
1815            .predict_proba(
1816                &SparseFeatures::from_sorted(vec![(0, 3.0), (1, 3.0), (2, 0.0)]).unwrap(),
1817            )
1818            .unwrap();
1819        let p_neg = model
1820            .predict_proba(
1821                &SparseFeatures::from_sorted(vec![(0, -3.0), (1, -3.0), (2, 0.0)]).unwrap(),
1822            )
1823            .unwrap();
1824        assert!(p_pos > 0.8);
1825        assert!(p_neg < 0.2);
1826    }
1827
1828    #[test]
1829    fn log_loss_converges() {
1830        // Average log loss should decrease over training.
1831        let mut model = FtrlClassifier::new(FtrlConfig {
1832            alpha: 0.5,
1833            beta: 1.0,
1834            l1: 0.0,
1835            l2: 0.0,
1836            max_features: None,
1837            new_feature_policy: NewFeaturePolicy::default(),
1838        })
1839        .unwrap();
1840        let loss_fn = crate::loss::log_loss::BinaryLogLoss::new();
1841        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(55);
1842        let mut first_loss = 0.0;
1843        let mut last_loss = 0.0;
1844        for i in 0..1000 {
1845            let x0 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
1846            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
1847            let y = x0 > 0.0;
1848            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
1849            let p = model.predict_proba(&sf).unwrap();
1850            let loss = loss_fn.loss(p, y);
1851            if i < 20 {
1852                first_loss += loss;
1853            }
1854            if i >= 980 {
1855                last_loss += loss;
1856            }
1857            model.learn(&sf, y).unwrap();
1858        }
1859        assert!(last_loss < first_loss, "log loss should decrease");
1860    }
1861
1862    // -----------------------------------------------------------------
1863    // FtrlClassifier: failure atomicity and overflow (ML-001/002)
1864    // -----------------------------------------------------------------
1865
1866    #[test]
1867    fn classifier_overflow_does_not_mutate_state() {
1868        let mut model = FtrlClassifier::new(FtrlConfig {
1869            alpha: 0.1,
1870            beta: 1.0,
1871            l1: 0.0,
1872            l2: 0.0,
1873            max_features: None,
1874            new_feature_policy: NewFeaturePolicy::default(),
1875        })
1876        .unwrap();
1877        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1878        // Cold-start probability is 0.5, grad = 0.5 - 0.0 = 0.5.
1879        // g for feature 1 = 0.5 * 1e300 = 5e299 (finite), g^2 = 2.5e599 = inf.
1880        let result = model.learn(&sf, false);
1881        assert!(result.is_err(), "expected overflow error, got {result:?}");
1882        assert_eq!(model.samples_seen(), 0);
1883        assert_eq!(model.feature_count(), 0);
1884        assert!(model.params.is_empty());
1885        assert_eq!(model.intercept.z, 0.0);
1886        assert_eq!(model.intercept.n, 0.0);
1887    }
1888
1889    #[test]
1890    fn classifier_partial_update_is_atomic() {
1891        let mut model = FtrlClassifier::new(FtrlConfig {
1892            alpha: 0.1,
1893            beta: 1.0,
1894            l1: 0.0,
1895            l2: 0.0,
1896            max_features: None,
1897            new_feature_policy: NewFeaturePolicy::default(),
1898        })
1899        .unwrap();
1900        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1e300)]).unwrap();
1901        assert!(model.learn(&sf, false).is_err());
1902        assert!(!model.params.contains_key(&0));
1903        assert!(!model.params.contains_key(&1));
1904        assert_eq!(model.samples_seen(), 0);
1905    }
1906
1907    #[test]
1908    #[cfg(feature = "serde")]
1909    fn classifier_samples_seen_overflow_is_atomic() {
1910        let json = format!(
1911            "{{\"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\":{}}}",
1912            u64::MAX
1913        );
1914        let mut model: FtrlClassifier = serde_json::from_str(&json).unwrap();
1915        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
1916        let result = model.learn(&sf, true);
1917        assert!(result.is_err(), "expected counter overflow");
1918        assert_eq!(model.samples_seen(), u64::MAX);
1919        assert_eq!(model.feature_count(), 0);
1920        assert_eq!(model.intercept.z, 0.0);
1921        assert_eq!(model.intercept.n, 0.0);
1922    }
1923
1924    // -----------------------------------------------------------------
1925    // FtrlClassifier: max_features boundary (ML-003)
1926    // -----------------------------------------------------------------
1927
1928    #[test]
1929    fn classifier_max_features_reject_at_limit() {
1930        let mut model = FtrlClassifier::new(FtrlConfig {
1931            alpha: 0.5,
1932            beta: 1.0,
1933            l1: 0.0,
1934            l2: 0.0,
1935            max_features: Some(2),
1936            new_feature_policy: NewFeaturePolicy::Reject,
1937        })
1938        .unwrap();
1939        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1940        model.learn(&sf, true).unwrap();
1941        assert_eq!(model.feature_count(), 2);
1942        let sf_new = SparseFeatures::from_sorted(vec![(2, 3.0)]).unwrap();
1943        assert!(model.learn(&sf_new, true).is_err());
1944        assert_eq!(model.feature_count(), 2);
1945        assert_eq!(model.samples_seen(), 1);
1946    }
1947
1948    #[test]
1949    fn classifier_max_features_ignore_skips_new() {
1950        let mut model = FtrlClassifier::new(FtrlConfig {
1951            alpha: 0.5,
1952            beta: 1.0,
1953            l1: 0.0,
1954            l2: 0.0,
1955            max_features: Some(2),
1956            new_feature_policy: NewFeaturePolicy::Ignore,
1957        })
1958        .unwrap();
1959        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
1960        model.learn(&sf, true).unwrap();
1961        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (2, 3.0)]).unwrap();
1962        model.learn(&sf_mixed, false).unwrap();
1963        assert_eq!(model.feature_count(), 2);
1964        assert!(!model.params.contains_key(&2));
1965        assert_eq!(model.samples_seen(), 2);
1966    }
1967
1968    #[test]
1969    fn classifier_max_features_multi_new_prejudge() {
1970        let mut model = FtrlClassifier::new(FtrlConfig {
1971            alpha: 0.5,
1972            beta: 1.0,
1973            l1: 0.0,
1974            l2: 0.0,
1975            max_features: Some(2),
1976            new_feature_policy: NewFeaturePolicy::Reject,
1977        })
1978        .unwrap();
1979        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0), (2, 3.0)]).unwrap();
1980        assert!(model.learn(&sf, true).is_err());
1981        assert_eq!(model.feature_count(), 0);
1982        assert_eq!(model.samples_seen(), 0);
1983    }
1984
1985    // -----------------------------------------------------------------
1986    // FtrlClassifier: serde validation (ML-004)
1987    // -----------------------------------------------------------------
1988
1989    #[test]
1990    #[cfg(feature = "serde")]
1991    fn classifier_serde_rejects_negative_n() {
1992        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}";
1993        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
1994        assert!(result.is_err(), "negative n must be rejected");
1995    }
1996
1997    #[test]
1998    #[cfg(feature = "serde")]
1999    fn classifier_serde_rejects_invalid_config() {
2000        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}";
2001        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2002        assert!(result.is_err(), "invalid alpha must be rejected");
2003    }
2004
2005    #[test]
2006    #[cfg(feature = "serde")]
2007    fn classifier_serde_accepts_missing_optional_fields() {
2008        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}";
2009        let model: FtrlClassifier = serde_json::from_str(json).unwrap();
2010        assert!(model.config().max_features.is_none());
2011        assert_eq!(model.config().new_feature_policy, NewFeaturePolicy::Reject);
2012    }
2013
2014    // -----------------------------------------------------------------
2015    // Second-audit: FTRL underflow / zero-denominator / Ignore order
2016    // -----------------------------------------------------------------
2017
2018    #[test]
2019    fn regressor_gradient_squared_underflow_is_atomic() {
2020        // Cold start: prediction = 0, target = -1e-200 → grad = 1e-200.
2021        // For feature 0 (value=1.0): g = 1e-200 (non-zero, finite).
2022        // g^2 = 1e-400 underflows to 0.0. Without the explicit underflow
2023        // check, n_new would stay at 0 while z advances, producing a state
2024        // whose next predict() divides by zero.
2025        let mut model = FtrlRegressor::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, 1.0)]).unwrap();
2035        let result = model.learn(&sf, -1e-200);
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        assert!(model.params.is_empty());
2040        assert_eq!(model.intercept.z, 0.0);
2041        assert_eq!(model.intercept.n, 0.0);
2042    }
2043
2044    #[test]
2045    fn classifier_gradient_squared_underflow_is_atomic() {
2046        // Cold start: probability = 0.5, target = false (y=0), grad = 0.5.
2047        // Feature 0 value = 1e-200: g = 0.5 * 1e-200 = 5e-201 (non-zero).
2048        // g^2 = 2.5e-401 underflows to 0.0.
2049        let mut model = FtrlClassifier::new(FtrlConfig {
2050            alpha: 1.0,
2051            beta: 0.0,
2052            l1: 0.0,
2053            l2: 0.0,
2054            max_features: None,
2055            new_feature_policy: NewFeaturePolicy::default(),
2056        })
2057        .unwrap();
2058        let sf = SparseFeatures::from_sorted(vec![(0, 1e-200)]).unwrap();
2059        let result = model.learn(&sf, false);
2060        assert!(result.is_err(), "expected underflow error, got {result:?}");
2061        assert_eq!(model.samples_seen(), 0);
2062        assert_eq!(model.feature_count(), 0);
2063    }
2064
2065    #[test]
2066    fn regressor_boundary_config_predict_after_learn_always_finite() {
2067        // beta=0, l2=0, l1=0 is the zero-denominator boundary. Every
2068        // successful learn must keep the model in a state where predict
2069        // returns a finite value.
2070        let mut model = FtrlRegressor::new(FtrlConfig {
2071            alpha: 1.0,
2072            beta: 0.0,
2073            l1: 0.0,
2074            l2: 0.0,
2075            max_features: None,
2076            new_feature_policy: NewFeaturePolicy::default(),
2077        })
2078        .unwrap();
2079        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(77);
2080        for _ in 0..100 {
2081            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2082            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2083            let y = 2.0 * x0 - x1;
2084            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2085            model.learn(&sf, y).unwrap();
2086            let pred = model.predict(&sf);
2087            assert!(
2088                pred.is_ok(),
2089                "predict failed after successful learn: {pred:?}"
2090            );
2091            assert!(
2092                pred.unwrap().is_finite(),
2093                "predict must return finite value after successful learn"
2094            );
2095        }
2096    }
2097
2098    #[test]
2099    fn classifier_boundary_config_predict_proba_after_learn_always_finite() {
2100        let mut model = FtrlClassifier::new(FtrlConfig {
2101            alpha: 1.0,
2102            beta: 0.0,
2103            l1: 0.0,
2104            l2: 0.0,
2105            max_features: None,
2106            new_feature_policy: NewFeaturePolicy::default(),
2107        })
2108        .unwrap();
2109        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(88);
2110        for _ in 0..100 {
2111            let x0 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2112            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
2113            let y = x0 > 0.0;
2114            let sf = SparseFeatures::from_sorted(vec![(0, x0), (1, x1)]).unwrap();
2115            model.learn(&sf, y).unwrap();
2116            let proba = model.predict_proba(&sf);
2117            assert!(proba.is_ok(), "predict_proba failed after learn: {proba:?}");
2118            let p = proba.unwrap();
2119            assert!(p.is_finite(), "probability must be finite, got {p}");
2120            assert!(
2121                (0.0..=1.0).contains(&p),
2122                "probability must be in [0,1], got {p}"
2123            );
2124        }
2125    }
2126
2127    #[test]
2128    #[cfg(feature = "serde")]
2129    fn regressor_serde_rejects_n_zero_z_nonzero() {
2130        // n=0, z=1.0: weight formula denominator = l2 + (beta + sqrt(0))/alpha.
2131        // With beta=0, l2=0, alpha=1: denominator = 0 → weight = inf.
2132        // validate() must reject this state.
2133        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}";
2134        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2135        assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2136    }
2137
2138    #[test]
2139    #[cfg(feature = "serde")]
2140    fn classifier_serde_rejects_n_zero_z_nonzero() {
2141        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}";
2142        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2143        assert!(result.is_err(), "n=0 && z!=0 must be rejected");
2144    }
2145
2146    #[test]
2147    #[cfg(feature = "serde")]
2148    fn regressor_predict_dot_plus_intercept_overflow() {
2149        // Craft a model where dot + intercept overflows f64.
2150        // config: alpha=1, beta=0, l1=0, l2=0 → weight = -z / sqrt(n)
2151        // z = -f64::MAX * 0.75, n = 1.0 → weight = f64::MAX * 0.75
2152        // dot = weight * 1.0 = f64::MAX * 0.75
2153        // intercept_weight = f64::MAX * 0.75
2154        // dot + intercept = f64::MAX * 1.5 → overflow to inf.
2155        let z = -f64::MAX * 0.75;
2156        let json = format!(
2157            "{{\"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}}",
2158            z
2159        );
2160        let model: FtrlRegressor = serde_json::from_str(&json).unwrap();
2161        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2162        let result = model.predict(&sf);
2163        assert!(
2164            result.is_err(),
2165            "expected dot+intercept overflow error, got {result:?}"
2166        );
2167    }
2168
2169    #[test]
2170    fn regressor_ignore_skips_overflowing_new_feature() {
2171        let mut model = FtrlRegressor::new(FtrlConfig {
2172            alpha: 0.5,
2173            beta: 1.0,
2174            l1: 0.0,
2175            l2: 0.0,
2176            max_features: Some(1),
2177            new_feature_policy: NewFeaturePolicy::Ignore,
2178        })
2179        .unwrap();
2180        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2181        model.learn(&sf, 1.0).unwrap();
2182        assert_eq!(model.feature_count(), 1);
2183
2184        // New feature 1 with huge value. grad is O(1), so
2185        // g = grad * f64::MAX is finite, but g^2 overflows to inf.
2186        // Under Ignore, the new feature must be skipped BEFORE the
2187        // multiplication; otherwise the whole learn() fails.
2188        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2189        let result = model.learn(&sf_mixed, 1.0);
2190        assert!(
2191            result.is_ok(),
2192            "Ignore must skip overflowing new feature, got {result:?}"
2193        );
2194        assert_eq!(model.feature_count(), 1);
2195        assert!(!model.params.contains_key(&1));
2196        assert_eq!(model.samples_seen(), 2);
2197        assert!(model.predict(&sf).is_ok());
2198    }
2199
2200    #[test]
2201    fn classifier_ignore_skips_overflowing_new_feature() {
2202        let mut model = FtrlClassifier::new(FtrlConfig {
2203            alpha: 0.5,
2204            beta: 1.0,
2205            l1: 0.0,
2206            l2: 0.0,
2207            max_features: Some(1),
2208            new_feature_policy: NewFeaturePolicy::Ignore,
2209        })
2210        .unwrap();
2211        let sf = SparseFeatures::from_sorted(vec![(0, 1.0)]).unwrap();
2212        model.learn(&sf, true).unwrap();
2213        assert_eq!(model.feature_count(), 1);
2214
2215        let sf_mixed = SparseFeatures::from_sorted(vec![(0, 1.0), (1, f64::MAX)]).unwrap();
2216        let result = model.learn(&sf_mixed, true);
2217        assert!(
2218            result.is_ok(),
2219            "Ignore must skip overflowing new feature, got {result:?}"
2220        );
2221        assert_eq!(model.feature_count(), 1);
2222        assert!(!model.params.contains_key(&1));
2223        assert_eq!(model.samples_seen(), 2);
2224        assert!(model.predict_proba(&sf).is_ok());
2225    }
2226
2227    // -----------------------------------------------------------------
2228    // §6.2: FTRL config-aware serde validation
2229    // -----------------------------------------------------------------
2230
2231    #[test]
2232    #[cfg(feature = "serde")]
2233    fn regressor_serde_rejects_config_dependent_zero_denominator() {
2234        // alpha = f64::MAX, beta = 0, l2 = 0, l1 = 0, n = 1e-300, z = 1.0.
2235        // sqrt(n) = 1e-150; denominator = (0 + 1e-150) / f64::MAX underflows
2236        // to 0.0. The basic validate() passes (n > 0, z finite) but
2237        // weight_checked must reject the zero denominator explicitly.
2238        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}";
2239        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2240        assert!(
2241            result.is_err(),
2242            "config-dependent zero denominator must be rejected"
2243        );
2244    }
2245
2246    #[test]
2247    #[cfg(feature = "serde")]
2248    fn classifier_serde_rejects_config_dependent_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\":{\"0\":{\"z\":1.0,\"n\":1e-300}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":0}";
2250        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2251        assert!(
2252            result.is_err(),
2253            "config-dependent zero denominator must be rejected"
2254        );
2255    }
2256
2257    #[test]
2258    #[cfg(feature = "serde")]
2259    fn regressor_serde_rejects_intercept_zero_denominator() {
2260        // Same dangerous config, but the malicious state is in the
2261        // intercept rather than a feature param.
2262        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}";
2263        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2264        assert!(
2265            result.is_err(),
2266            "intercept zero denominator must be rejected"
2267        );
2268    }
2269
2270    #[test]
2271    #[cfg(feature = "serde")]
2272    fn classifier_serde_rejects_intercept_zero_denominator() {
2273        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}";
2274        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2275        assert!(
2276            result.is_err(),
2277            "intercept zero denominator must be rejected"
2278        );
2279    }
2280
2281    #[test]
2282    #[cfg(feature = "serde")]
2283    fn regressor_valid_boundary_state_roundtrips() {
2284        // Boundary config (beta=0, l2=0) with legitimately trained state
2285        // must round-trip through serde without being rejected by the new
2286        // config-aware checks.
2287        let mut model = FtrlRegressor::new(FtrlConfig {
2288            alpha: 1.0,
2289            beta: 0.0,
2290            l1: 0.0,
2291            l2: 0.0,
2292            max_features: None,
2293            new_feature_policy: NewFeaturePolicy::default(),
2294        })
2295        .unwrap();
2296        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2297        model.learn(&sf, 3.0).unwrap();
2298        model.learn(&sf, 5.0).unwrap();
2299        let json = serde_json::to_string(&model).unwrap();
2300        let restored: FtrlRegressor = serde_json::from_str(&json).unwrap();
2301        assert_eq!(restored.samples_seen(), model.samples_seen());
2302        assert_eq!(restored.feature_count(), model.feature_count());
2303        let p1 = model.predict(&sf).unwrap();
2304        let p2 = restored.predict(&sf).unwrap();
2305        assert!((p1 - p2).abs() < 1e-12);
2306        // weights() must not produce Infinity on the restored model.
2307        for (_, w) in restored.weights() {
2308            assert!(w.is_finite(), "restored weight must be finite, got {w}");
2309        }
2310        assert!(restored.intercept().is_finite());
2311    }
2312
2313    #[test]
2314    #[cfg(feature = "serde")]
2315    fn classifier_valid_boundary_state_roundtrips() {
2316        let mut model = FtrlClassifier::new(FtrlConfig {
2317            alpha: 1.0,
2318            beta: 0.0,
2319            l1: 0.0,
2320            l2: 0.0,
2321            max_features: None,
2322            new_feature_policy: NewFeaturePolicy::default(),
2323        })
2324        .unwrap();
2325        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 2.0)]).unwrap();
2326        model.learn(&sf, true).unwrap();
2327        model.learn(&sf, false).unwrap();
2328        let json = serde_json::to_string(&model).unwrap();
2329        let restored: FtrlClassifier = serde_json::from_str(&json).unwrap();
2330        assert_eq!(restored.samples_seen(), model.samples_seen());
2331        assert_eq!(restored.feature_count(), model.feature_count());
2332        let p1 = model.predict_proba(&sf).unwrap();
2333        let p2 = restored.predict_proba(&sf).unwrap();
2334        assert!((p1 - p2).abs() < 1e-12);
2335        for (_, w) in restored.weights() {
2336            assert!(w.is_finite(), "restored weight must be finite, got {w}");
2337        }
2338        assert!(restored.intercept().is_finite());
2339    }
2340
2341    // -----------------------------------------------------------------
2342    // Fourth-stage audit: FTRL max_features serde invariant (4-A-02)
2343    // -----------------------------------------------------------------
2344
2345    #[test]
2346    #[cfg(feature = "serde")]
2347    fn regressor_serde_rejects_params_above_max_features() {
2348        // config.max_features = 1, but two stored params. This violates
2349        // the model contract and must be rejected at deserialisation.
2350        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":1,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1.0},\"1\":{\"z\":1.0,\"n\":1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":1}";
2351        let result: Result<FtrlRegressor, _> = serde_json::from_str(json);
2352        let err = match result {
2353            Ok(_) => panic!("expected serde error, got Ok"),
2354            Err(e) => e,
2355        };
2356        let msg = err.to_string();
2357        assert!(
2358            msg.contains("max_features") && msg.contains("feature count"),
2359            "error must mention feature count / max_features, got: {msg}"
2360        );
2361    }
2362
2363    #[test]
2364    #[cfg(feature = "serde")]
2365    fn classifier_serde_rejects_params_above_max_features() {
2366        let json = "{\"config\":{\"alpha\":0.1,\"beta\":1.0,\"l1\":1.0,\"l2\":1.0,\"max_features\":1,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":1.0,\"n\":1.0},\"1\":{\"z\":1.0,\"n\":1.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":1}";
2367        let result: Result<FtrlClassifier, _> = serde_json::from_str(json);
2368        let err = match result {
2369            Ok(_) => panic!("expected serde error, got Ok"),
2370            Err(e) => e,
2371        };
2372        let msg = err.to_string();
2373        assert!(
2374            msg.contains("max_features") && msg.contains("feature count"),
2375            "error must mention feature count / max_features, got: {msg}"
2376        );
2377    }
2378
2379    #[test]
2380    #[cfg(feature = "serde")]
2381    fn regressor_serde_accepts_params_equal_to_max_features() {
2382        // params.len() == max_features is legal. Round-trip must succeed
2383        // and weights()/predict() must remain finite.
2384        let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":2,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":3}";
2385        let model: FtrlRegressor =
2386            serde_json::from_str(json).expect("equal count must be accepted");
2387        assert_eq!(model.feature_count(), 2);
2388        assert_eq!(model.samples_seen(), 3);
2389        for (_, w) in model.weights() {
2390            assert!(w.is_finite(), "weight must be finite, got {w}");
2391        }
2392        assert!(model.intercept().is_finite());
2393        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2394        let pred = model.predict(&sf).expect("predict must succeed");
2395        assert!(pred.is_finite(), "prediction must be finite, got {pred}");
2396    }
2397
2398    #[test]
2399    #[cfg(feature = "serde")]
2400    fn classifier_serde_accepts_params_equal_to_max_features() {
2401        let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":2,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":3}";
2402        let model: FtrlClassifier =
2403            serde_json::from_str(json).expect("equal count must be accepted");
2404        assert_eq!(model.feature_count(), 2);
2405        assert_eq!(model.samples_seen(), 3);
2406        for (_, w) in model.weights() {
2407            assert!(w.is_finite(), "weight must be finite, got {w}");
2408        }
2409        assert!(model.intercept().is_finite());
2410        let sf = SparseFeatures::from_sorted(vec![(0, 1.0), (1, 1.0)]).unwrap();
2411        let p = model
2412            .predict_proba(&sf)
2413            .expect("predict_proba must succeed");
2414        assert!(p.is_finite(), "probability must be finite, got {p}");
2415        assert!(
2416            (0.0..=1.0).contains(&p),
2417            "probability must be in [0,1], got {p}"
2418        );
2419    }
2420
2421    #[test]
2422    #[cfg(feature = "serde")]
2423    fn regressor_serde_allows_unbounded_params_when_max_features_none() {
2424        // max_features = None must not constrain params.len(). Three
2425        // stored params with no cap must round-trip cleanly.
2426        let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0},\"2\":{\"z\":1.5,\"n\":3.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":6}";
2427        let model: FtrlRegressor =
2428            serde_json::from_str(json).expect("unbounded state must be accepted");
2429        assert_eq!(model.feature_count(), 3);
2430        for (_, w) in model.weights() {
2431            assert!(w.is_finite(), "weight must be finite, got {w}");
2432        }
2433        assert!(model.intercept().is_finite());
2434    }
2435
2436    #[test]
2437    #[cfg(feature = "serde")]
2438    fn classifier_serde_allows_unbounded_params_when_max_features_none() {
2439        let json = "{\"config\":{\"alpha\":0.5,\"beta\":1.0,\"l1\":0.0,\"l2\":0.0,\"max_features\":null,\"new_feature_policy\":\"Reject\"},\"params\":{\"0\":{\"z\":0.5,\"n\":1.0},\"1\":{\"z\":-0.25,\"n\":2.0},\"2\":{\"z\":1.5,\"n\":3.0}},\"intercept\":{\"z\":0.0,\"n\":0.0},\"samples_seen\":6}";
2440        let model: FtrlClassifier =
2441            serde_json::from_str(json).expect("unbounded state must be accepted");
2442        assert_eq!(model.feature_count(), 3);
2443        for (_, w) in model.weights() {
2444            assert!(w.is_finite(), "weight must be finite, got {w}");
2445        }
2446        assert!(model.intercept().is_finite());
2447    }
2448}