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