Skip to main content

gam_models/bms/
alo_replay.rs

1use std::sync::Arc;
2
3use super::deviation_runtime::{AnchorComponentTag, InstalledFlexBlock};
4use super::family::{BernoulliMarginalSlopeFamily, bernoulli_marginal_link_map};
5use super::gradient_paths::rigid_standard_normal_row_kernel;
6use super::hessian_paths::{
7    block_slices, new_cell_moment_cache_stats, new_cell_moment_lru_cache, primary_slices,
8};
9use super::{DeviationRuntime, LatentMeasureKind};
10use crate::inference::model::{SavedAnchorKind, SavedCompiledFlexBlock};
11use gam_linalg::matrix::DesignMatrix;
12use gam_problem::InverseLink;
13use gam_problem::ParameterBlockState;
14use ndarray::{Array1, Array2};
15
16/// Complete saved-row state for exact rigid Bernoulli marginal-slope ALO replay.
17pub struct BernoulliMarginalSlopeAloRowInput<'a> {
18    pub base_link: &'a InverseLink,
19    pub marginal_eta: f64,
20    pub slope: f64,
21    pub latent_z: f64,
22    pub response: f64,
23    pub prior_weight: f64,
24    pub probit_frailty_scale: f64,
25}
26
27/// Negative-log-likelihood derivatives in the affine fitted coordinates
28/// `[marginal eta, slope]` for one saved Bernoulli marginal-slope row.
29#[derive(Clone, Debug, PartialEq)]
30pub struct BernoulliMarginalSlopeAloRowGeometry {
31    pub negative_log_likelihood: f64,
32    pub nll_score: [f64; 2],
33    pub observed_hessian: [[f64; 2]; 2],
34}
35
36/// Replay the exact rigid standard-normal row program used by fitting.
37///
38/// The latent score supplied here must already be in the fitted normalized and
39/// calibrated coordinate system. Gaussian-shift frailty is represented by the
40/// persisted probit scale, so no prediction-time approximation enters the
41/// score or observed Hessian.
42pub fn bernoulli_marginal_slope_alo_row_geometry(
43    input: BernoulliMarginalSlopeAloRowInput<'_>,
44) -> Result<BernoulliMarginalSlopeAloRowGeometry, String> {
45    let marginal = bernoulli_marginal_link_map(input.base_link, input.marginal_eta)?;
46    let (negative_log_likelihood, nll_score, observed_hessian) = rigid_standard_normal_row_kernel(
47        marginal,
48        input.slope,
49        input.latent_z,
50        input.response,
51        input.prior_weight,
52        input.probit_frailty_scale,
53    )?;
54    Ok(BernoulliMarginalSlopeAloRowGeometry {
55        negative_log_likelihood,
56        nll_score,
57        observed_hessian,
58    })
59}
60
61/// Exact saved-row geometry in the full local primary frame
62/// `[marginal eta, slope, score-warp coefficients..., link-deviation
63/// coefficients...]`.
64#[derive(Clone, Debug)]
65pub struct BernoulliMarginalSlopeSavedAloRowGeometry {
66    pub nll_score: Array1<f64>,
67    pub observed_hessian: Array2<f64>,
68    pub coordinate_values: Array1<f64>,
69}
70
71/// Row-aligned replay of the fitted Bernoulli marginal-slope likelihood.
72#[derive(Clone, Debug)]
73pub struct BernoulliMarginalSlopeSavedAloReplay {
74    pub rows: Vec<BernoulliMarginalSlopeSavedAloRowGeometry>,
75    pub score_warp_dimension: usize,
76    pub link_deviation_dimension: usize,
77}
78
79pub(crate) struct BernoulliMarginalSlopeSavedAloReplayInput<'a> {
80    pub base_link: &'a InverseLink,
81    pub marginal_design: &'a DesignMatrix,
82    pub logslope_design: &'a DesignMatrix,
83    pub marginal_beta: &'a Array1<f64>,
84    pub logslope_beta: &'a Array1<f64>,
85    pub score_warp_beta: Option<&'a Array1<f64>>,
86    pub link_deviation_beta: Option<&'a Array1<f64>>,
87    pub marginal_eta: &'a Array1<f64>,
88    pub slope: &'a Array1<f64>,
89    pub latent_z: &'a Array1<f64>,
90    pub response: &'a Array1<f64>,
91    pub prior_weights: &'a Array1<f64>,
92    pub latent_measure: LatentMeasureKind,
93    pub gaussian_frailty_sd: Option<f64>,
94    pub score_warp_runtime: Option<&'a SavedCompiledFlexBlock>,
95    pub link_deviation_runtime: Option<&'a SavedCompiledFlexBlock>,
96    pub score_warp_anchor_rows: Option<&'a Array2<f64>>,
97    pub link_deviation_anchor_rows: Option<&'a Array2<f64>>,
98}
99
100fn dense_saved_table(
101    rows: &[Vec<f64>],
102    n_spans: usize,
103    basis_dim: usize,
104    label: &str,
105) -> Result<Array2<f64>, String> {
106    if rows.len() != n_spans || rows.iter().any(|row| row.len() != basis_dim) {
107        return Err(format!(
108            "saved {label} table is ragged or mis-sized: rows={}, expected={n_spans}, basis_dim={basis_dim}",
109            rows.len(),
110        ));
111    }
112    let values = rows
113        .iter()
114        .flat_map(|row| row.iter().copied())
115        .collect::<Vec<_>>();
116    Array2::from_shape_vec((n_spans, basis_dim), values)
117        .map_err(|error| format!("saved {label} table shape: {error}"))
118}
119
120fn dense_anchor_correction(rows: &[Vec<f64>], basis_dim: usize) -> Result<Array2<f64>, String> {
121    let nrows = rows.len();
122    if rows.iter().any(|row| row.len() != basis_dim) {
123        return Err(format!(
124            "saved anchor correction is ragged or has a row outside basis dimension {basis_dim}"
125        ));
126    }
127    Array2::from_shape_vec(
128        (nrows, basis_dim),
129        rows.iter().flat_map(|row| row.iter().copied()).collect(),
130    )
131    .map_err(|error| format!("saved anchor correction shape: {error}"))
132}
133
134pub(crate) fn exact_runtime_from_saved(
135    saved: &SavedCompiledFlexBlock,
136    anchor_rows: Option<&Array2<f64>>,
137    label: &str,
138) -> Result<DeviationRuntime, String> {
139    saved
140        .validate_exact_replay_contract()
141        .map_err(|error| format!("{label}: {error}"))?;
142    let n_spans = saved.breakpoints.len() - 1;
143    let c0 = dense_saved_table(
144        &saved.span_c0,
145        n_spans,
146        saved.basis_dim,
147        &format!("{label} c0"),
148    )?;
149    let c1 = dense_saved_table(
150        &saved.span_c1,
151        n_spans,
152        saved.basis_dim,
153        &format!("{label} c1"),
154    )?;
155    let c2 = dense_saved_table(
156        &saved.span_c2,
157        n_spans,
158        saved.basis_dim,
159        &format!("{label} c2"),
160    )?;
161    let c3 = dense_saved_table(
162        &saved.span_c3,
163        n_spans,
164        saved.basis_dim,
165        &format!("{label} c3"),
166    )?;
167    let installed = match saved.anchor_correction.as_ref() {
168        Some(correction) => {
169            let anchor_rows = anchor_rows.ok_or_else(|| {
170                format!(
171                    "saved {label} has a cross-block anchor map but no row-aligned anchor design"
172                )
173            })?;
174            let anchor_components = saved
175                .anchor_components
176                .iter()
177                .map(|component| match &component.kind {
178                    SavedAnchorKind::Parametric { block, ncols } => {
179                        AnchorComponentTag::Parametric {
180                            block: *block,
181                            ncols: *ncols,
182                        }
183                    }
184                    SavedAnchorKind::FlexEvaluation { ncols } => {
185                        AnchorComponentTag::FlexEvaluation { ncols: *ncols }
186                    }
187                })
188                .collect::<Vec<_>>();
189            let expected_anchor_columns = anchor_components
190                .iter()
191                .map(|component| match component {
192                    AnchorComponentTag::Parametric { ncols, .. }
193                    | AnchorComponentTag::FlexEvaluation { ncols } => *ncols,
194                })
195                .sum::<usize>();
196            if expected_anchor_columns == 0 {
197                return Err(format!("saved {label} anchor map has no anchor components"));
198            }
199            if anchor_rows.ncols() != expected_anchor_columns {
200                return Err(format!(
201                    "saved {label} anchor design has {} columns; component layout requires {expected_anchor_columns}",
202                    anchor_rows.ncols(),
203                ));
204            }
205            Some(InstalledFlexBlock {
206                anchor_correction: dense_anchor_correction(correction, saved.basis_dim)?,
207                anchor_components,
208            })
209        }
210        None => {
211            if anchor_rows.is_some_and(|rows| rows.ncols() != 0) {
212                return Err(format!(
213                    "saved {label} received anchor rows without a persisted anchor map"
214                ));
215            }
216            None
217        }
218    };
219    DeviationRuntime::from_exact_cubic_tables(
220        Array1::from_vec(saved.breakpoints.clone()),
221        c0,
222        c1,
223        c2,
224        c3,
225        installed,
226        anchor_rows.cloned(),
227    )
228}
229
230fn validate_optional_flex_block(
231    runtime: Option<&SavedCompiledFlexBlock>,
232    beta: Option<&Array1<f64>>,
233    label: &str,
234) -> Result<usize, String> {
235    match (runtime, beta) {
236        (None, None) => Ok(0),
237        (Some(runtime), Some(beta)) if runtime.basis_dim == beta.len() => Ok(beta.len()),
238        (Some(runtime), Some(beta)) => Err(format!(
239            "saved {label} runtime has basis dimension {}; beta has {} entries",
240            runtime.basis_dim,
241            beta.len(),
242        )),
243        (Some(_), None) => Err(format!(
244            "saved {label} runtime has no fitted coefficient block"
245        )),
246        (None, Some(_)) => Err(format!("saved {label} coefficients have no exact runtime")),
247    }
248}
249
250/// Replay the exact fit-time BMS row program from frozen saved state.
251///
252/// No basis compilation, fitting, numerical differentiation, or alternate
253/// likelihood is permitted here.  The saved cubic tables are rehydrated in
254/// their fitted coefficient frame and passed to the same observed-Hessian
255/// row authority used by the optimizer.
256pub(crate) fn replay_saved_bernoulli_marginal_slope_alo(
257    input: BernoulliMarginalSlopeSavedAloReplayInput<'_>,
258) -> Result<BernoulliMarginalSlopeSavedAloReplay, String> {
259    let n = input.response.len();
260    if n == 0
261        || input.prior_weights.len() != n
262        || input.marginal_design.nrows() != n
263        || input.logslope_design.nrows() != n
264        || input.marginal_eta.len() != n
265        || input.slope.len() != n
266        || input.latent_z.len() != n
267    {
268        return Err(format!(
269            "saved BMS ALO row mismatch: response={n}, weights={}, marginal_design={}, logslope_design={}, marginal_eta={}, slope={}, z={}",
270            input.prior_weights.len(),
271            input.marginal_design.nrows(),
272            input.logslope_design.nrows(),
273            input.marginal_eta.len(),
274            input.slope.len(),
275            input.latent_z.len(),
276        ));
277    }
278    if input.marginal_design.ncols() != input.marginal_beta.len()
279        || input.logslope_design.ncols() != input.logslope_beta.len()
280    {
281        return Err(format!(
282            "saved BMS ALO affine frame mismatch: marginal design/beta={}/{}, logslope design/beta={}/{}",
283            input.marginal_design.ncols(),
284            input.marginal_beta.len(),
285            input.logslope_design.ncols(),
286            input.logslope_beta.len(),
287        ));
288    }
289    if let Some((row, weight)) = input
290        .prior_weights
291        .iter()
292        .copied()
293        .enumerate()
294        .find(|(_, weight)| !weight.is_finite() || *weight < 0.0)
295    {
296        return Err(format!(
297            "saved BMS ALO prior weight[{row}] must be finite and non-negative, got {weight}"
298        ));
299    }
300    if let Some((row, response)) = input
301        .response
302        .iter()
303        .copied()
304        .enumerate()
305        .find(|(_, response)| *response != 0.0 && *response != 1.0)
306    {
307        return Err(format!(
308            "saved BMS ALO response[{row}] must be exactly 0 or 1, got {response}"
309        ));
310    }
311    for (label, values) in [
312        ("marginal eta", input.marginal_eta),
313        ("slope", input.slope),
314        ("latent z", input.latent_z),
315        ("marginal beta", input.marginal_beta),
316        ("logslope beta", input.logslope_beta),
317    ] {
318        if let Some((row, value)) = values
319            .iter()
320            .copied()
321            .enumerate()
322            .find(|(_, value)| !value.is_finite())
323        {
324            return Err(format!(
325                "saved BMS ALO {label}[{row}] must be finite, got {value}"
326            ));
327        }
328    }
329    for (label, beta) in [
330        ("score-warp", input.score_warp_beta),
331        ("link-deviation", input.link_deviation_beta),
332    ] {
333        if let Some((coordinate, value)) = beta.and_then(|beta| {
334            beta.iter()
335                .copied()
336                .enumerate()
337                .find(|(_, value)| !value.is_finite())
338        }) {
339            return Err(format!(
340                "saved BMS ALO {label} beta[{coordinate}] must be finite, got {value}"
341            ));
342        }
343    }
344    for (label, rows) in [
345        ("score-warp", input.score_warp_anchor_rows),
346        ("link-deviation", input.link_deviation_anchor_rows),
347    ] {
348        if let Some(rows) = rows
349            && rows.nrows() != n
350        {
351            return Err(format!(
352                "saved BMS ALO {label} anchor design has {} rows; expected {n}",
353                rows.nrows(),
354            ));
355        }
356    }
357    input
358        .latent_measure
359        .validate("saved BMS ALO latent measure")?;
360    let score_warp_dimension = validate_optional_flex_block(
361        input.score_warp_runtime,
362        input.score_warp_beta,
363        "score-warp",
364    )?;
365    let link_deviation_dimension = validate_optional_flex_block(
366        input.link_deviation_runtime,
367        input.link_deviation_beta,
368        "link-deviation",
369    )?;
370    let score_warp = input
371        .score_warp_runtime
372        .map(|runtime| {
373            exact_runtime_from_saved(runtime, input.score_warp_anchor_rows, "score-warp")
374        })
375        .transpose()?;
376    let link_dev = input
377        .link_deviation_runtime
378        .map(|runtime| {
379            exact_runtime_from_saved(runtime, input.link_deviation_anchor_rows, "link-deviation")
380        })
381        .transpose()?;
382
383    let policy = gam_runtime::resource::ResourcePolicy::default_library();
384    let family = BernoulliMarginalSlopeFamily {
385        y: Arc::new(input.response.clone()),
386        weights: Arc::new(input.prior_weights.clone()),
387        z: Arc::new(input.latent_z.clone()),
388        latent_measure: input.latent_measure,
389        gaussian_frailty_sd: input.gaussian_frailty_sd,
390        base_link: input.base_link.clone(),
391        marginal_design: input.marginal_design.clone(),
392        logslope_design: input.logslope_design.clone(),
393        score_warp,
394        link_dev,
395        policy: policy.clone(),
396        cell_moment_lru: new_cell_moment_lru_cache(&policy),
397        cell_moment_cache_stats: new_cell_moment_cache_stats(),
398        intercept_warm_starts: None,
399        auto_subsample_phase_counter: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
400        auto_subsample_last_rho: Arc::new(std::sync::Mutex::new(None)),
401    };
402    let slices = block_slices(&family);
403    let primary = primary_slices(&slices);
404    let mut block_states = vec![
405        ParameterBlockState {
406            beta: input.marginal_beta.clone(),
407            eta: input.marginal_eta.clone(),
408        },
409        ParameterBlockState {
410            beta: input.logslope_beta.clone(),
411            eta: input.slope.clone(),
412        },
413    ];
414    if let Some(beta) = input.score_warp_beta {
415        block_states.push(ParameterBlockState {
416            beta: beta.clone(),
417            // The exact row program consumes the fitted flex coefficient
418            // vector directly.  Its framework-level block eta is shape-only
419            // state and never enters the likelihood geometry.
420            eta: Array1::zeros(n),
421        });
422    }
423    if let Some(beta) = input.link_deviation_beta {
424        block_states.push(ParameterBlockState {
425            beta: beta.clone(),
426            eta: Array1::zeros(n),
427        });
428    }
429    family.validate_exact_block_state_shapes(&block_states)?;
430
431    let mut rows = Vec::with_capacity(n);
432    for row in 0..n {
433        let row_context = family.build_row_exact_context_with_stats_and_cell_cache(
434            row,
435            &block_states,
436            None,
437            false,
438        )?;
439        let (negative_log_likelihood, nll_score, observed_hessian) = family
440            .compute_row_primary_gradient_hessian(row, &block_states, &primary, &row_context)?;
441        if nll_score.len() != primary.total
442            || observed_hessian.dim() != (primary.total, primary.total)
443            || !negative_log_likelihood.is_finite()
444            || nll_score.iter().any(|value| !value.is_finite())
445            || observed_hessian.iter().any(|value| !value.is_finite())
446        {
447            return Err(format!(
448                "saved BMS ALO row {row} returned invalid local geometry: nll={negative_log_likelihood}, score={}, hessian={}x{}, expected primary width {}",
449                nll_score.len(),
450                observed_hessian.nrows(),
451                observed_hessian.ncols(),
452                primary.total,
453            ));
454        }
455        let mut coordinate_values = Array1::<f64>::zeros(primary.total);
456        coordinate_values[primary.q] = input.marginal_eta[row];
457        coordinate_values[primary.logslope] = input.slope[row];
458        if let (Some(range), Some(beta)) = (primary.h.as_ref(), input.score_warp_beta) {
459            coordinate_values
460                .slice_mut(ndarray::s![range.clone()])
461                .assign(beta);
462        }
463        if let (Some(range), Some(beta)) = (primary.w.as_ref(), input.link_deviation_beta) {
464            coordinate_values
465                .slice_mut(ndarray::s![range.clone()])
466                .assign(beta);
467        }
468        rows.push(BernoulliMarginalSlopeSavedAloRowGeometry {
469            nll_score,
470            observed_hessian,
471            coordinate_values,
472        });
473    }
474    Ok(BernoulliMarginalSlopeSavedAloReplay {
475        rows,
476        score_warp_dimension,
477        link_deviation_dimension,
478    })
479}
480
481#[cfg(test)]
482mod tests {
483    use super::*;
484    use gam_linalg::matrix::DenseDesignMatrix;
485    use gam_math::probability::{normal_cdf, normal_pdf};
486    use gam_problem::StandardLink;
487
488    fn assert_close(label: &str, actual: f64, expected: f64, tolerance: f64) {
489        assert!(
490            (actual - expected).abs() <= tolerance,
491            "{label}: actual={actual:.16e}, expected={expected:.16e}, tolerance={tolerance:.3e}"
492        );
493    }
494
495    #[test]
496    fn rigid_saved_alo_geometry_matches_independent_probit_chain_rule() {
497        let marginal_eta: f64 = 0.35;
498        let slope: f64 = -0.6;
499        let latent_z: f64 = 0.8;
500        let response: f64 = 1.0;
501        let weight: f64 = 1.7;
502        let scale: f64 = 0.75;
503        let geometry =
504            bernoulli_marginal_slope_alo_row_geometry(BernoulliMarginalSlopeAloRowInput {
505                base_link: &InverseLink::Standard(StandardLink::Probit),
506                marginal_eta,
507                slope,
508                latent_z,
509                response,
510                prior_weight: weight,
511                probit_frailty_scale: scale,
512            })
513            .expect("rigid saved marginal-slope row must replay");
514
515        // Independent closed form for eta(q, g) = q sqrt(1 + (s g)^2) + s g z.
516        // At an interior probit marginal map q == marginal_eta exactly.
517        let sg = scale * slope;
518        let c = (1.0 + sg * sg).sqrt();
519        let eta = marginal_eta * c + sg * latent_z;
520        let sign = 2.0 * response - 1.0;
521        let margin = sign * eta;
522        let cdf = normal_cdf(margin);
523        let mills = normal_pdf(margin) / cdf;
524        let nll_first_eta = -weight * sign * mills;
525        let nll_second_eta = weight * mills * (margin + mills);
526
527        let eta_q = c;
528        let eta_g = marginal_eta * scale * scale * slope / c + scale * latent_z;
529        let eta_qg = scale * scale * slope / c;
530        let eta_gg = marginal_eta * scale * scale / c.powi(3);
531        let expected_score = [nll_first_eta * eta_q, nll_first_eta * eta_g];
532        let expected_hessian = [
533            [
534                nll_second_eta * eta_q * eta_q,
535                nll_second_eta * eta_q * eta_g + nll_first_eta * eta_qg,
536            ],
537            [
538                nll_second_eta * eta_q * eta_g + nll_first_eta * eta_qg,
539                nll_second_eta * eta_g * eta_g + nll_first_eta * eta_gg,
540            ],
541        ];
542
543        assert_close(
544            "negative log likelihood",
545            geometry.negative_log_likelihood,
546            -weight * cdf.ln(),
547            2e-13,
548        );
549        for axis in 0..2 {
550            assert_close(
551                &format!("score[{axis}]"),
552                geometry.nll_score[axis],
553                expected_score[axis],
554                2e-12,
555            );
556            for other in 0..2 {
557                assert_close(
558                    &format!("hessian[{axis},{other}]"),
559                    geometry.observed_hessian[axis][other],
560                    expected_hessian[axis][other],
561                    3e-12,
562                );
563            }
564        }
565
566        let score_meat = geometry.nll_score[0] * geometry.nll_score[0];
567        assert!(
568            (geometry.observed_hessian[0][0] - score_meat).abs() > 1e-3,
569            "observed Hessian and empirical score meat must remain distinct"
570        );
571    }
572
573    fn independent_empirical_score_warp_nll(point: [f64; 3]) -> f64 {
574        let [marginal_eta, slope, score_beta] = point;
575        let nodes = [-0.8_f64, 0.9_f64];
576        let grid_weights = [0.35_f64, 0.65_f64];
577        let target = normal_cdf(marginal_eta);
578        let calibration = |intercept: f64| {
579            nodes
580                .iter()
581                .zip(grid_weights.iter())
582                .map(|(&z, &weight)| weight * normal_cdf(intercept + slope * (z + score_beta * z)))
583                .sum::<f64>()
584                - target
585        };
586        let mut lower = -40.0_f64;
587        let mut upper = 40.0_f64;
588        assert!(calibration(lower) < 0.0 && calibration(upper) > 0.0);
589        for _iteration in 0..180 {
590            let midpoint = 0.5 * (lower + upper);
591            if calibration(midpoint) < 0.0 {
592                lower = midpoint;
593            } else {
594                upper = midpoint;
595            }
596        }
597        let intercept = 0.5 * (lower + upper);
598        let observed_z = 0.25_f64;
599        let observed_eta = intercept + slope * (observed_z + score_beta * observed_z);
600        -1.3 * normal_cdf(observed_eta).ln()
601    }
602
603    #[test]
604    fn empirical_flex_saved_alo_matches_independent_resolved_likelihood_oracle() {
605        let marginal_eta = 0.2_f64;
606        let slope = -0.35_f64;
607        let score_beta = 0.12_f64;
608        let score_runtime = SavedCompiledFlexBlock {
609            kernel: crate::cubic_cell_kernel::ANCHORED_DEVIATION_KERNEL.to_string(),
610            breakpoints: vec![-2.0, 2.0],
611            basis_dim: 1,
612            // One frozen local cubic representing h(z)=z on the entire
613            // empirical/observed support.  The independent oracle above uses
614            // the global expression directly and never calls this runtime.
615            span_c0: vec![vec![-2.0]],
616            span_c1: vec![vec![1.0]],
617            span_c2: vec![vec![0.0]],
618            span_c3: vec![vec![0.0]],
619            anchor_correction: None,
620            anchor_components: Vec::new(),
621        };
622        let marginal_design = DesignMatrix::Dense(DenseDesignMatrix::from(Array2::ones((1, 1))));
623        let logslope_design = DesignMatrix::Dense(DenseDesignMatrix::from(Array2::ones((1, 1))));
624        let marginal_beta = Array1::from_vec(vec![marginal_eta]);
625        let logslope_beta = Array1::from_vec(vec![slope]);
626        let score_warp_beta = Array1::from_vec(vec![score_beta]);
627        let marginal_rows = Array1::from_vec(vec![marginal_eta]);
628        let slope_rows = Array1::from_vec(vec![slope]);
629        let latent_z = Array1::from_vec(vec![0.25]);
630        let response = Array1::from_vec(vec![1.0]);
631        let prior_weights = Array1::from_vec(vec![1.3]);
632        let replay =
633            replay_saved_bernoulli_marginal_slope_alo(BernoulliMarginalSlopeSavedAloReplayInput {
634                base_link: &InverseLink::Standard(StandardLink::Probit),
635                marginal_design: &marginal_design,
636                logslope_design: &logslope_design,
637                marginal_beta: &marginal_beta,
638                logslope_beta: &logslope_beta,
639                score_warp_beta: Some(&score_warp_beta),
640                link_deviation_beta: None,
641                marginal_eta: &marginal_rows,
642                slope: &slope_rows,
643                latent_z: &latent_z,
644                response: &response,
645                prior_weights: &prior_weights,
646                latent_measure: LatentMeasureKind::GlobalEmpirical {
647                    grid: super::super::EmpiricalZGrid::new(
648                        vec![-0.8, 0.9],
649                        vec![0.35, 0.65],
650                        "saved ALO empirical-flex oracle",
651                    )
652                    .expect("valid empirical grid"),
653                },
654                gaussian_frailty_sd: None,
655                score_warp_runtime: Some(&score_runtime),
656                link_deviation_runtime: None,
657                score_warp_anchor_rows: None,
658                link_deviation_anchor_rows: None,
659            })
660            .expect("saved empirical-flex row must replay");
661        assert_eq!(replay.score_warp_dimension, 1);
662        assert_eq!(replay.link_deviation_dimension, 0);
663        let row = &replay.rows[0];
664        assert_eq!(
665            row.coordinate_values.to_vec(),
666            vec![marginal_eta, slope, score_beta]
667        );
668
669        let point = [marginal_eta, slope, score_beta];
670        let gradient_step = 2.0e-5_f64;
671        for axis in 0..3 {
672            let mut plus = point;
673            let mut minus = point;
674            plus[axis] += gradient_step;
675            minus[axis] -= gradient_step;
676            let expected = (independent_empirical_score_warp_nll(plus)
677                - independent_empirical_score_warp_nll(minus))
678                / (2.0 * gradient_step);
679            assert_close(
680                &format!("empirical-flex score[{axis}]"),
681                row.nll_score[axis],
682                expected,
683                3.0e-7,
684            );
685        }
686
687        let hessian_step = 3.0e-4_f64;
688        let center = independent_empirical_score_warp_nll(point);
689        for first in 0..3 {
690            for second in first..3 {
691                let expected = if first == second {
692                    let mut plus = point;
693                    let mut minus = point;
694                    plus[first] += hessian_step;
695                    minus[first] -= hessian_step;
696                    (independent_empirical_score_warp_nll(plus) - 2.0 * center
697                        + independent_empirical_score_warp_nll(minus))
698                        / hessian_step.powi(2)
699                } else {
700                    let mut plus_plus = point;
701                    let mut plus_minus = point;
702                    let mut minus_plus = point;
703                    let mut minus_minus = point;
704                    plus_plus[first] += hessian_step;
705                    plus_plus[second] += hessian_step;
706                    plus_minus[first] += hessian_step;
707                    plus_minus[second] -= hessian_step;
708                    minus_plus[first] -= hessian_step;
709                    minus_plus[second] += hessian_step;
710                    minus_minus[first] -= hessian_step;
711                    minus_minus[second] -= hessian_step;
712                    (independent_empirical_score_warp_nll(plus_plus)
713                        - independent_empirical_score_warp_nll(plus_minus)
714                        - independent_empirical_score_warp_nll(minus_plus)
715                        + independent_empirical_score_warp_nll(minus_minus))
716                        / (4.0 * hessian_step.powi(2))
717                };
718                assert_close(
719                    &format!("empirical-flex Hessian[{first},{second}]"),
720                    row.observed_hessian[[first, second]],
721                    expected,
722                    4.0e-5,
723                );
724                assert_close(
725                    &format!("empirical-flex symmetry[{second},{first}]"),
726                    row.observed_hessian[[second, first]],
727                    expected,
728                    4.0e-5,
729                );
730            }
731        }
732        let score_meat = row.nll_score[0] * row.nll_score[0];
733        assert!(
734            (row.observed_hessian[[0, 0]] - score_meat).abs() > 1.0e-3,
735            "observed Hessian W and empirical score meat C must remain separate"
736        );
737    }
738}