Skip to main content

gam_models/bms/
deviation_runtime.rs

1use crate::cubic_cell_kernel as exact_kernel;
2use crate::util::span::span_index_for_breakpoints;
3use gam_linalg::faer_ndarray::{FaerEigh, fast_ab};
4use gam_solve::pirls::LinearInequalityConstraints;
5use gam_terms::basis::create_ispline_derivative_dense;
6use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
7
8/// Require a breakpoint sequence suitable for BMS span lookup: finite,
9/// strictly increasing, and long enough to define at least one span.
10fn validate_breakpoints(breakpoints: &[f64], label: &str) -> Result<(), String> {
11    if breakpoints.len() < 2 {
12        return Err(format!("{label} requires at least two breakpoints"));
13    }
14    if let Some((idx, window)) = breakpoints.windows(2).enumerate().find(|(_, window)| {
15        !window[0].is_finite() || !window[1].is_finite() || window[0] >= window[1]
16    }) {
17        return Err(format!(
18            "{label} requires strictly increasing finite breakpoints; breakpoints[{idx}]={:.6}, breakpoints[{}]={:.6}",
19            window[0],
20            idx + 1,
21            window[1]
22        ));
23    }
24    Ok::<(), _>(())
25}
26
27/// Deduplicate an ordered BMS knot sequence into strictly increasing
28/// breakpoints.
29fn breakpoints_from_knots(knots: &[f64], label: &str) -> Result<Vec<f64>, String> {
30    let mut breakpoints = Vec::new();
31    for &knot in knots {
32        if breakpoints
33            .last()
34            .is_none_or(|prev: &f64| (knot - *prev).abs() > 1e-12)
35        {
36            breakpoints.push(knot);
37        }
38    }
39    validate_breakpoints(&breakpoints, label)?;
40    Ok(breakpoints)
41}
42
43/// Round-off tolerance on the minimum monotonicity-derivative slack. The
44/// constraints are constructed with a positive required margin
45/// (`monotonicity_eps`); this separate, tiny negative bound only absorbs the
46/// finite-precision accumulation in evaluating the slack at the I-spline
47/// breakpoints, so a coefficient that is feasible up to a few ulps is not
48/// spuriously rejected. Anything more negative is a genuine violation.
49pub(crate) const MONOTONICITY_SLACK_ROUNDOFF_TOL: f64 = -1e-10;
50
51/// Typed errors emitted by the deviation runtime construction and evaluation
52/// helpers in this module.
53///
54/// Each variant carries a pre-formatted `reason` string so `Display` is
55/// byte-equivalent to the original `format!(...)` outputs the module used
56/// before the typed-error migration. The category split lets callers
57/// pattern-match on the failure kind without parsing the message.
58#[derive(Debug, Clone)]
59pub enum DeviationRuntimeError {
60    /// A scalar configuration value, index, derivative order, runtime value,
61    /// or required metadata bundle did not satisfy the contract (out-of-range
62    /// index, non-finite value, missing support points, span width <= 0).
63    InvalidInput { reason: String },
64    /// A matrix / vector shape did not match an expected dimension while
65    /// composing transforms, validating anchors, or accepting beta vectors.
66    DimensionMismatch { reason: String },
67    /// A numerical kernel (eigendecomposition, I-spline construction,
68    /// monotonicity slack search) failed or produced no usable output.
69    NumericalFailure { reason: String },
70}
71
72impl_reason_error_boilerplate! {
73    DeviationRuntimeError {
74        InvalidInput,
75        DimensionMismatch,
76        NumericalFailure,
77    }
78}
79
80/// Installed cross-block flex block on the runtime.
81///
82/// Direct on-runtime image of `identifiability::families::compiler::CompiledBlock`:
83/// `anchor_correction` = `compiled.anchor_correction` (the d × k matrix M),
84/// `anchor_components` = the per-anchor predict-time tags (the parent
85/// predictor uses them to rebuild `n_row` at predict-time rows). The
86/// post-residualisation row evaluator is
87///
88///   design_row(x) = pure_span_row(x) − n_row · M
89///
90/// The compiler bakes the orthonormalising rotation into M, so no
91/// separate rotation matrix is stored on the install state.
92#[derive(Clone, Debug)]
93pub struct InstalledFlexBlock {
94    /// Anchor correction matrix `M ∈ R^{d × k}` from
95    /// `CompiledBlock::anchor_correction`. The design evaluator subtracts
96    /// `n_row · M` per row.
97    pub anchor_correction: Array2<f64>,
98    /// Per-anchor predict-time tags, in the order the anchors were stacked
99    /// (parametric before flex). `sum(ncols)` equals the row dimension of
100    /// `anchor_correction`.
101    pub anchor_components: Vec<AnchorComponentTag>,
102}
103
104#[derive(Clone, Debug)]
105pub enum AnchorComponentTag {
106    /// Parametric anchor — at predict time the parent predictor reconstructs
107    /// the per-row vector from the saved marginal/logslope blocks; the
108    /// runtime only needs to know which block and how many columns. The
109    /// `block` tag is consumed by the serde plumbing in
110    /// `inference::model::SavedAnchorComponent`.
111    Parametric {
112        block: ParametricAnchorBlock,
113        ncols: usize,
114    },
115    /// Flex-evaluation anchor — a sibling flex block's design at training
116    /// rows (post-reparameterisation, in the same coordinate frame the
117    /// predictor will use at predict time). The number of columns equals
118    /// the sibling block's reparameterised basis dimension.
119    FlexEvaluation { ncols: usize },
120}
121
122#[derive(Clone, Copy, Debug, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
123pub enum ParametricAnchorBlock {
124    Marginal,
125    Logslope,
126}
127
128pub(crate) fn integrate_polynomial_product(left: &[f64], right: &[f64], width: f64) -> f64 {
129    let mut total = 0.0;
130    for (left_power, &left_coeff) in left.iter().enumerate() {
131        for (right_power, &right_coeff) in right.iter().enumerate() {
132            let power = left_power + right_power + 1;
133            total += left_coeff * right_coeff * width.powi(power as i32) / power as f64;
134        }
135    }
136    total
137}
138
139/// Precomputed per-span polynomial coefficient matrices for a structurally
140/// monotone anchored deviation basis.
141///
142/// Raw coefficients are monotone I-spline coefficients. The deviation
143/// derivative `w'(x)` is a nonnegative quadratic B-spline combination, so
144/// `w(x)` is a cubic I-spline combination with `C2` continuity at knots and
145/// constant tails. Zero coefficients still mean the identity map. The fitted
146/// coefficients live in the configured moment-anchor nullspace and are mapped
147/// back to these raw coefficients for monotonicity.
148///
149/// Monotonicity of the full transform `x + w(x)` is enforced by lower bounds
150/// on each span's quadratic Bernstein controls for `w'(x)`.
151#[derive(Clone, Debug)]
152pub struct DeviationRuntime {
153    pub(crate) degree: usize,
154    pub(crate) value_span_degree: usize,
155    pub(crate) basis_dim: usize,
156    pub(crate) monotonicity_eps: f64,
157    pub(crate) endpoint_points: Array1<f64>,
158    pub(crate) span_c0: Array2<f64>,
159    pub(crate) span_c1: Array2<f64>,
160    pub(crate) span_c2: Array2<f64>,
161    pub(crate) span_c3: Array2<f64>,
162    pub(crate) monotonicity_constraint_rows: Array2<f64>,
163    /// Deviation basis values at the rightmost breakpoint (1 × basis_dim).
164    /// Used for constant-tail continuation outside support: the deviation
165    /// saturates at this value for all z > right endpoint.
166    pub(crate) right_boundary_value_row: Array1<f64>,
167    /// Cross-block installed flex block. `None` until
168    /// `install_compiled_flex_block` is called.
169    pub(crate) installed_flex_block: Option<InstalledFlexBlock>,
170    /// Stacked parametric-anchor rows at training rows (n × d). Used by
171    /// `design_at_training_with_residual` to rebuild `block.design` after
172    /// orthogonalisation. Dropped before serialisation; predict-time
173    /// reconstruction rebuilds anchor rows fresh at the predict-time
174    /// feature rows.
175    pub(crate) anchor_rows_at_training: Option<Array2<f64>>,
176}
177
178/// Build the integrated derivative penalty matrix `P` on the *raw* I-spline
179/// coefficients (before any null-space transform), where
180/// `P_{ij} = ∫ b_i^(k)(x) b_j^(k)(x) dx` integrated piecewise over the knot
181/// support. The null space of `P` is the function-space null space of the
182/// k-th-derivative penalty: polynomials of degree < k. For k = 1 this is
183/// {constants}; for k = 2 it is {constants, linears}; for k = 3 it is
184/// {constants, linears, quadratics}. Dropping these directions from the
185/// basis at construction time is what gives the link-deviation block
186/// β-independent identifiability (the location block's intercept and any
187/// unpenalized location-linear absorb constants/linears in η; β_dev contains
188/// only the wiggle).
189///
190/// Mirrors `integrated_derivative_penalty_with_nullity` but operates on the
191/// raw cubic span coefficients, so it can be evaluated *before* the basis
192/// transform `Z` is constructed (which is what we need to compute `Z`
193/// itself).
194pub(crate) fn raw_integrated_derivative_penalty(
195    endpoint_points: &Array1<f64>,
196    raw_span_c0: &Array2<f64>,
197    raw_span_c1: &Array2<f64>,
198    raw_span_c2: &Array2<f64>,
199    raw_span_c3: &Array2<f64>,
200    derivative_order: usize,
201) -> Result<Array2<f64>, String> {
202    let raw_dim = raw_span_c0.ncols();
203    let n_spans = endpoint_points.len().saturating_sub(1);
204    if raw_span_c1.ncols() != raw_dim
205        || raw_span_c2.ncols() != raw_dim
206        || raw_span_c3.ncols() != raw_dim
207    {
208        return Err("raw smoothness penalty: span coefficient column dimensions disagree".into());
209    }
210    let mut penalty = Array2::<f64>::zeros((raw_dim, raw_dim));
211    for span_idx in 0..n_spans {
212        let left = endpoint_points[span_idx];
213        let right = endpoint_points[span_idx + 1];
214        let width = right - left;
215        if !width.is_finite() || width <= 0.0 {
216            return Err(format!(
217                "raw smoothness penalty span {span_idx} has invalid width {width}"
218            ));
219        }
220        for i in 0..raw_dim {
221            let ci = raw_span_derivative_polynomial_coefficients(
222                span_idx,
223                i,
224                derivative_order,
225                raw_span_c0,
226                raw_span_c1,
227                raw_span_c2,
228                raw_span_c3,
229            );
230            for j in i..raw_dim {
231                let cj = raw_span_derivative_polynomial_coefficients(
232                    span_idx,
233                    j,
234                    derivative_order,
235                    raw_span_c0,
236                    raw_span_c1,
237                    raw_span_c2,
238                    raw_span_c3,
239                );
240                let contribution = integrate_polynomial_product(&ci, &cj, width);
241                penalty[[i, j]] += contribution;
242                if i != j {
243                    penalty[[j, i]] += contribution;
244                }
245            }
246        }
247    }
248    Ok(penalty)
249}
250
251/// Per-span polynomial coefficients of the `derivative_order`-th derivative
252/// of raw basis function `basis_idx` on its parametric coordinate `t`. Mirrors
253/// `DeviationRuntime::span_derivative_polynomial_coefficients` but on raw
254/// coefficients so it's callable before `Z` exists.
255pub(crate) fn raw_span_derivative_polynomial_coefficients(
256    span_idx: usize,
257    basis_idx: usize,
258    derivative_order: usize,
259    raw_span_c0: &Array2<f64>,
260    raw_span_c1: &Array2<f64>,
261    raw_span_c2: &Array2<f64>,
262    raw_span_c3: &Array2<f64>,
263) -> Vec<f64> {
264    let c0 = raw_span_c0[[span_idx, basis_idx]];
265    let c1 = raw_span_c1[[span_idx, basis_idx]];
266    let c2 = raw_span_c2[[span_idx, basis_idx]];
267    let c3 = raw_span_c3[[span_idx, basis_idx]];
268    match derivative_order {
269        0 => vec![c0, c1, c2, c3],
270        1 => vec![c1, 2.0 * c2, 3.0 * c3],
271        2 => vec![2.0 * c2, 6.0 * c3],
272        3 => vec![6.0 * c3],
273        _ => Vec::new(),
274    }
275}
276
277/// Compute `Z` = orthonormal columns spanning the orthogonal complement of
278/// the null space of `P_raw` (the integrated derivative penalty in raw
279/// coordinates). Eigenvectors with strictly-positive eigenvalues are taken;
280/// near-zero eigenvalues correspond to functions with zero `derivative_order`-
281/// th derivative, i.e., polynomials of degree `< derivative_order` evaluated
282/// in the raw basis.
283///
284/// Returned `Z` has shape `raw_dim × (raw_dim − nullity)`. After applying it
285/// (`raw_basis · Z`), the transformed basis cannot represent any polynomial
286/// of degree < `derivative_order` — that direction is structurally absent
287/// from the parameterization. This is the β-independent identifiability
288/// constraint that replaces the data-distribution-dependent moment anchor.
289pub(crate) fn smoothness_nullspace_orthogonal_complement(
290    raw_penalty: &Array2<f64>,
291) -> Result<Array2<f64>, String> {
292    let n = raw_penalty.nrows();
293    if raw_penalty.ncols() != n {
294        return Err("smoothness penalty matrix must be square for null-space drop".to_string());
295    }
296    let (eigenvalues, eigenvectors) = raw_penalty
297        .eigh(faer::Side::Lower)
298        .map_err(|e| format!("raw smoothness penalty eigendecomposition failed: {e}"))?;
299    let evals = eigenvalues
300        .as_slice()
301        .ok_or_else(|| "raw smoothness penalty eigenvalues are not contiguous".to_string())?;
302    let threshold =
303        gam_solve::estimate::reml::reml_outer_engine::positive_eigenvalue_threshold(evals);
304    let kept: Vec<usize> = evals
305        .iter()
306        .enumerate()
307        .filter_map(|(i, &v)| (v > threshold).then_some(i))
308        .collect();
309    if kept.is_empty() {
310        return Err(
311            "smoothness penalty has no positive eigenvalues; basis is entirely in the penalty's \
312             null space and cannot be identified after the smoothness null-space drop"
313                .to_string(),
314        );
315    }
316    if kept.len() == n {
317        return Err(
318            "smoothness penalty has no null directions; nothing to drop. The link-deviation \
319             basis was expected to carry a non-trivial null space (constants/linears) for \
320             absorption by the location block — check the configured penalty derivative order"
321                .to_string(),
322        );
323    }
324    let mut z = Array2::<f64>::zeros((n, kept.len()));
325    for (col_out, &col_in) in kept.iter().enumerate() {
326        z.column_mut(col_out).assign(&eigenvectors.column(col_in));
327    }
328    Ok(z)
329}
330
331pub(crate) fn build_quadratic_derivative_bernstein_constraints(
332    endpoint_points: &Array1<f64>,
333    span_c1: &Array2<f64>,
334    span_c2: &Array2<f64>,
335    span_c3: &Array2<f64>,
336) -> Result<Array2<f64>, String> {
337    let n_spans = endpoint_points.len().saturating_sub(1);
338    let basis_dim = span_c1.ncols();
339    let mut rows = Array2::<f64>::zeros((3 * n_spans, basis_dim));
340    for span_idx in 0..n_spans {
341        let width = endpoint_points[span_idx + 1] - endpoint_points[span_idx];
342        if !width.is_finite() || width <= 0.0 {
343            return Err(DeviationRuntimeError::InvalidInput {
344                reason: format!(
345                    "DeviationRuntime monotonicity span {span_idx} has invalid width {width}"
346                ),
347            }
348            .into());
349        }
350        let left_row = 3 * span_idx;
351        let mid_row = left_row + 1;
352        let right_row = left_row + 2;
353        for basis_idx in 0..basis_dim {
354            let c1 = span_c1[[span_idx, basis_idx]];
355            let c2 = span_c2[[span_idx, basis_idx]];
356            let c3 = span_c3[[span_idx, basis_idx]];
357            // For w(t)=c0+c1*t+c2*t^2+c3*t^3 on t in [0,h],
358            // w'(t)=c1+2*c2*t+3*c3*t^2. In quadratic Bernstein form over
359            // s=t/h, the controls are:
360            //   b0 = c1
361            //   b1 = c1 + c2*h
362            //   b2 = c1 + 2*c2*h + 3*c3*h^2
363            // Since Bernstein basis functions are non-negative and sum to 1,
364            // b_k >= eps-1 is a linear certificate for x + w(x) monotonicity.
365            // `exact_monotonicity_min_slack` below still checks the true
366            // quadratic minimum, including the interior vertex.
367            rows[[left_row, basis_idx]] = c1;
368            rows[[mid_row, basis_idx]] = c1 + c2 * width;
369            rows[[right_row, basis_idx]] = c1 + 2.0 * c2 * width + 3.0 * c3 * width * width;
370        }
371    }
372    Ok(rows)
373}
374
375impl DeviationRuntime {
376    /// Rehydrate the exact post-compilation cubic tables carried by a saved
377    /// model for likelihood replay.
378    ///
379    /// This is deliberately not a spline constructor: rebuilding a runtime
380    /// from knots would rerun rank selection and cross-block
381    /// orthogonalisation, potentially changing both the coefficient frame and
382    /// the function.  Saved-model inference must instead consume the frozen
383    /// span coefficients and anchor map byte-for-byte.  The caller is
384    /// responsible for validating the saved schema marker before entering
385    /// this constructor.
386    pub(crate) fn from_exact_cubic_tables(
387        breakpoints: Array1<f64>,
388        span_c0: Array2<f64>,
389        span_c1: Array2<f64>,
390        span_c2: Array2<f64>,
391        span_c3: Array2<f64>,
392        installed_flex_block: Option<InstalledFlexBlock>,
393        anchor_rows_at_training: Option<Array2<f64>>,
394    ) -> Result<Self, String> {
395        validate_breakpoints(
396            breakpoints.as_slice().ok_or_else(|| {
397                String::from(DeviationRuntimeError::InvalidInput {
398                    reason: "saved deviation breakpoints are not contiguous".to_string(),
399                })
400            })?,
401            "saved deviation replay breakpoints",
402        )?;
403        let n_spans = breakpoints.len() - 1;
404        let basis_dim = span_c0.ncols();
405        if basis_dim == 0 {
406            return Err(DeviationRuntimeError::DimensionMismatch {
407                reason: "saved deviation replay requires at least one basis column".to_string(),
408            }
409            .into());
410        }
411        let expected = (n_spans, basis_dim);
412        for (label, coefficients) in [
413            ("c0", &span_c0),
414            ("c1", &span_c1),
415            ("c2", &span_c2),
416            ("c3", &span_c3),
417        ] {
418            if coefficients.dim() != expected {
419                return Err(DeviationRuntimeError::DimensionMismatch {
420                    reason: format!(
421                        "saved deviation replay {label} table is {}x{}; expected {}x{}",
422                        coefficients.nrows(),
423                        coefficients.ncols(),
424                        expected.0,
425                        expected.1,
426                    ),
427                }
428                .into());
429            }
430            if let Some(((row, column), value)) = coefficients
431                .indexed_iter()
432                .find(|(_, value)| !value.is_finite())
433            {
434                return Err(DeviationRuntimeError::InvalidInput {
435                    reason: format!(
436                        "saved deviation replay {label}[{row},{column}] is non-finite ({value})"
437                    ),
438                }
439                .into());
440            }
441        }
442
443        let final_span = n_spans - 1;
444        let width = breakpoints[n_spans] - breakpoints[final_span];
445        let mut right_boundary_value_row = Array1::<f64>::zeros(basis_dim);
446        for basis in 0..basis_dim {
447            right_boundary_value_row[basis] = span_c0[[final_span, basis]]
448                + width
449                    * (span_c1[[final_span, basis]]
450                        + width
451                            * (span_c2[[final_span, basis]]
452                                + width * span_c3[[final_span, basis]]));
453        }
454        if let Some((basis, value)) = right_boundary_value_row
455            .iter()
456            .copied()
457            .enumerate()
458            .find(|(_, value)| !value.is_finite())
459        {
460            return Err(DeviationRuntimeError::InvalidInput {
461                reason: format!(
462                    "saved deviation replay right-boundary value[{basis}] is non-finite ({value})"
463                ),
464            }
465            .into());
466        }
467        let monotonicity_constraint_rows = build_quadratic_derivative_bernstein_constraints(
468            &breakpoints,
469            &span_c1,
470            &span_c2,
471            &span_c3,
472        )?;
473
474        match (&installed_flex_block, &anchor_rows_at_training) {
475            (Some(installed), Some(rows)) => {
476                if rows.ncols() != installed.anchor_correction.nrows() {
477                    return Err(DeviationRuntimeError::DimensionMismatch {
478                        reason: format!(
479                            "saved deviation replay anchor rows have {} columns; anchor correction requires {}",
480                            rows.ncols(),
481                            installed.anchor_correction.nrows(),
482                        ),
483                    }
484                    .into());
485                }
486                if installed.anchor_correction.ncols() != basis_dim {
487                    return Err(DeviationRuntimeError::DimensionMismatch {
488                        reason: format!(
489                            "saved deviation replay anchor correction has {} columns; basis has {basis_dim}",
490                            installed.anchor_correction.ncols(),
491                        ),
492                    }
493                    .into());
494                }
495            }
496            (Some(_), None) => {
497                return Err(DeviationRuntimeError::DimensionMismatch {
498                    reason: "saved deviation replay has an anchor correction but no row-aligned anchor design"
499                        .to_string(),
500                }
501                .into());
502            }
503            (None, Some(rows)) if rows.ncols() != 0 => {
504                return Err(DeviationRuntimeError::DimensionMismatch {
505                    reason: format!(
506                        "saved deviation replay has {} anchor columns but no anchor correction",
507                        rows.ncols()
508                    ),
509                }
510                .into());
511            }
512            // No anchor correction was saved, so a zero-column anchor design
513            // (or no anchor design at all) is the consistent no-anchor replay
514            // and needs no cross-check.
515            (None, None) | (None, Some(_)) => {}
516        }
517
518        Ok(Self {
519            degree: 2,
520            value_span_degree: 3,
521            basis_dim,
522            // The persisted tables are already the fitted constrained
523            // function.  Replay never re-solves feasibility, so there is no
524            // configuration-space epsilon to reconstruct here.
525            monotonicity_eps: 0.0,
526            endpoint_points: breakpoints,
527            span_c0,
528            span_c1,
529            span_c2,
530            span_c3,
531            monotonicity_constraint_rows,
532            right_boundary_value_row,
533            installed_flex_block,
534            anchor_rows_at_training,
535        })
536    }
537
538    /// Construct the link-deviation runtime with a smoothness-null-space-drop
539    /// basis transform. `max_penalty_derivative_order` is the highest
540    /// derivative order of any penalty that will subsequently be applied to
541    /// this block (computed by the caller from its `DeviationBlockConfig`).
542    /// The returned basis structurally excludes polynomials of degree
543    /// `< max_penalty_derivative_order`, so the configured smoothness
544    /// penalties have no null space on the transformed basis and the
545    /// joint Hessian + penalty system is well-conditioned at every PIRLS
546    /// iteration regardless of how β shifts the linear predictor distribution.
547    ///
548    /// This replaces the previous data-distribution moment anchor (at the
549    /// rigid-pilot η₀), which gave a β-dependent identifiability constraint
550    /// that drifted out of alignment with η_current during PIRLS and produced
551    /// near-singular joint Hessians (σ_min ≈ ridge_floor).
552    pub(crate) fn try_new(
553        knots: Array1<f64>,
554        monotonicity_eps: f64,
555        max_penalty_derivative_order: usize,
556    ) -> Result<Self, String> {
557        Self::try_new_with_smoothness_drop(knots, monotonicity_eps, max_penalty_derivative_order)
558    }
559
560    pub(super) fn try_new_with_smoothness_drop(
561        knots: Array1<f64>,
562        monotonicity_eps: f64,
563        max_penalty_derivative_order: usize,
564    ) -> Result<Self, String> {
565        if !monotonicity_eps.is_finite() || monotonicity_eps < 0.0 {
566            return Err(DeviationRuntimeError::InvalidInput {
567                reason: format!(
568                    "DeviationRuntime monotonicity_eps must be finite and non-negative, got {monotonicity_eps}"
569                ),
570            }
571            .into());
572        }
573
574        let bkpts = breakpoints_from_knots(
575            knots.as_slice().ok_or_else(|| {
576                String::from(DeviationRuntimeError::InvalidInput {
577                    reason: "DeviationRuntime knots are not contiguous".to_string(),
578                })
579            })?,
580            "DeviationRuntime breakpoints",
581        )?;
582        let endpoint_points = Array1::from_vec(bkpts);
583        if endpoint_points.len() < 3 {
584            return Err(DeviationRuntimeError::InvalidInput {
585                reason:
586                    "DeviationRuntime requires at least two active knot spans and one interior node"
587                        .to_string(),
588            }
589            .into());
590        }
591        let n_spans = endpoint_points.len() - 1;
592        for span_idx in 0..n_spans {
593            let left = endpoint_points[span_idx];
594            let right = endpoint_points[span_idx + 1];
595            let width = right - left;
596            if !width.is_finite() || width <= 0.0 {
597                return Err(DeviationRuntimeError::InvalidInput {
598                    reason: format!(
599                        "DeviationRuntime requires strictly increasing span endpoints at span {span_idx}: left={left}, right={right}"
600                    ),
601                }
602                .into());
603            }
604        }
605        let span_lefts = Array1::from_iter((0..n_spans).map(|idx| endpoint_points[idx]));
606        let span_midpoints = Array1::from_iter(
607            (0..n_spans).map(|idx| 0.5 * (endpoint_points[idx] + endpoint_points[idx + 1])),
608        );
609        let right_endpoint = Array1::from_vec(vec![endpoint_points[n_spans]]);
610        let internal_degree = 2usize;
611        let raw_span_c0 =
612            create_ispline_derivative_dense(span_lefts.view(), &knots, internal_degree, 0)
613                .map_err(|e| {
614                    String::from(DeviationRuntimeError::NumericalFailure {
615                        reason: format!("DeviationRuntime cubic I-spline values failed: {e}"),
616                    })
617                })?;
618        let raw_span_c1 =
619            create_ispline_derivative_dense(span_lefts.view(), &knots, internal_degree, 1)
620                .map_err(|e| {
621                    String::from(DeviationRuntimeError::NumericalFailure {
622                        reason: format!(
623                            "DeviationRuntime cubic I-spline first derivatives failed: {e}"
624                        ),
625                    })
626                })?;
627        let raw_span_c2 =
628            create_ispline_derivative_dense(span_lefts.view(), &knots, internal_degree, 2)
629                .map_err(|e| {
630                    String::from(DeviationRuntimeError::NumericalFailure {
631                        reason: format!(
632                            "DeviationRuntime cubic I-spline second derivatives failed: {e}"
633                        ),
634                    })
635                })?
636                .mapv(|value| 0.5 * value);
637        let raw_span_c3 =
638            create_ispline_derivative_dense(span_midpoints.view(), &knots, internal_degree, 3)
639                .map_err(|e| {
640                    String::from(DeviationRuntimeError::NumericalFailure {
641                        reason: format!(
642                            "DeviationRuntime cubic I-spline third derivatives failed: {e}"
643                        ),
644                    })
645                })?
646                .mapv(|value| value / 6.0);
647        let raw_right_boundary_values =
648            create_ispline_derivative_dense(right_endpoint.view(), &knots, internal_degree, 0)
649                .map_err(|e| {
650                    String::from(DeviationRuntimeError::NumericalFailure {
651                        reason: format!(
652                            "DeviationRuntime cubic I-spline right boundary failed: {e}"
653                        ),
654                    })
655                })?;
656        let raw_right_boundary_value_row = raw_right_boundary_values.row(0).to_owned();
657
658        if max_penalty_derivative_order == 0 {
659            return Err(
660                "DeviationRuntime requires max_penalty_derivative_order >= 1 so the basis can \
661                 drop the corresponding smoothness null space; an order-0 (mass) penalty alone \
662                 has no null space and would not require any drop"
663                    .to_string(),
664            );
665        }
666        if max_penalty_derivative_order > 3 {
667            return Err(format!(
668                "DeviationRuntime cubic basis supports derivative orders up to 3; got max \
669                 penalty derivative order {max_penalty_derivative_order}"
670            ));
671        }
672        let raw_smoothness_penalty = raw_integrated_derivative_penalty(
673            &endpoint_points,
674            &raw_span_c0,
675            &raw_span_c1,
676            &raw_span_c2,
677            &raw_span_c3,
678            max_penalty_derivative_order,
679        )?;
680        let coefficient_transform =
681            smoothness_nullspace_orthogonal_complement(&raw_smoothness_penalty)?;
682        let basis_dim = coefficient_transform.ncols();
683        let span_c0 = fast_ab(&raw_span_c0, &coefficient_transform);
684        let span_c1 = fast_ab(&raw_span_c1, &coefficient_transform);
685        let span_c2 = fast_ab(&raw_span_c2, &coefficient_transform);
686        let span_c3 = fast_ab(&raw_span_c3, &coefficient_transform);
687        let right_boundary_value_row = raw_right_boundary_value_row.dot(&coefficient_transform);
688        let monotonicity_constraint_rows = build_quadratic_derivative_bernstein_constraints(
689            &endpoint_points,
690            &span_c1,
691            &span_c2,
692            &span_c3,
693        )?;
694
695        Ok(Self {
696            degree: 3,
697            value_span_degree: 3,
698            basis_dim,
699            monotonicity_eps,
700            endpoint_points,
701            span_c0,
702            span_c1,
703            span_c2,
704            span_c3,
705            monotonicity_constraint_rows,
706            right_boundary_value_row,
707            installed_flex_block: None,
708            anchor_rows_at_training: None,
709        })
710    }
711
712    // The per-block `smoothness_nullspace_orthogonal_complement` transform
713    // above eliminates within-block polynomial aliasing (constants/linears in
714    // η_pilot) so the location block can carry the intercept. That handles
715    // single-flex-block configurations. When two flex blocks of η_pilot are
716    // simultaneously active (score-warp + linkwiggle), each is individually
717    // orthogonal to constants, but their column spans still overlap inside
718    // the orthogonal complement of constants — both are cubic I-spline bases
719    // of the same scalar argument. The overlap manifests as a near-null
720    // direction in the joint penalized Hessian: a linear combination of
721    // β_score_warp and β_link_dev that produces zero net η-contribution at
722    // the rigid-pilot training points yet costs only the (penalised) basis
723    // norm, so Newton steps along that direction blow up.
724    //
725    // Compose an external column transform `T` (shape `basis_dim × new_dim`)
726    // into the cubic span tables and monotonicity constraints. After this
727    // call every `design(...)`-style method returns matrices in the new
728    // `new_dim`-column parameterisation: `runtime.design(values) ==
729    // old_runtime.design(values) · T`. Penalties built later via
730    // `integrated_derivative_penalty_with_nullity` are also expressed in
731    // the new parameterisation.
732    //
733    // Used by `install_compiled_flex_block_into_runtime` to
734    // enforce the joint-design identifiability invariant in the W-metric
735    // (W = p(1−p) at training rows). With `A_train` the stacked parametric
736    // anchors and `C_train = span_eval(values)` the candidate basis at the
737    // training rows, the residualised candidate is
738    //
739    //     C̃_train = (I − P_A^{(W)}) C_train,    P_A^{(W)} = A(AᵀWA)⁻¹AᵀW
740    //
741    // and the kept directions are the eigenvectors of `C̃ᵀ W C̃` above the
742    // numerical noise floor. The block-triangular reparameterisation
743    // `Aβ_A + Cβ_C = A(β_A + Bβ_C) + (C − AB)β_C` with `B = (AᵀWA)⁻¹AᵀWC`
744    // means dropping a direction in C̃ drops *exactly* a direction
745    // span(C) shares with span(A) under W, leaving no aliasing in the
746    // joint design `[X_loc | X_logslope | A | C·V − N·M]` (full column
747    // rank up to numerical tolerance, so `σ_min(joint H+S) ≥ λ_min(S₊)`
748    // regardless of how β shifts the linear-predictor distribution).
749    //
750    // The old `T = null(A_trainᵀ C_train)` algorithm was wrong: that
751    // null-space is the candidate directions *already* exactly W-orthogonal
752    // to A (Gram = 0), not the directions left after projecting A out.
753    // `null(AᵀC) = ∅` does NOT imply `span(C) ⊆ span(A)` — counterexample
754    // `A = e₁`, `C = e₁ + e₂` has `AᵀC = 1 ≠ 0` (empty null space) yet
755    // `(I − P_A) C = e₂ ≠ 0`. Whenever the anchor is wider than the
756    // candidate (d ≥ p_c) the old test generically returned ∅ even when
757    // the residualised candidate had full rank, prompting a spurious
758    // "fully aliased" hard-error. The current code residualises and keeps
759    // exactly the surviving rank.
760    /// Compose a rank-reveal right-selector and an optional anchor-residual.
761    /// After this call, `design(x)` returns
762    ///   design_row(x) = span_eval(x) · V  −  n_row(x) · installed.anchor_correction
763    /// where V is `right_selector` (applied via right-multiplication into
764    /// `span_c{0..3}`). Only the `design()` path (derivative_order=0) subtracts
765    /// the residual: the anchor argument is a different scalar variable than
766    /// the candidate argument, so d/dx of `n_row(x)` w.r.t. the candidate
767    /// argument is identically zero.
768    pub(crate) fn compose_anchor_orthogonalisation(
769        &mut self,
770        right_selector: &Array2<f64>,
771        installed_flex_block: Option<InstalledFlexBlock>,
772    ) -> Result<(), String> {
773        let old_dim = self.basis_dim;
774        if right_selector.nrows() != old_dim {
775            return Err(DeviationRuntimeError::DimensionMismatch {
776                reason: format!(
777                    "DeviationRuntime cross-block transform shape mismatch: \
778                     transform rows={}, expected basis_dim={}",
779                    right_selector.nrows(),
780                    old_dim,
781                ),
782            }
783            .into());
784        }
785        let new_dim = right_selector.ncols();
786        if new_dim == 0 {
787            return Err(DeviationRuntimeError::DimensionMismatch {
788                reason: "DeviationRuntime cross-block transform reduces basis dim to 0; \
789                 the candidate's column span is fully aliased by the anchor block"
790                    .to_string(),
791            }
792            .into());
793        }
794        if new_dim > old_dim {
795            return Err(DeviationRuntimeError::DimensionMismatch {
796                reason: format!(
797                    "DeviationRuntime cross-block transform must not increase basis dim; \
798                     got new_dim={} from old_dim={}",
799                    new_dim, old_dim,
800                ),
801            }
802            .into());
803        }
804        if let Some(ref installed) = installed_flex_block {
805            let d_expected: usize = installed
806                .anchor_components
807                .iter()
808                .map(|c| match c {
809                    AnchorComponentTag::Parametric { ncols, .. } => *ncols,
810                    AnchorComponentTag::FlexEvaluation { ncols } => *ncols,
811                })
812                .sum();
813            if installed.anchor_correction.nrows() != d_expected {
814                return Err(DeviationRuntimeError::DimensionMismatch {
815                    reason: format!(
816                        "DeviationRuntime installed flex block: anchor_correction rows={}, expected sum-of-component-ncols={}",
817                        installed.anchor_correction.nrows(),
818                        d_expected,
819                    ),
820                }
821                .into());
822            }
823            if installed.anchor_correction.ncols() != new_dim {
824                return Err(DeviationRuntimeError::DimensionMismatch {
825                    reason: format!(
826                        "DeviationRuntime installed flex block: anchor_correction cols={}, expected new basis dim {}",
827                        installed.anchor_correction.ncols(),
828                        new_dim,
829                    ),
830                }
831                .into());
832            }
833        }
834        self.span_c0 = fast_ab(&self.span_c0, right_selector);
835        self.span_c1 = fast_ab(&self.span_c1, right_selector);
836        self.span_c2 = fast_ab(&self.span_c2, right_selector);
837        self.span_c3 = fast_ab(&self.span_c3, right_selector);
838        // `right_boundary_value_row` is a 1-D row vector of length basis_dim;
839        // right-multiplying by V (basis_dim × new_dim) gives the new row.
840        self.right_boundary_value_row = self.right_boundary_value_row.dot(right_selector);
841        // Monotonicity rows (n_constraints × basis_dim) follow the same
842        // right-multiplication. The constraint inequality `A β ≥ ε - 1`
843        // becomes `(A · V) β_new ≥ ε - 1` under the reparameterisation
844        // β = V β_new, so the row matrix is right-multiplied directly.
845        self.monotonicity_constraint_rows =
846            fast_ab(&self.monotonicity_constraint_rows, right_selector);
847        self.basis_dim = new_dim;
848        self.installed_flex_block = installed_flex_block;
849        Ok(())
850    }
851
852    /// Accessor for the installed flex block set via
853    /// `install_compiled_flex_block`. Save-time code uses this to snapshot
854    /// the install state into the saved model; predict-time code reconstructs
855    /// the per-row η correction `n_row · anchor_correction · β`.
856    pub fn installed_flex_block(&self) -> Option<&InstalledFlexBlock> {
857        self.installed_flex_block.as_ref()
858    }
859
860    /// Single-step install of a compiled flex block from
861    /// `identifiability::families::compiler::compile`.
862    ///
863    /// Semantics:
864    /// - `compiled.t_lw` is the right-selector `V` applied to `span_c{0..3}`,
865    ///   `right_boundary_value_row`, and `monotonicity_constraint_rows`.
866    /// - `compiled.anchor_correction` (always `Some` for non-empty anchor
867    ///   unions) is the d×k correction `M`.
868    /// - `anchor_components` records the per-anchor predict-time tags so
869    ///   the saved-model rebuild can replay the anchor row map.
870    /// - `n_train_at_training` is cached for
871    ///   `design_at_training_with_residual`.
872    pub(crate) fn install_compiled_flex_block(
873        &mut self,
874        compiled: &gam_identifiability::families::compiler::CompiledBlock,
875        anchor_components: Vec<AnchorComponentTag>,
876        n_train_at_training: Array2<f64>,
877    ) -> Result<(), String> {
878        let m = compiled.anchor_correction.as_ref().ok_or_else(|| {
879            "DeviationRuntime::install_compiled_flex_block: compiled block has no \
880             anchor_correction — install requires a non-empty anchor union"
881                .to_string()
882        })?;
883        let installed = InstalledFlexBlock {
884            anchor_correction: m.clone(),
885            anchor_components,
886        };
887        self.anchor_rows_at_training = Some(n_train_at_training);
888        self.compose_anchor_orthogonalisation(&compiled.t_lw, Some(installed))
889    }
890
891    /// Cached parametric-anchor matrix at training rows, installed by
892    /// `install_compiled_flex_block_into_runtime` when the
893    /// runtime is reparameterised against the parametric anchor union.
894    /// Used by per-row link-deviation evaluators that need the row's
895    /// anchor slice to apply `design_with_anchor_rows` correctly. Returns
896    /// `None` for runtimes that have not been reparameterised.
897    pub fn anchor_rows_at_training(&self) -> Option<&Array2<f64>> {
898        self.anchor_rows_at_training.as_ref()
899    }
900
901    /// Evaluate `design(values) - anchor_rows · M` where `anchor_rows` is
902    /// the n × d parametric-anchor matrix at the same rows as `values`.
903    /// Mandatory when an installed flex block is present; for runtimes
904    /// without one this is equivalent to `design(values)` and `anchor_rows`
905    /// must be `n × 0`.
906    pub fn design_with_anchor_rows(
907        &self,
908        values: &Array1<f64>,
909        anchor_rows: ArrayView2<f64>,
910    ) -> Result<Array2<f64>, String> {
911        let mut out = self.evaluate_span_polynomial_design_raw(values, 0)?;
912        if let Some(installed) = &self.installed_flex_block {
913            if anchor_rows.nrows() != values.len() {
914                return Err(DeviationRuntimeError::DimensionMismatch {
915                    reason: format!(
916                        "design_with_anchor_rows: anchor_rows has {} rows, expected {} (matching values)",
917                        anchor_rows.nrows(),
918                        values.len(),
919                    ),
920                }
921                .into());
922            }
923            if anchor_rows.ncols() != installed.anchor_correction.nrows() {
924                return Err(DeviationRuntimeError::DimensionMismatch {
925                    reason: format!(
926                        "design_with_anchor_rows: anchor_rows has {} cols, expected {} (sum of component ncols)",
927                        anchor_rows.ncols(),
928                        installed.anchor_correction.nrows(),
929                    ),
930                }
931                .into());
932            }
933            let subtract = anchor_rows.dot(&installed.anchor_correction);
934            out = out - subtract;
935        } else if anchor_rows.ncols() != 0 {
936            // Permit empty 0-col anchor rows without complaint; otherwise
937            // hard-error so callers don't silently pass mismatched rows.
938            return Err(DeviationRuntimeError::DimensionMismatch {
939                reason: format!(
940                    "design_with_anchor_rows: runtime has no installed flex block but anchor_rows has {} cols",
941                    anchor_rows.ncols(),
942                ),
943            }
944            .into());
945        }
946        Ok(out)
947    }
948
949    /// Rebuild the training-row design after orthogonalisation, using
950    /// `anchor_rows_at_training` cached at `install_compiled_flex_block` time.
951    pub(crate) fn design_at_training_with_residual(
952        &self,
953        values: &Array1<f64>,
954    ) -> Result<Array2<f64>, String> {
955        if let Some(rows) = self.anchor_rows_at_training.as_ref() {
956            self.design_with_anchor_rows(values, rows.view())
957        } else if self.installed_flex_block.is_some() {
958            Err(
959                "design_at_training_with_residual: runtime has installed_flex_block but no cached training anchor rows"
960                    .to_string(),
961            )
962        } else {
963            self.design(values)
964        }
965    }
966
967    // ── public field accessors ──
968
969    pub fn degree(&self) -> usize {
970        self.degree
971    }
972
973    pub fn value_span_degree(&self) -> usize {
974        self.value_span_degree
975    }
976
977    pub fn basis_dim(&self) -> usize {
978        self.basis_dim
979    }
980
981    pub fn monotonicity_eps(&self) -> f64 {
982        self.monotonicity_eps
983    }
984
985    pub fn span_c0(&self) -> &Array2<f64> {
986        &self.span_c0
987    }
988
989    pub fn span_c1(&self) -> &Array2<f64> {
990        &self.span_c1
991    }
992
993    pub fn span_c2(&self) -> &Array2<f64> {
994        &self.span_c2
995    }
996
997    pub fn span_c3(&self) -> &Array2<f64> {
998        &self.span_c3
999    }
1000
1001    // ── design evaluation ──
1002
1003    pub(super) fn validate_beta_shape(
1004        &self,
1005        beta: ArrayView1<'_, f64>,
1006        label: &str,
1007    ) -> Result<(), String> {
1008        if beta.len() != self.basis_dim {
1009            return Err(DeviationRuntimeError::DimensionMismatch {
1010                reason: format!(
1011                    "{label} length mismatch: got {}, expected {}",
1012                    beta.len(),
1013                    self.basis_dim
1014                ),
1015            }
1016            .into());
1017        }
1018        Ok::<(), _>(())
1019    }
1020
1021    /// Raw cubic-span polynomial design evaluation, without any
1022    /// anchor-residual subtraction. Internal — callers that need the
1023    /// residualised design must go through `design()` (which asserts no
1024    /// residual) or `design_with_anchor_rows()`.
1025    pub(super) fn evaluate_span_polynomial_design_raw(
1026        &self,
1027        values: &Array1<f64>,
1028        derivative_order: usize,
1029    ) -> Result<Array2<f64>, String> {
1030        let (left_ep, right_ep) = self.support_interval()?;
1031        let mut out = Array2::<f64>::zeros((values.len(), self.basis_dim));
1032        for (row_idx, &value) in values.iter().enumerate() {
1033            if !value.is_finite() {
1034                return Err(DeviationRuntimeError::InvalidInput {
1035                    reason: format!(
1036                        "deviation runtime design value at row {row_idx} is non-finite ({value})"
1037                    ),
1038                }
1039                .into());
1040            }
1041            if value < left_ep {
1042                if derivative_order == 0 {
1043                    out.row_mut(row_idx).assign(&self.span_c0.row(0));
1044                }
1045                continue;
1046            }
1047            if value > right_ep {
1048                if derivative_order == 0 {
1049                    out.row_mut(row_idx)
1050                        .assign(&self.right_boundary_value_row.view());
1051                }
1052                continue;
1053            }
1054            let span_idx = self.left_biased_span_index_for(value)?;
1055            let left = self.endpoint_points[span_idx];
1056            let t = value - left;
1057            for basis_idx in 0..self.basis_dim {
1058                let c0 = self.span_c0[[span_idx, basis_idx]];
1059                let c1 = self.span_c1[[span_idx, basis_idx]];
1060                let c2 = self.span_c2[[span_idx, basis_idx]];
1061                let c3 = self.span_c3[[span_idx, basis_idx]];
1062                out[[row_idx, basis_idx]] = match derivative_order {
1063                    0 => c0 + c1 * t + c2 * t * t + c3 * t * t * t,
1064                    1 => c1 + 2.0 * c2 * t + 3.0 * c3 * t * t,
1065                    2 => 2.0 * c2 + 6.0 * c3 * t,
1066                    3 => 6.0 * c3,
1067                    4 => 0.0,
1068                    other => {
1069                        return Err(DeviationRuntimeError::InvalidInput {
1070                            reason: format!(
1071                                "deviation runtime only supports derivative orders up to 4, got {other}"
1072                            ),
1073                        }
1074                        .into());
1075                    }
1076                };
1077            }
1078        }
1079        Ok(out)
1080    }
1081
1082    /// Pure-span design (no anchor-residual subtraction). Callers must
1083    /// ensure the runtime has no anchor residual; otherwise use
1084    /// `design_with_anchor_rows`. Derivative paths are unaffected: the
1085    /// residual subtraction `n_row · M` is constant in the candidate
1086    /// argument, so its derivatives are identically zero.
1087    pub fn design(&self, values: &Array1<f64>) -> Result<Array2<f64>, String> {
1088        assert!(
1089            self.installed_flex_block.is_none(),
1090            "DeviationRuntime::design called on a runtime with an installed flex block; \
1091             use design_with_anchor_rows or design_at_training_with_residual instead"
1092        );
1093        self.evaluate_span_polynomial_design_raw(values, 0)
1094    }
1095
1096    pub fn first_derivative_design(&self, values: &Array1<f64>) -> Result<Array2<f64>, String> {
1097        self.evaluate_span_polynomial_design_raw(values, 1)
1098    }
1099
1100    pub fn second_derivative_design(&self, values: &Array1<f64>) -> Result<Array2<f64>, String> {
1101        self.evaluate_span_polynomial_design_raw(values, 2)
1102    }
1103
1104    pub fn third_derivative_design(&self, values: &Array1<f64>) -> Result<Array2<f64>, String> {
1105        self.evaluate_span_polynomial_design_raw(values, 3)
1106    }
1107
1108    pub(crate) fn integrated_derivative_penalty_with_nullity(
1109        &self,
1110        derivative_order: usize,
1111    ) -> Result<(Array2<f64>, usize), String> {
1112        if derivative_order > self.value_span_degree {
1113            return Err(DeviationRuntimeError::InvalidInput {
1114                reason: format!(
1115                    "deviation penalty derivative order {derivative_order} exceeds value-basis degree {}",
1116                    self.value_span_degree
1117                ),
1118            }
1119            .into());
1120        }
1121        let mut penalty = Array2::<f64>::zeros((self.basis_dim, self.basis_dim));
1122        for span_idx in 0..self.span_count() {
1123            let (left, right) = self.span_interval(span_idx)?;
1124            let width = right - left;
1125            if !width.is_finite() || width <= 0.0 {
1126                return Err(DeviationRuntimeError::InvalidInput {
1127                    reason: format!("deviation penalty span {span_idx} has invalid width {width}"),
1128                }
1129                .into());
1130            }
1131            for i in 0..self.basis_dim {
1132                let ci =
1133                    self.span_derivative_polynomial_coefficients(span_idx, i, derivative_order)?;
1134                for j in i..self.basis_dim {
1135                    let cj = self.span_derivative_polynomial_coefficients(
1136                        span_idx,
1137                        j,
1138                        derivative_order,
1139                    )?;
1140                    let contribution = integrate_polynomial_product(&ci, &cj, width);
1141                    penalty[[i, j]] += contribution;
1142                    if i != j {
1143                        penalty[[j, i]] += contribution;
1144                    }
1145                }
1146            }
1147        }
1148        let (evals, _) = penalty.eigh(faer::Side::Lower).map_err(|e| {
1149            String::from(DeviationRuntimeError::NumericalFailure {
1150                reason: format!("deviation integrated penalty eigendecomposition failed: {e}"),
1151            })
1152        })?;
1153        let threshold = gam_solve::estimate::reml::reml_outer_engine::positive_eigenvalue_threshold(
1154            evals.as_slice().ok_or_else(|| {
1155                String::from(DeviationRuntimeError::NumericalFailure {
1156                    reason: "deviation penalty eigenvalues are not contiguous".to_string(),
1157                })
1158            })?,
1159        );
1160        let rank = evals.iter().filter(|&&value| value > threshold).count();
1161        let nullity = self.basis_dim.saturating_sub(rank);
1162        Ok((penalty, nullity))
1163    }
1164
1165    pub(crate) fn structural_monotonicity_constraints(&self) -> LinearInequalityConstraints {
1166        LinearInequalityConstraints {
1167            a: self.monotonicity_constraint_rows.clone(),
1168            b: Array1::from_elem(
1169                self.monotonicity_constraint_rows.nrows(),
1170                self.monotonicity_eps - 1.0,
1171            ),
1172        }
1173    }
1174
1175    // ── span geometry ──
1176
1177    pub(super) fn span_count(&self) -> usize {
1178        self.endpoint_points.len().saturating_sub(1)
1179    }
1180
1181    pub fn breakpoints(&self) -> &Array1<f64> {
1182        &self.endpoint_points
1183    }
1184
1185    pub(super) fn span_interval(&self, span_idx: usize) -> Result<(f64, f64), String> {
1186        if span_idx >= self.span_count() {
1187            return Err(DeviationRuntimeError::InvalidInput {
1188                reason: format!(
1189                    "deviation span index {} out of range for {} spans",
1190                    span_idx,
1191                    self.span_count()
1192                ),
1193            }
1194            .into());
1195        }
1196        Ok((
1197            self.endpoint_points[span_idx],
1198            self.endpoint_points[span_idx + 1],
1199        ))
1200    }
1201
1202    pub(super) fn span_index_for(&self, value: f64) -> Result<usize, String> {
1203        span_index_for_breakpoints(
1204            self.endpoint_points.as_slice().ok_or_else(|| {
1205                String::from(DeviationRuntimeError::InvalidInput {
1206                    reason: "deviation runtime breakpoints are not contiguous".to_string(),
1207                })
1208            })?,
1209            value,
1210            "deviation span lookup",
1211        )
1212    }
1213
1214    pub(super) fn left_biased_span_index_for(&self, value: f64) -> Result<usize, String> {
1215        let mut span_idx = self.span_index_for(value)?;
1216        // Bias to the LEFT-hand span at internal breakpoints. The cubic basis
1217        // is C², so value, first derivative, and second derivative are
1218        // unchanged; only the span-local third derivative needs a convention.
1219        if span_idx > 0 && value == self.endpoint_points[span_idx] {
1220            span_idx -= 1;
1221        }
1222        Ok(span_idx)
1223    }
1224
1225    pub(super) fn span_derivative_polynomial_coefficients(
1226        &self,
1227        span_idx: usize,
1228        basis_idx: usize,
1229        derivative_order: usize,
1230    ) -> Result<Vec<f64>, String> {
1231        if span_idx >= self.span_count() {
1232            return Err(DeviationRuntimeError::InvalidInput {
1233                reason: format!(
1234                    "deviation span index {} out of range for {} spans",
1235                    span_idx,
1236                    self.span_count()
1237                ),
1238            }
1239            .into());
1240        }
1241        if basis_idx >= self.basis_dim {
1242            return Err(DeviationRuntimeError::InvalidInput {
1243                reason: format!(
1244                    "deviation basis index {} out of range for {} coefficients",
1245                    basis_idx, self.basis_dim
1246                ),
1247            }
1248            .into());
1249        }
1250        let c0 = self.span_c0[[span_idx, basis_idx]];
1251        let c1 = self.span_c1[[span_idx, basis_idx]];
1252        let c2 = self.span_c2[[span_idx, basis_idx]];
1253        let c3 = self.span_c3[[span_idx, basis_idx]];
1254        match derivative_order {
1255            0 => Ok(vec![c0, c1, c2, c3]),
1256            1 => Ok(vec![c1, 2.0 * c2, 3.0 * c3]),
1257            2 => Ok(vec![2.0 * c2, 6.0 * c3]),
1258            3 => Ok(vec![6.0 * c3]),
1259            other => Err(DeviationRuntimeError::InvalidInput {
1260                reason: format!(
1261                    "deviation polynomial coefficients only support derivative orders up to 3, got {other}"
1262                ),
1263            }
1264            .into()),
1265        }
1266    }
1267
1268    // ── cubic Taylor extraction ──
1269
1270    pub(crate) fn local_cubic_on_span(
1271        &self,
1272        beta: ArrayView1<'_, f64>,
1273        span_idx: usize,
1274    ) -> Result<exact_kernel::LocalSpanCubic, String> {
1275        self.validate_beta_shape(beta.view(), "deviation local cubic coefficients")?;
1276        let (left, right) = self.span_interval(span_idx)?;
1277        Ok(exact_kernel::LocalSpanCubic {
1278            left,
1279            right,
1280            c0: self.span_c0.row(span_idx).dot(&beta),
1281            c1: self.span_c1.row(span_idx).dot(&beta),
1282            c2: self.span_c2.row(span_idx).dot(&beta),
1283            c3: self.span_c3.row(span_idx).dot(&beta),
1284        })
1285    }
1286
1287    pub fn basis_span_cubic(
1288        &self,
1289        span_idx: usize,
1290        basis_idx: usize,
1291    ) -> Result<exact_kernel::LocalSpanCubic, String> {
1292        if basis_idx >= self.basis_dim {
1293            return Err(DeviationRuntimeError::InvalidInput {
1294                reason: format!(
1295                    "deviation basis index {} out of range for {} coefficients",
1296                    basis_idx, self.basis_dim
1297                ),
1298            }
1299            .into());
1300        }
1301        let (left, right) = self.span_interval(span_idx)?;
1302        Ok(exact_kernel::LocalSpanCubic {
1303            left,
1304            right,
1305            c0: self.span_c0[[span_idx, basis_idx]],
1306            c1: self.span_c1[[span_idx, basis_idx]],
1307            c2: self.span_c2[[span_idx, basis_idx]],
1308            c3: self.span_c3[[span_idx, basis_idx]],
1309        })
1310    }
1311
1312    /// Return the correct per-basis `LocalSpanCubic` for any evaluation
1313    /// point. Strictly outside the knot support, returns a constant cubic
1314    /// (c1=c2=c3=0) at the saturated tail value. Interior breakpoints use the
1315    /// left span so span-local third derivatives match derivative designs.
1316    pub fn basis_cubic_at(
1317        &self,
1318        basis_idx: usize,
1319        value: f64,
1320    ) -> Result<exact_kernel::LocalSpanCubic, String> {
1321        if basis_idx >= self.basis_dim {
1322            return Err(DeviationRuntimeError::InvalidInput {
1323                reason: format!(
1324                    "deviation basis index {} out of range for {} coefficients",
1325                    basis_idx, self.basis_dim
1326                ),
1327            }
1328            .into());
1329        }
1330        let (left_ep, right_ep) = self.support_interval()?;
1331        if value < left_ep {
1332            return Ok(exact_kernel::LocalSpanCubic {
1333                left: left_ep,
1334                right: left_ep + 1.0,
1335                c0: self.span_c0[[0, basis_idx]],
1336                c1: 0.0,
1337                c2: 0.0,
1338                c3: 0.0,
1339            });
1340        }
1341        if value > right_ep {
1342            return Ok(exact_kernel::LocalSpanCubic {
1343                left: right_ep,
1344                right: right_ep + 1.0,
1345                c0: self.right_boundary_value_row[basis_idx],
1346                c1: 0.0,
1347                c2: 0.0,
1348                c3: 0.0,
1349            });
1350        }
1351        let span_idx = self.left_biased_span_index_for(value)?;
1352        self.basis_span_cubic(span_idx, basis_idx)
1353    }
1354
1355    pub fn for_each_basis_cubic_at<F>(&self, value: f64, mut visit: F) -> Result<(), String>
1356    where
1357        F: FnMut(usize, exact_kernel::LocalSpanCubic) -> Result<(), String>,
1358    {
1359        let (left_ep, right_ep) = self.support_interval()?;
1360        if value < left_ep {
1361            for basis_idx in 0..self.basis_dim {
1362                visit(
1363                    basis_idx,
1364                    exact_kernel::LocalSpanCubic {
1365                        left: left_ep,
1366                        right: left_ep + 1.0,
1367                        c0: self.span_c0[[0, basis_idx]],
1368                        c1: 0.0,
1369                        c2: 0.0,
1370                        c3: 0.0,
1371                    },
1372                )?;
1373            }
1374            return Ok(());
1375        }
1376        if value > right_ep {
1377            for basis_idx in 0..self.basis_dim {
1378                visit(
1379                    basis_idx,
1380                    exact_kernel::LocalSpanCubic {
1381                        left: right_ep,
1382                        right: right_ep + 1.0,
1383                        c0: self.right_boundary_value_row[basis_idx],
1384                        c1: 0.0,
1385                        c2: 0.0,
1386                        c3: 0.0,
1387                    },
1388                )?;
1389            }
1390            return Ok(());
1391        }
1392
1393        let span_idx = self.left_biased_span_index_for(value)?;
1394        let (left, right) = self.span_interval(span_idx)?;
1395        for basis_idx in 0..self.basis_dim {
1396            visit(
1397                basis_idx,
1398                exact_kernel::LocalSpanCubic {
1399                    left,
1400                    right,
1401                    c0: self.span_c0[[span_idx, basis_idx]],
1402                    c1: self.span_c1[[span_idx, basis_idx]],
1403                    c2: self.span_c2[[span_idx, basis_idx]],
1404                    c3: self.span_c3[[span_idx, basis_idx]],
1405                },
1406            )?;
1407        }
1408        Ok(())
1409    }
1410
1411    /// Return the correct composite `LocalSpanCubic` for any evaluation
1412    /// point. Strictly outside the knot support, returns a constant cubic
1413    /// (c1=c2=c3=0) at the saturated tail value. Interior breakpoints use the
1414    /// left span so span-local third derivatives match derivative designs.
1415    pub(crate) fn local_cubic_at(
1416        &self,
1417        beta: ArrayView1<'_, f64>,
1418        value: f64,
1419    ) -> Result<exact_kernel::LocalSpanCubic, String> {
1420        self.validate_beta_shape(beta.view(), "deviation local cubic")?;
1421        let (left_ep, right_ep) = self.support_interval()?;
1422        if value < left_ep {
1423            return Ok(exact_kernel::LocalSpanCubic {
1424                left: left_ep,
1425                right: left_ep + 1.0,
1426                c0: self.left_tail_value(beta.view()),
1427                c1: 0.0,
1428                c2: 0.0,
1429                c3: 0.0,
1430            });
1431        }
1432        if value > right_ep {
1433            return Ok(exact_kernel::LocalSpanCubic {
1434                left: right_ep,
1435                right: right_ep + 1.0,
1436                c0: self.right_tail_value(beta.view()),
1437                c1: 0.0,
1438                c2: 0.0,
1439                c3: 0.0,
1440            });
1441        }
1442        let span_idx = self.left_biased_span_index_for(value)?;
1443        self.local_cubic_on_span(beta, span_idx)
1444    }
1445
1446    // ── tail value helpers ──
1447
1448    /// Left-tail constant: deviation value at the leftmost breakpoint.
1449    /// For anchored I-spline bases this is the anchor value (typically 0).
1450    pub(super) fn left_tail_value(&self, beta: ArrayView1<'_, f64>) -> f64 {
1451        self.span_c0.row(0).dot(&beta)
1452    }
1453
1454    /// Right-tail constant: deviation value at the rightmost breakpoint.
1455    /// For I-spline bases this is the saturated integral value.
1456    pub(super) fn right_tail_value(&self, beta: ArrayView1<'_, f64>) -> f64 {
1457        self.right_boundary_value_row.dot(&beta)
1458    }
1459
1460    /// Conservative L1 sup-norm bound for the deviation value basis.
1461    ///
1462    /// For every evaluation point `x`, this returns a finite `K` such that
1463    /// `|B(x)·β| <= K * ||β||_∞`.  Each basis column is a cubic on each
1464    /// finite span and constant in the two tails, so the supremum is attained
1465    /// at a span endpoint, an interior root of the derivative, or a tail
1466    /// value.  Summing per-column suprema gives a conservative row-wise L1
1467    /// bound that is independent of `x`.
1468    pub(crate) fn value_basis_l1_sup_norm(&self) -> f64 {
1469        let mut total = 0.0;
1470        for basis_idx in 0..self.basis_dim {
1471            let mut col_sup = self.span_c0[[0, basis_idx]]
1472                .abs()
1473                .max(self.right_boundary_value_row[basis_idx].abs());
1474            for span_idx in 0..self.span_count() {
1475                let left = self.endpoint_points[span_idx];
1476                let right = self.endpoint_points[span_idx + 1];
1477                let width = right - left;
1478                if !width.is_finite() || width <= 0.0 {
1479                    continue;
1480                }
1481                let c0 = self.span_c0[[span_idx, basis_idx]];
1482                let c1 = self.span_c1[[span_idx, basis_idx]];
1483                let c2 = self.span_c2[[span_idx, basis_idx]];
1484                let c3 = self.span_c3[[span_idx, basis_idx]];
1485                let eval_abs = |t: f64| (c0 + c1 * t + c2 * t * t + c3 * t * t * t).abs();
1486                col_sup = col_sup.max(eval_abs(0.0)).max(eval_abs(width));
1487                let a = 3.0 * c3;
1488                let b = 2.0 * c2;
1489                let c = c1;
1490                if a.abs() <= f64::EPSILON {
1491                    if b.abs() > f64::EPSILON {
1492                        let t = -c / b;
1493                        if t > 0.0 && t < width {
1494                            col_sup = col_sup.max(eval_abs(t));
1495                        }
1496                    }
1497                } else {
1498                    let disc = b * b - 4.0 * a * c;
1499                    if disc >= 0.0 {
1500                        let sqrt_disc = disc.sqrt();
1501                        for t in [(-b - sqrt_disc) / (2.0 * a), (-b + sqrt_disc) / (2.0 * a)] {
1502                            if t > 0.0 && t < width {
1503                                col_sup = col_sup.max(eval_abs(t));
1504                            }
1505                        }
1506                    }
1507                }
1508            }
1509            total += col_sup;
1510        }
1511        total
1512    }
1513
1514    // ── monotonicity enforcement ──
1515
1516    pub(super) fn support_interval(&self) -> Result<(f64, f64), String> {
1517        match (self.endpoint_points.first(), self.endpoint_points.last()) {
1518            (Some(&left), Some(&right)) => Ok((left, right)),
1519            _ => Err(DeviationRuntimeError::InvalidInput {
1520                reason: "deviation runtime is missing monotonicity support points".to_string(),
1521            }
1522            .into()),
1523        }
1524    }
1525
1526    pub(crate) fn exact_monotonicity_min_slack(&self, beta: &Array1<f64>) -> Result<f64, String> {
1527        if beta.len() != self.basis_dim {
1528            return Err(DeviationRuntimeError::DimensionMismatch {
1529                reason: format!(
1530                    "deviation monotonicity length mismatch: got {}, expected {}",
1531                    beta.len(),
1532                    self.basis_dim
1533                ),
1534            }
1535            .into());
1536        }
1537        if beta.iter().any(|value| !value.is_finite()) {
1538            let bad = beta
1539                .iter()
1540                .enumerate()
1541                .find(|(_, value)| !value.is_finite())
1542                .map(|(idx, value)| format!("deviation coefficient {idx} is non-finite ({value})"))
1543                .unwrap_or_else(|| "deviation coefficient is non-finite".to_string());
1544            return Err(DeviationRuntimeError::InvalidInput { reason: bad }.into());
1545        }
1546
1547        let mut min_slack = f64::INFINITY;
1548        for span_idx in 0..self.span_count() {
1549            let left = self.endpoint_points[span_idx];
1550            let right = self.endpoint_points[span_idx + 1];
1551            let width = right - left;
1552            if !width.is_finite() || width <= 0.0 {
1553                continue;
1554            }
1555            let c1 = self.span_c1.row(span_idx).dot(beta);
1556            let c2 = self.span_c2.row(span_idx).dot(beta);
1557            let c3 = self.span_c3.row(span_idx).dot(beta);
1558            let d1_left = c1;
1559            let d1_right = c1 + 2.0 * c2 * width + 3.0 * c3 * width * width;
1560            let d2_left = 2.0 * c2;
1561            let d3 = 6.0 * c3;
1562            let left_slack = 1.0 + d1_left - self.monotonicity_eps;
1563            let right_slack = 1.0 + d1_right - self.monotonicity_eps;
1564            min_slack = min_slack.min(left_slack.min(right_slack));
1565
1566            if d3 > 0.0 {
1567                let t_star = -d2_left / d3;
1568                if t_star > 0.0 && t_star < width {
1569                    let interior = 1.0 + d1_left + d2_left * t_star + 0.5 * d3 * t_star * t_star
1570                        - self.monotonicity_eps;
1571                    min_slack = min_slack.min(interior);
1572                }
1573            }
1574        }
1575        if min_slack.is_finite() {
1576            Ok(min_slack)
1577        } else {
1578            Err(DeviationRuntimeError::NumericalFailure {
1579                reason: "deviation monotonicity slack computation produced no active spans"
1580                    .to_string(),
1581            }
1582            .into())
1583        }
1584    }
1585
1586    pub(crate) fn monotonicity_feasible(
1587        &self,
1588        beta: &Array1<f64>,
1589        context: &str,
1590    ) -> Result<(), String> {
1591        let slack = self.exact_monotonicity_min_slack(beta)?;
1592        if slack >= MONOTONICITY_SLACK_ROUNDOFF_TOL {
1593            Ok(())
1594        } else {
1595            let (left, right) = self.support_interval()?;
1596            Err(DeviationRuntimeError::NumericalFailure {
1597                reason: format!(
1598                    "{context} violates exact monotonicity on [{left:.6}, {right:.6}] (minimum derivative slack {slack:.3e}, eps={:.3e})",
1599                    self.monotonicity_eps
1600                ),
1601            }
1602            .into())
1603        }
1604    }
1605}