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