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
16pub 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#[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
36pub 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#[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#[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
250pub(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 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 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 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}