use super::*;
use crate::bms::signed_probit_neglog_unary_stack;
use crate::survival::marginal_slope::timewiggle_geometry::{
TimewiggleBasisDerivativeRows, TimewiggleQBaseValues, timewiggle_q_from_basis_derivative_rows,
};
use gam_math::jet_scalar::{
DynamicJetArena, DynamicOneSeed, DynamicOrder2, DynamicOrder2Accumulator, DynamicOrder2Term,
OneSeed, Order2, RuntimeJetScalar,
};
use gam_math::jet_tower::Tower2;
#[derive(Clone, Debug)]
pub(crate) struct FlexTimepointBasePack {
pub(crate) eta: f64,
pub(crate) chi: f64,
pub(crate) d: f64,
pub(crate) eta_u: Vec<f64>,
pub(crate) eta_uv: Vec<f64>,
pub(crate) chi_u: Vec<f64>,
pub(crate) chi_uv: Vec<f64>,
pub(crate) d_u: Vec<f64>,
pub(crate) d_uv: Vec<f64>,
}
#[derive(Clone, Debug)]
pub(crate) struct FlexTimepointDirectionalPack {
pub(crate) eta_uv_dir: Vec<f64>,
pub(crate) eta_u_dir: Vec<f64>,
pub(crate) chi_u_dir: Vec<f64>,
pub(crate) chi_uv_dir: Vec<f64>,
pub(crate) d_u_dir: Vec<f64>,
pub(crate) d_uv_dir: Vec<f64>,
}
#[derive(Clone, Debug)]
pub(crate) struct FlexTimepointBidirectionalPack {
pub(crate) eta_uv_uv: Vec<f64>,
pub(crate) chi_uv_uv: Vec<f64>,
pub(crate) d_uv_uv: Vec<f64>,
}
#[derive(Clone, Copy, Debug, Default)]
pub(crate) struct FlexFamilyRowDirection {
pub(crate) entry: f64,
pub(crate) exit: f64,
pub(crate) derivative_exit: f64,
pub(crate) probit_scale: f64,
}
#[derive(Clone, Debug)]
pub(crate) struct FlexFamilyCoefficientTerms {
pub(crate) objective: f64,
pub(crate) gradient: Array1<f64>,
pub(crate) hessian: Array2<f64>,
}
#[derive(Clone, Debug)]
pub(crate) struct FlexFamilyDirectionRowTerms {
pub(crate) first: FlexFamilyCoefficientTerms,
pub(crate) second: FlexFamilyCoefficientTerms,
pub(crate) directional: Option<FlexFamilyCoefficientTerms>,
}
pub(crate) fn pack_flex_timepoint_base(base: &SurvivalFlexTimepointExact) -> FlexTimepointBasePack {
FlexTimepointBasePack {
eta: base.eta,
chi: base.chi,
d: base.d,
eta_u: base.eta_u.to_vec(),
eta_uv: base.eta_uv.iter().copied().collect(),
chi_u: base.chi_u.to_vec(),
chi_uv: base.chi_uv.iter().copied().collect(),
d_u: base.d_u.to_vec(),
d_uv: base.d_uv.iter().copied().collect(),
}
}
thread_local! {
static FLEX_THIRD_JET_ARENA: std::cell::RefCell<DynamicJetArena> =
std::cell::RefCell::new(DynamicJetArena::new());
}
pub(crate) fn with_flex_third_jet_arena<R>(evaluate: impl FnOnce(&mut DynamicJetArena) -> R) -> R {
FLEX_THIRD_JET_ARENA.with(|workspace| {
let mut arena = workspace.borrow_mut();
arena.reset();
let result = evaluate(&mut arena);
arena.reset();
result
})
}
#[inline]
fn surv_stack(eta: f64) -> Result<[f64; 5], String> {
let signed_margin = -eta;
if signed_margin != f64::INFINITY && !signed_margin.is_finite() {
return Err(format!(
"non-finite signed margin in exact probit derivative helper: {signed_margin}"
));
}
let m_stack = signed_probit_neglog_unary_stack(signed_margin, -1.0);
Ok([m_stack[0], -m_stack[1], m_stack[2], -m_stack[3], m_stack[4]])
}
#[inline]
fn ln_stack(x: f64) -> [f64; 5] {
let inv = 1.0 / x;
let inv2 = inv * inv;
[x.ln(), inv, -inv2, 2.0 * inv2 * inv, -6.0 * inv2 * inv2]
}
trait FlexJet: JetField + Clone {
const ORDER: usize;
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self;
}
impl FlexJet for f64 {
const ORDER: usize = 0;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
factors[0] * self
}
}
impl<J: FlexJet> FlexJet for Dual2<J> {
const ORDER: usize = if J::ORDER >= 2 { 4 } else { J::ORDER + 2 };
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
let shifted = |outer_order: usize| {
std::array::from_fn(|inner_order| {
factors
.get(inner_order + outer_order)
.copied()
.unwrap_or(0.0)
})
};
Self {
v: self.v.scale_homogeneous_orders(shifted(0)),
g: self.g.scale_homogeneous_orders(shifted(1)),
h: self.h.scale_homogeneous_orders(shifted(2)),
}
}
}
const FLEX_OUTER_SOURCE_COUNT: usize = 6;
const FLEX_OUTER_TERM_COUNT: usize = 7;
#[derive(Clone, Copy)]
#[repr(usize)]
enum FlexOuterSource {
Eta0 = 0,
Eta1 = 1,
Q1 = 2,
Chi1 = 3,
D1 = 4,
Qd1 = 5,
}
impl FlexOuterSource {
#[inline(always)]
fn index(self) -> usize {
self as usize
}
}
#[derive(Clone, Copy)]
enum FlexOuterTransform {
Compose([f64; 5]),
Square,
}
#[derive(Clone, Copy)]
enum FlexOuterCombine {
Add,
Subtract,
}
#[derive(Clone, Copy)]
struct FlexOuterTerm {
source: FlexOuterSource,
transform: FlexOuterTransform,
scale: f64,
combine: FlexOuterCombine,
}
impl FlexOuterTerm {
#[inline(always)]
fn evaluate<J: FlexJet>(&self, sources: &FlexJetSources<'_, J>) -> J {
let source = sources.get(self.source);
let transformed = match self.transform {
FlexOuterTransform::Compose(stack) => source.compose_unary(stack),
FlexOuterTransform::Square => source.mul(source),
};
transformed.scale(self.scale)
}
#[inline(always)]
fn order2_channels(&self, source_value: f64) -> [f64; 3] {
let mut channels = match self.transform {
FlexOuterTransform::Compose(stack) => [stack[0], stack[1], stack[2]],
FlexOuterTransform::Square => [
source_value * source_value,
source_value + source_value,
2.0,
],
};
for channel in &mut channels {
*channel *= self.scale;
}
channels
}
}
struct FlexJetSources<'a, J> {
eta0: &'a J,
eta1: &'a J,
q1: &'a J,
chi1: &'a J,
d1: &'a J,
qd1: &'a J,
}
impl<'a, J> FlexJetSources<'a, J> {
#[inline(always)]
fn get(&self, source: FlexOuterSource) -> &'a J {
match source {
FlexOuterSource::Eta0 => self.eta0,
FlexOuterSource::Eta1 => self.eta1,
FlexOuterSource::Q1 => self.q1,
FlexOuterSource::Chi1 => self.chi1,
FlexOuterSource::D1 => self.d1,
FlexOuterSource::Qd1 => self.qd1,
}
}
}
struct FlexOuterPlan {
terms: [FlexOuterTerm; FLEX_OUTER_TERM_COUNT],
}
impl FlexOuterPlan {
#[inline(always)]
fn new(
chi1: f64,
d1: f64,
qd1: f64,
surv0: [f64; 5],
surv1: [f64; 5],
wi: f64,
di: f64,
) -> Self {
let wd = wi * di;
Self {
terms: [
FlexOuterTerm {
source: FlexOuterSource::Eta0,
transform: FlexOuterTransform::Compose(surv0),
scale: wi,
combine: FlexOuterCombine::Add,
},
FlexOuterTerm {
source: FlexOuterSource::Eta1,
transform: FlexOuterTransform::Compose(surv1),
scale: -wi * (1.0 - di),
combine: FlexOuterCombine::Add,
},
FlexOuterTerm {
source: FlexOuterSource::Eta1,
transform: FlexOuterTransform::Square,
scale: 0.5 * wd,
combine: FlexOuterCombine::Add,
},
FlexOuterTerm {
source: FlexOuterSource::Q1,
transform: FlexOuterTransform::Square,
scale: 0.5 * wd,
combine: FlexOuterCombine::Add,
},
FlexOuterTerm {
source: FlexOuterSource::Chi1,
transform: FlexOuterTransform::Compose(ln_stack(chi1)),
scale: wd,
combine: FlexOuterCombine::Subtract,
},
FlexOuterTerm {
source: FlexOuterSource::D1,
transform: FlexOuterTransform::Compose(ln_stack(d1)),
scale: wd,
combine: FlexOuterCombine::Add,
},
FlexOuterTerm {
source: FlexOuterSource::Qd1,
transform: FlexOuterTransform::Compose(ln_stack(qd1)),
scale: wd,
combine: FlexOuterCombine::Subtract,
},
],
}
}
#[inline(always)]
fn evaluate<J: FlexJet>(&self, sources: &FlexJetSources<'_, J>) -> J {
let mut output = self.terms[0].evaluate(sources);
for term in &self.terms[1..] {
let contribution = term.evaluate(sources);
output = match term.combine {
FlexOuterCombine::Add => output.add(&contribution),
FlexOuterCombine::Subtract => output.sub(&contribution),
};
}
output
}
#[inline(always)]
fn compile_order2(&self, source_values: [f64; FLEX_OUTER_SOURCE_COUNT]) -> FlexOrder2Plan {
let mut value = 0.0;
let mut derivatives = [[0.0; 2]; FLEX_OUTER_SOURCE_COUNT];
for term in &self.terms {
let channels = term.order2_channels(source_values[term.source.index()]);
let sign = match term.combine {
FlexOuterCombine::Add => 1.0,
FlexOuterCombine::Subtract => -1.0,
};
value += sign * channels[0];
derivatives[term.source.index()][0] += sign * channels[1];
derivatives[term.source.index()][1] += sign * channels[2];
}
FlexOrder2Plan { value, derivatives }
}
}
#[inline]
fn flex_row_nll<J: FlexJet>(
eta0: &J,
eta1: &J,
chi1: &J,
d1: &J,
q1: &J,
qd1: &J,
surv0: [f64; 5],
surv1: [f64; 5],
wi: f64,
di: f64,
) -> J {
FlexOuterPlan::new(chi1.value(), d1.value(), qd1.value(), surv0, surv1, wi, di).evaluate(
&FlexJetSources {
eta0,
eta1,
q1,
chi1,
d1,
qd1,
},
)
}
struct FlexOrder2Plan {
value: f64,
derivatives: [[f64; 2]; FLEX_OUTER_SOURCE_COUNT],
}
struct FlexOrder2View<'a> {
value: f64,
gradient: ndarray::ArrayView1<'a, f64>,
hessian: ndarray::ArrayView2<'a, f64>,
}
struct FlexOrder2Inputs<'a> {
eta0: FlexOrder2View<'a>,
eta1: FlexOrder2View<'a>,
q1: (f64, usize),
chi1: FlexOrder2View<'a>,
d1: FlexOrder2View<'a>,
qd1: (f64, usize),
}
enum FlexOrder2Term<'a> {
Dense {
outer: [f64; 2],
gradient: ndarray::ArrayView1<'a, f64>,
hessian: ndarray::ArrayView2<'a, f64>,
},
Axis {
outer: [f64; 2],
axis: usize,
},
}
impl DynamicOrder2Term for FlexOrder2Term<'_> {
#[inline(always)]
fn outer_first(&self) -> f64 {
match self {
Self::Dense { outer, .. } | Self::Axis { outer, .. } => outer[0],
}
}
#[inline(always)]
fn outer_second(&self) -> f64 {
match self {
Self::Dense { outer, .. } | Self::Axis { outer, .. } => outer[1],
}
}
#[inline(always)]
fn inner_gradient(&self, axis: usize) -> f64 {
match self {
Self::Dense { gradient, .. } => gradient[axis],
Self::Axis { axis: source, .. } => {
if axis == *source {
1.0
} else {
0.0
}
}
}
}
#[inline(always)]
fn inner_hessian(&self, row: usize, column: usize) -> f64 {
match self {
Self::Dense { hessian, .. } => hessian[[row, column]],
Self::Axis { .. } => 0.0,
}
}
}
#[inline]
fn lower_flex_outer_plan_order2(
plan: &FlexOuterPlan,
inputs: FlexOrder2Inputs<'_>,
dimension: usize,
) -> (f64, Vec<f64>, Vec<f64>) {
let source_values = [
inputs.eta0.value,
inputs.eta1.value,
inputs.q1.0,
inputs.chi1.value,
inputs.d1.value,
inputs.qd1.0,
];
let compiled = plan.compile_order2(source_values);
let terms = [
FlexOrder2Term::Dense {
outer: compiled.derivatives[FlexOuterSource::Eta0.index()],
gradient: inputs.eta0.gradient,
hessian: inputs.eta0.hessian,
},
FlexOrder2Term::Dense {
outer: compiled.derivatives[FlexOuterSource::Eta1.index()],
gradient: inputs.eta1.gradient,
hessian: inputs.eta1.hessian,
},
FlexOrder2Term::Axis {
outer: compiled.derivatives[FlexOuterSource::Q1.index()],
axis: inputs.q1.1,
},
FlexOrder2Term::Dense {
outer: compiled.derivatives[FlexOuterSource::Chi1.index()],
gradient: inputs.chi1.gradient,
hessian: inputs.chi1.hessian,
},
FlexOrder2Term::Dense {
outer: compiled.derivatives[FlexOuterSource::D1.index()],
gradient: inputs.d1.gradient,
hessian: inputs.d1.hessian,
},
FlexOrder2Term::Axis {
outer: compiled.derivatives[FlexOuterSource::Qd1.index()],
axis: inputs.qd1.1,
},
];
DynamicOrder2Accumulator::from_composed_sum(dimension, compiled.value, &terms).into_channels()
}
#[derive(Clone)]
struct Jet2 {
v: f64,
g: Vec<f64>,
h: Vec<f64>,
}
impl Jet2 {
fn from_parts(v: f64, g: &[f64], h: &[f64]) -> Self {
let p = g.len();
let hv = if h.is_empty() {
vec![0.0; p * p]
} else {
assert_eq!(h.len(), p * p, "Jet2::from_parts Hessian length");
h.to_vec()
};
Jet2 {
v,
g: g.to_vec(),
h: hv,
}
}
fn primary(x: f64, axis: usize, p: usize) -> Self {
let mut g = vec![0.0; p];
if axis < p {
g[axis] = 1.0;
}
Jet2 {
v: x,
g,
h: vec![0.0; p * p],
}
}
#[inline]
fn p(&self) -> usize {
self.g.len()
}
#[inline]
fn scale_homogeneous_from(&self, offset: usize, factors: [f64; 5]) -> Self {
assert!(offset + 2 < factors.len());
Jet2 {
v: factors[offset] * self.v,
g: self
.g
.iter()
.map(|&channel| factors[offset + 1] * channel)
.collect(),
h: self
.h
.iter()
.map(|&channel| factors[offset + 2] * channel)
.collect(),
}
}
}
impl JetField for Jet2 {
#[inline]
fn value(&self) -> f64 {
self.v
}
fn add(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
let mut h = vec![0.0; p * p];
for i in 0..p {
g[i] = self.g[i] + o.g[i];
}
for k in 0..p * p {
h[k] = self.h[k] + o.h[k];
}
Jet2 {
v: self.v + o.v,
g,
h,
}
}
fn sub(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
let mut h = vec![0.0; p * p];
for i in 0..p {
g[i] = self.g[i] - o.g[i];
}
for k in 0..p * p {
h[k] = self.h[k] - o.h[k];
}
Jet2 {
v: self.v - o.v,
g,
h,
}
}
fn mul(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
let mut h = vec![0.0; p * p];
for i in 0..p {
g[i] = self.v * o.g[i] + self.g[i] * o.v;
}
for i in 0..p {
for j in 0..p {
h[i * p + j] = self.v * o.h[i * p + j]
+ self.g[i] * o.g[j]
+ self.g[j] * o.g[i]
+ self.h[i * p + j] * o.v;
}
}
Jet2 {
v: self.v * o.v,
g,
h,
}
}
fn scale(&self, s: f64) -> Self {
Jet2 {
v: self.v * s,
g: self.g.iter().map(|&x| x * s).collect(),
h: self.h.iter().map(|&x| x * s).collect(),
}
}
#[inline]
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
let p = self.p();
let (f, f1, f2) = (d[0], d[1], d[2]);
let mut g = vec![0.0; p];
let mut h = vec![0.0; p * p];
for i in 0..p {
g[i] = f1 * self.g[i];
}
for i in 0..p {
for j in 0..p {
h[i * p + j] = f2 * self.g[i] * self.g[j] + f1 * self.h[i * p + j];
}
}
Jet2 { v: f, g, h }
}
}
impl FlexJet for Jet2 {
const ORDER: usize = 2;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
self.scale_homogeneous_from(0, factors)
}
}
#[derive(Clone)]
struct Jet1 {
v: f64,
g: Vec<f64>,
}
impl Jet1 {
fn from_view(v: f64, g: ndarray::ArrayView1<'_, f64>) -> Self {
Jet1 {
v,
g: g.iter().copied().collect(),
}
}
fn primary(x: f64, axis: usize, p: usize) -> Self {
let mut g = vec![0.0; p];
if axis < p {
g[axis] = 1.0;
}
Jet1 { v: x, g }
}
#[inline]
fn p(&self) -> usize {
self.g.len()
}
}
impl JetField for Jet1 {
#[inline]
fn value(&self) -> f64 {
self.v
}
fn add(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
for i in 0..p {
g[i] = self.g[i] + o.g[i];
}
Jet1 { v: self.v + o.v, g }
}
fn sub(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
for i in 0..p {
g[i] = self.g[i] - o.g[i];
}
Jet1 { v: self.v - o.v, g }
}
fn mul(&self, o: &Self) -> Self {
let p = self.p();
let mut g = vec![0.0; p];
for i in 0..p {
g[i] = self.v * o.g[i] + self.g[i] * o.v;
}
Jet1 { v: self.v * o.v, g }
}
fn scale(&self, s: f64) -> Self {
Jet1 {
v: self.v * s,
g: self.g.iter().map(|&x| x * s).collect(),
}
}
#[inline]
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
let p = self.p();
let (f, f1) = (d[0], d[1]);
let mut g = vec![0.0; p];
for i in 0..p {
g[i] = f1 * self.g[i];
}
Jet1 { v: f, g }
}
}
impl FlexJet for Jet1 {
const ORDER: usize = 1;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
Jet1 {
v: factors[0] * self.v,
g: self.g.iter().map(|&channel| factors[1] * channel).collect(),
}
}
}
#[derive(Clone)]
struct Jet3 {
base: Jet2,
eps: Jet2,
}
impl Jet3 {
fn primary(x: f64, axis: usize, p: usize, dir_axis: f64) -> Self {
Jet3 {
base: Jet2::primary(x, axis, p),
eps: Jet2::from_parts(dir_axis, &vec![0.0; p], &[]),
}
}
fn contracted_third(&self) -> Vec<f64> {
self.eps.h.clone()
}
}
impl JetField for Jet3 {
#[inline]
fn value(&self) -> f64 {
self.base.v
}
fn add(&self, o: &Self) -> Self {
Jet3 {
base: self.base.add(&o.base),
eps: self.eps.add(&o.eps),
}
}
fn sub(&self, o: &Self) -> Self {
Jet3 {
base: self.base.sub(&o.base),
eps: self.eps.sub(&o.eps),
}
}
fn mul(&self, o: &Self) -> Self {
Jet3 {
base: self.base.mul(&o.base),
eps: self.base.mul(&o.eps).add(&self.eps.mul(&o.base)),
}
}
fn scale(&self, s: f64) -> Self {
Jet3 {
base: self.base.scale(s),
eps: self.eps.scale(s),
}
}
#[inline]
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
let base = self.base.compose_unary([d[0], d[1], d[2], d[3], d[4]]);
let fprime = self.base.compose_unary([d[1], d[2], d[3], d[4], d[4]]);
let eps = fprime.mul(&self.eps);
Jet3 { base, eps }
}
}
impl FlexJet for Jet3 {
const ORDER: usize = 3;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
Jet3 {
base: self.base.scale_homogeneous_from(0, factors),
eps: self.eps.scale_homogeneous_from(1, factors),
}
}
}
trait FlexCoefficientJet: FlexJet {
fn affine(
value: f64,
gradient: Vec<f64>,
directional_value: f64,
directional_gradient: Vec<f64>,
) -> Self;
fn constant(value: f64, dimension: usize) -> Self {
Self::affine(value, vec![0.0; dimension], 0.0, vec![0.0; dimension])
}
fn owned_base_terms(&self) -> FlexFamilyCoefficientTerms;
fn owned_directional_terms(&self) -> Option<FlexFamilyCoefficientTerms>;
}
fn owned_jet2_terms(jet: &Jet2) -> FlexFamilyCoefficientTerms {
let dimension = jet.g.len();
FlexFamilyCoefficientTerms {
objective: jet.v,
gradient: Array1::from_vec(jet.g.clone()),
hessian: Array2::from_shape_vec((dimension, dimension), jet.h.clone())
.expect("Jet2 coefficient Hessian shape invariant"),
}
}
impl FlexCoefficientJet for Jet2 {
fn affine(
value: f64,
gradient: Vec<f64>,
directional_value: f64,
directional_gradient: Vec<f64>,
) -> Self {
assert!(
directional_value == 0.0 && directional_gradient.iter().all(|value| *value == 0.0),
"Jet2 cannot erase a coefficient-direction seed; use Jet3"
);
Self::from_parts(value, &gradient, &[])
}
fn owned_base_terms(&self) -> FlexFamilyCoefficientTerms {
owned_jet2_terms(self)
}
fn owned_directional_terms(&self) -> Option<FlexFamilyCoefficientTerms> {
None
}
}
impl FlexCoefficientJet for Jet3 {
fn affine(
value: f64,
gradient: Vec<f64>,
directional_value: f64,
directional_gradient: Vec<f64>,
) -> Self {
Self {
base: Jet2::from_parts(value, &gradient, &[]),
eps: Jet2::from_parts(directional_value, &directional_gradient, &[]),
}
}
fn owned_base_terms(&self) -> FlexFamilyCoefficientTerms {
owned_jet2_terms(&self.base)
}
fn owned_directional_terms(&self) -> Option<FlexFamilyCoefficientTerms> {
Some(owned_jet2_terms(&self.eps))
}
}
#[derive(Clone, Copy)]
struct ArenaJet3<'arena> {
arena: &'arena DynamicJetArena,
inner: DynamicOneSeed<'arena>,
}
impl<'arena> ArenaJet3<'arena> {
#[inline]
fn primary(
x: f64,
axis: usize,
p: usize,
direction: f64,
arena: &'arena DynamicJetArena,
) -> Self {
let inner = if axis < p {
DynamicOneSeed::seed_direction(x, axis, direction, p, arena)
} else {
DynamicOneSeed {
base: DynamicOrder2::constant(x, p, arena),
eps: DynamicOrder2::constant(direction, p, arena),
}
};
Self { arena, inner }
}
}
impl JetField for ArenaJet3<'_> {
#[inline(always)]
fn value(&self) -> f64 {
self.inner.value()
}
#[inline(always)]
fn add(&self, other: &Self) -> Self {
Self {
arena: self.arena,
inner: self.inner.add(&other.inner),
}
}
#[inline(always)]
fn sub(&self, other: &Self) -> Self {
Self {
arena: self.arena,
inner: self.inner.sub(&other.inner),
}
}
#[inline(always)]
fn mul(&self, other: &Self) -> Self {
Self {
arena: self.arena,
inner: self.inner.mul(&other.inner),
}
}
#[inline(always)]
fn neg(&self) -> Self {
Self {
arena: self.arena,
inner: self.inner.neg(),
}
}
#[inline(always)]
fn scale(&self, scale: f64) -> Self {
Self {
arena: self.arena,
inner: self.inner.scale(scale),
}
}
#[inline(always)]
fn compose_unary(&self, derivatives: [f64; 5]) -> Self {
Self {
arena: self.arena,
inner: self.inner.compose_unary(derivatives),
}
}
}
impl FlexJet for ArenaJet3<'_> {
const ORDER: usize = 3;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
let dimension = self.inner.base.g.len();
let scale_order2 = |channels: &DynamicOrder2<'_>, offset: usize| {
DynamicOrder2::from_channel_functions(
factors[offset] * channels.v,
dimension,
self.arena,
|axis| factors[offset + 1] * channels.g[axis],
|row, column| factors[offset + 2] * channels.h[row * dimension + column],
)
};
Self {
arena: self.arena,
inner: DynamicOneSeed {
base: scale_order2(&self.inner.base, 0),
eps: scale_order2(&self.inner.eps, 1),
},
}
}
}
#[derive(Clone, Copy)]
struct FixedJet3<const K: usize> {
inner: OneSeed<K>,
}
impl<const K: usize> FixedJet3<K> {
#[inline]
fn primary(x: f64, axis: usize, p: usize, direction: f64) -> Self {
assert_eq!(p, K, "fixed FLEX Jet3 width mismatch");
let inner = if axis < K {
OneSeed::seed_direction(x, axis, direction)
} else {
OneSeed {
base: <Order2<K> as gam_math::jet_scalar::JetScalar<K>>::constant(x),
eps: <Order2<K> as gam_math::jet_scalar::JetScalar<K>>::constant(direction),
}
};
Self { inner }
}
}
impl<const K: usize> JetField for FixedJet3<K> {
#[inline(always)]
fn value(&self) -> f64 {
self.inner.value()
}
#[inline(always)]
fn add(&self, other: &Self) -> Self {
Self {
inner: self.inner.add(&other.inner),
}
}
#[inline(always)]
fn sub(&self, other: &Self) -> Self {
Self {
inner: self.inner.sub(&other.inner),
}
}
#[inline(always)]
fn mul(&self, other: &Self) -> Self {
Self {
inner: self.inner.mul(&other.inner),
}
}
#[inline(always)]
fn neg(&self) -> Self {
Self {
inner: self.inner.neg(),
}
}
#[inline(always)]
fn scale(&self, scale: f64) -> Self {
Self {
inner: self.inner.scale(scale),
}
}
#[inline(always)]
fn compose_unary(&self, derivatives: [f64; 5]) -> Self {
Self {
inner: self.inner.compose_unary(derivatives),
}
}
}
impl<const K: usize> FlexJet for FixedJet3<K> {
const ORDER: usize = 3;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
let scale_order2 = |channels: &Order2<K>, offset: usize| {
let mut scaled = Tower2::<K>::zero();
scaled.v = factors[offset] * channels.0.v;
for axis in 0..K {
scaled.g[axis] = factors[offset + 1] * channels.0.g[axis];
}
for row in 0..K {
for column in 0..K {
scaled.h[row][column] = factors[offset + 2] * channels.0.h[row][column];
}
}
Order2(scaled)
};
Self {
inner: OneSeed {
base: scale_order2(&self.inner.base, 0),
eps: scale_order2(&self.inner.eps, 1),
},
}
}
}
trait FlexThirdOutput: FlexJet + MomentTerm {
fn pack_timepoint_outputs(
eta: &Self,
chi: &Self,
d: &Self,
) -> (FlexTimepointBasePack, FlexTimepointDirectionalPack);
}
impl FlexThirdOutput for ArenaJet3<'_> {
fn pack_timepoint_outputs(
eta: &Self,
chi: &Self,
d: &Self,
) -> (FlexTimepointBasePack, FlexTimepointDirectionalPack) {
(
FlexTimepointBasePack {
eta: eta.inner.base.v,
chi: chi.inner.base.v,
d: d.inner.base.v,
eta_u: eta.inner.base.g.to_vec(),
eta_uv: eta.inner.base.h.to_vec(),
chi_u: chi.inner.base.g.to_vec(),
chi_uv: chi.inner.base.h.to_vec(),
d_u: d.inner.base.g.to_vec(),
d_uv: d.inner.base.h.to_vec(),
},
FlexTimepointDirectionalPack {
eta_u_dir: eta.inner.eps.g.to_vec(),
eta_uv_dir: eta.inner.eps.h.to_vec(),
chi_u_dir: chi.inner.eps.g.to_vec(),
chi_uv_dir: chi.inner.eps.h.to_vec(),
d_u_dir: d.inner.eps.g.to_vec(),
d_uv_dir: d.inner.eps.h.to_vec(),
},
)
}
}
impl<const K: usize> FlexThirdOutput for FixedJet3<K> {
fn pack_timepoint_outputs(
eta: &Self,
chi: &Self,
d: &Self,
) -> (FlexTimepointBasePack, FlexTimepointDirectionalPack) {
let flatten =
|matrix: &[[f64; K]; K]| matrix.iter().flat_map(|row| row.iter().copied()).collect();
(
FlexTimepointBasePack {
eta: eta.inner.base.0.v,
chi: chi.inner.base.0.v,
d: d.inner.base.0.v,
eta_u: eta.inner.base.0.g.to_vec(),
eta_uv: flatten(&eta.inner.base.0.h),
chi_u: chi.inner.base.0.g.to_vec(),
chi_uv: flatten(&chi.inner.base.0.h),
d_u: d.inner.base.0.g.to_vec(),
d_uv: flatten(&d.inner.base.0.h),
},
FlexTimepointDirectionalPack {
eta_u_dir: eta.inner.eps.0.g.to_vec(),
eta_uv_dir: flatten(&eta.inner.eps.0.h),
chi_u_dir: chi.inner.eps.0.g.to_vec(),
chi_uv_dir: flatten(&chi.inner.eps.0.h),
d_u_dir: d.inner.eps.0.g.to_vec(),
d_uv_dir: flatten(&d.inner.eps.0.h),
},
)
}
}
#[derive(Clone)]
struct Jet4 {
base: Jet2,
eps: Jet2,
del: Jet2,
eps_del: Jet2,
}
impl Jet4 {
fn primary(x: f64, axis: usize, p: usize, du: f64, dv: f64) -> Self {
let zero = vec![0.0; p];
Jet4 {
base: Jet2::primary(x, axis, p),
eps: Jet2::from_parts(du, &zero, &[]),
del: Jet2::from_parts(dv, &zero, &[]),
eps_del: Jet2::from_parts(0.0, &zero, &[]),
}
}
fn contracted_fourth(&self) -> Vec<f64> {
self.eps_del.h.clone()
}
}
impl JetField for Jet4 {
#[inline]
fn value(&self) -> f64 {
self.base.v
}
fn add(&self, o: &Self) -> Self {
Jet4 {
base: self.base.add(&o.base),
eps: self.eps.add(&o.eps),
del: self.del.add(&o.del),
eps_del: self.eps_del.add(&o.eps_del),
}
}
fn sub(&self, o: &Self) -> Self {
Jet4 {
base: self.base.sub(&o.base),
eps: self.eps.sub(&o.eps),
del: self.del.sub(&o.del),
eps_del: self.eps_del.sub(&o.eps_del),
}
}
fn mul(&self, o: &Self) -> Self {
let base = self.base.mul(&o.base);
let eps = self.base.mul(&o.eps).add(&self.eps.mul(&o.base));
let del = self.base.mul(&o.del).add(&self.del.mul(&o.base));
let eps_del = self
.base
.mul(&o.eps_del)
.add(&self.eps.mul(&o.del))
.add(&self.del.mul(&o.eps))
.add(&self.eps_del.mul(&o.base));
Jet4 {
base,
eps,
del,
eps_del,
}
}
fn scale(&self, s: f64) -> Self {
Jet4 {
base: self.base.scale(s),
eps: self.eps.scale(s),
del: self.del.scale(s),
eps_del: self.eps_del.scale(s),
}
}
#[inline]
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
let base = self.base.compose_unary([d[0], d[1], d[2], d[3], d[4]]);
let fprime = self.base.compose_unary([d[1], d[2], d[3], d[4], d[4]]);
let fsecond = self.base.compose_unary([d[2], d[3], d[4], d[4], d[4]]);
let eps = fprime.mul(&self.eps);
let del = fprime.mul(&self.del);
let eps_del = fsecond
.mul(&self.eps)
.mul(&self.del)
.add(&fprime.mul(&self.eps_del));
Jet4 {
base,
eps,
del,
eps_del,
}
}
}
impl FlexJet for Jet4 {
const ORDER: usize = 4;
#[inline]
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
Jet4 {
base: self.base.scale_homogeneous_from(0, factors),
eps: self.eps.scale_homogeneous_from(1, factors),
del: self.del.scale_homogeneous_from(1, factors),
eps_del: self.eps_del.scale_homogeneous_from(2, factors),
}
}
}
#[inline]
fn dot(x: &[f64], y: &[f64]) -> f64 {
x.iter().zip(y.iter()).map(|(&a, &b)| a * b).sum()
}
fn mat_vec(m: &[f64], v: &[f64], p: usize) -> Vec<f64> {
let mut out = vec![0.0; p];
for i in 0..p {
let mut acc = 0.0;
for j in 0..p {
acc += m[i * p + j] * v[j];
}
out[i] = acc;
}
out
}
fn quad_form(m: &[f64], v1: &[f64], v2: &[f64], p: usize) -> f64 {
let mut acc = 0.0;
for i in 0..p {
let mi = &m[i * p..i * p + p];
acc += v1[i] * dot(mi, v2);
}
acc
}
pub(crate) struct FlexRowJet2Channels<'a> {
pub eta0_v: f64,
pub eta0_g: ndarray::ArrayView1<'a, f64>,
pub eta0_h: Option<ndarray::ArrayView2<'a, f64>>,
pub eta1_v: f64,
pub eta1_g: ndarray::ArrayView1<'a, f64>,
pub eta1_h: Option<ndarray::ArrayView2<'a, f64>>,
pub chi1_v: f64,
pub chi1_g: ndarray::ArrayView1<'a, f64>,
pub chi1_h: Option<ndarray::ArrayView2<'a, f64>>,
pub d1_v: f64,
pub d1_g: ndarray::ArrayView1<'a, f64>,
pub d1_h: Option<ndarray::ArrayView2<'a, f64>>,
}
pub(crate) struct FlexThirdPacks<'a> {
pub entry_base: &'a FlexTimepointBasePack,
pub exit_base: &'a FlexTimepointBasePack,
pub entry_ext: &'a FlexTimepointDirectionalPack,
pub exit_ext: &'a FlexTimepointDirectionalPack,
}
pub(crate) struct FlexFourthPacks<'a> {
pub entry_base: &'a FlexTimepointBasePack,
pub exit_base: &'a FlexTimepointBasePack,
pub entry_ext_u: &'a FlexTimepointDirectionalPack,
pub exit_ext_u: &'a FlexTimepointDirectionalPack,
pub entry_ext_v: &'a FlexTimepointDirectionalPack,
pub exit_ext_v: &'a FlexTimepointDirectionalPack,
pub entry_bi: &'a FlexTimepointBidirectionalPack,
pub exit_bi: &'a FlexTimepointBidirectionalPack,
}
impl SurvivalMarginalSlopeFamily {
pub(crate) fn flex_row_nll_value_grad_hess(
&self,
row: usize,
primary: &FlexPrimarySlices,
q1: f64,
qd1: f64,
ch: FlexRowJet2Channels<'_>,
) -> Result<(f64, Array1<f64>, Array2<f64>), String> {
let FlexRowJet2Channels {
eta0_v,
eta0_g,
eta0_h,
eta1_v,
eta1_g,
eta1_h,
chi1_v,
chi1_g,
chi1_h,
d1_v,
d1_g,
d1_h,
} = ch;
let p = primary.total;
let wi = self.weights[row];
let di = self.event[row];
let surv0 = surv_stack(eta0_v)?;
let surv1 = surv_stack(eta1_v)?;
let want_hess = eta1_h.is_some();
if !want_hess {
let eta0 = Jet1::from_view(eta0_v, eta0_g);
let eta1 = Jet1::from_view(eta1_v, eta1_g);
let chi1 = Jet1::from_view(chi1_v, chi1_g);
let d1 = Jet1::from_view(d1_v, d1_g);
let q1j = Jet1::primary(q1, primary.q1, p);
let qd1j = Jet1::primary(qd1, primary.qd1, p);
let out = flex_row_nll(&eta0, &eta1, &chi1, &d1, &q1j, &qd1j, surv0, surv1, wi, di);
let value = out.v + wi * di * std::f64::consts::TAU.ln();
let grad = Array1::from(out.g);
return Ok((value, grad, Array2::zeros((p, p))));
}
let eta0_h = eta0_h.ok_or("flex order-two lowering: missing eta0 Hessian")?;
let chi1_h = chi1_h.ok_or("flex order-two lowering: missing chi1 Hessian")?;
let d1_h = d1_h.ok_or("flex order-two lowering: missing d1 Hessian")?;
let eta1_h = eta1_h.ok_or("flex order-two lowering: missing eta1 Hessian")?;
let plan = FlexOuterPlan::new(chi1_v, d1_v, qd1, surv0, surv1, wi, di);
let (row_value, row_gradient, row_hessian) = lower_flex_outer_plan_order2(
&plan,
FlexOrder2Inputs {
eta0: FlexOrder2View {
value: eta0_v,
gradient: eta0_g,
hessian: eta0_h,
},
eta1: FlexOrder2View {
value: eta1_v,
gradient: eta1_g,
hessian: eta1_h,
},
q1: (q1, primary.q1),
chi1: FlexOrder2View {
value: chi1_v,
gradient: chi1_g,
hessian: chi1_h,
},
d1: FlexOrder2View {
value: d1_v,
gradient: d1_g,
hessian: d1_h,
},
qd1: (qd1, primary.qd1),
},
p,
);
let value = row_value + wi * di * std::f64::consts::TAU.ln();
let grad = Array1::from(row_gradient);
let hess = Array2::from_shape_vec((p, p), row_hessian).map_err(|e| e.to_string())?;
Ok((value, grad, hess))
}
pub(crate) fn flex_row_nll_third_contracted(
&self,
row: usize,
primary: &FlexPrimarySlices,
q1: f64,
qd1: f64,
dir: &[f64],
packs: FlexThirdPacks<'_>,
) -> Result<Array2<f64>, String> {
let FlexThirdPacks {
entry_base,
exit_base,
entry_ext,
exit_ext,
} = packs;
let p = primary.total;
let wi = self.weights[row];
let di = self.event[row];
let surv0 = surv_stack(entry_base.eta)?;
let surv1 = surv_stack(exit_base.eta)?;
let mk =
|base_v: f64, base_g: &[f64], base_h: &[f64], ext_g: &[f64], ext_h: &[f64]| -> Jet3 {
Jet3 {
base: Jet2::from_parts(base_v, base_g, base_h),
eps: Jet2::from_parts(dot(base_g, dir), ext_g, ext_h),
}
};
let eta0 = mk(
entry_base.eta,
&entry_base.eta_u,
&entry_base.eta_uv,
&entry_ext.eta_u_dir,
&entry_ext.eta_uv_dir,
);
let eta1 = mk(
exit_base.eta,
&exit_base.eta_u,
&exit_base.eta_uv,
&exit_ext.eta_u_dir,
&exit_ext.eta_uv_dir,
);
let chi1 = mk(
exit_base.chi,
&exit_base.chi_u,
&exit_base.chi_uv,
&exit_ext.chi_u_dir,
&exit_ext.chi_uv_dir,
);
let d1 = mk(
exit_base.d,
&exit_base.d_u,
&exit_base.d_uv,
&exit_ext.d_u_dir,
&exit_ext.d_uv_dir,
);
let q1j = Jet3::primary(q1, primary.q1, p, dir[primary.q1]);
let qd1j = Jet3::primary(qd1, primary.qd1, p, dir[primary.qd1]);
let out = flex_row_nll(&eta0, &eta1, &chi1, &d1, &q1j, &qd1j, surv0, surv1, wi, di);
Array2::from_shape_vec((p, p), out.contracted_third()).map_err(|e| e.to_string())
}
pub(crate) fn flex_row_nll_fourth_contracted(
&self,
row: usize,
primary: &FlexPrimarySlices,
q1: f64,
qd1: f64,
dir_u: &[f64],
dir_v: &[f64],
packs: FlexFourthPacks<'_>,
) -> Result<Array2<f64>, String> {
let FlexFourthPacks {
entry_base,
exit_base,
entry_ext_u,
exit_ext_u,
entry_ext_v,
exit_ext_v,
entry_bi,
exit_bi,
} = packs;
let p = primary.total;
let wi = self.weights[row];
let di = self.event[row];
let surv0 = surv_stack(entry_base.eta)?;
let surv1 = surv_stack(exit_base.eta)?;
let mk = |base_v: f64,
base_g: &[f64],
base_h: &[f64],
ext_u_g: &[f64],
ext_u_h: &[f64],
ext_v_g: &[f64],
ext_v_h: &[f64],
bi_h: &[f64]|
-> Jet4 {
let eps_del_v = quad_form(base_h, dir_u, dir_v, p);
let eps_del_g = mat_vec(ext_u_h, dir_v, p);
Jet4 {
base: Jet2::from_parts(base_v, base_g, base_h),
eps: Jet2::from_parts(dot(base_g, dir_u), ext_u_g, ext_u_h),
del: Jet2::from_parts(dot(base_g, dir_v), ext_v_g, ext_v_h),
eps_del: Jet2::from_parts(eps_del_v, &eps_del_g, bi_h),
}
};
let eta0 = mk(
entry_base.eta,
&entry_base.eta_u,
&entry_base.eta_uv,
&entry_ext_u.eta_u_dir,
&entry_ext_u.eta_uv_dir,
&entry_ext_v.eta_u_dir,
&entry_ext_v.eta_uv_dir,
&entry_bi.eta_uv_uv,
);
let eta1 = mk(
exit_base.eta,
&exit_base.eta_u,
&exit_base.eta_uv,
&exit_ext_u.eta_u_dir,
&exit_ext_u.eta_uv_dir,
&exit_ext_v.eta_u_dir,
&exit_ext_v.eta_uv_dir,
&exit_bi.eta_uv_uv,
);
let chi1 = mk(
exit_base.chi,
&exit_base.chi_u,
&exit_base.chi_uv,
&exit_ext_u.chi_u_dir,
&exit_ext_u.chi_uv_dir,
&exit_ext_v.chi_u_dir,
&exit_ext_v.chi_uv_dir,
&exit_bi.chi_uv_uv,
);
let d1 = mk(
exit_base.d,
&exit_base.d_u,
&exit_base.d_uv,
&exit_ext_u.d_u_dir,
&exit_ext_u.d_uv_dir,
&exit_ext_v.d_u_dir,
&exit_ext_v.d_uv_dir,
&exit_bi.d_uv_uv,
);
let q1j = Jet4::primary(q1, primary.q1, p, dir_u[primary.q1], dir_v[primary.q1]);
let qd1j = Jet4::primary(qd1, primary.qd1, p, dir_u[primary.qd1], dir_v[primary.qd1]);
let out = flex_row_nll(&eta0, &eta1, &chi1, &d1, &q1j, &qd1j, surv0, surv1, wi, di);
Array2::from_shape_vec((p, p), out.contracted_fourth()).map_err(|e| e.to_string())
}
}
fn recip<J: FlexJet>(x: &J) -> J {
let v = x.value();
let inv = 1.0 / v;
let inv2 = inv * inv;
x.compose_unary([
inv,
-inv2,
2.0 * inv2 * inv,
-6.0 * inv2 * inv2,
24.0 * inv2 * inv2 * inv,
])
}
fn exp_jet<J: FlexJet>(x: &J) -> J {
let e = x.value().exp();
x.compose_unary([e, e, e, e, e])
}
fn add_const<J: FlexJet>(x: &J, c: f64) -> J {
x.compose_unary([x.value() + c, 1.0, 0.0, 0.0, 0.0])
}
trait MomentTerm: FlexJet {
fn moment_term(&self, moment: &Self) -> Self {
const EULER: [f64; 5] = [0.0, 1.0, 2.0, 3.0, 4.0];
const EULER_INVERSE: [f64; 5] = [0.0, 1.0, 0.5, 1.0 / 3.0, 0.25];
self.scale_homogeneous_orders(EULER)
.mul(moment)
.scale_homogeneous_orders(EULER_INVERSE)
}
}
impl<J: FlexJet> MomentTerm for J {}
fn base_moment_jets<J: FlexJet>(
c: &[J; 4],
z_left: &J,
left_finite: bool,
z_right: &J,
right_finite: bool,
numeric_moments: &[f64],
) -> [J; 5] {
assert!(
(1..=4).contains(&J::ORDER),
"base_moment_jets supports exact derivative orders 1 through 4"
);
let required_moments = 5 + 6 * J::ORDER;
assert!(
numeric_moments.len() >= required_moments,
"order-{} base-moment jet requires numeric M_0..M_{}, got only {} moments",
J::ORDER,
required_moments - 1,
numeric_moments.len()
);
let c0_const: [J; 4] = std::array::from_fn(|k| const_jet_like(&c[k], c[k].value()));
let conv = |lhs: &[J], rhs: &[J]| -> Vec<J> {
let mut out: Vec<J> = (0..lhs.len() + rhs.len() - 1)
.map(|_| const_jet_like(&c[0], 0.0))
.collect();
for (i, li) in lhs.iter().enumerate() {
for (j, rj) in rhs.iter().enumerate() {
out[i + j] = out[i + j].add(&li.mul(rj));
}
}
out
};
let eta_sq = conv(c, c);
let eta0_sq = conv(&c0_const, &c0_const);
let neg_dq: Vec<J> = eta_sq
.iter()
.zip(eta0_sq.iter())
.map(|(a, b)| a.sub(b).scale(-0.5))
.collect();
let mut s_poly: Vec<J> = vec![const_jet_like(&c[0], 1.0)];
let mut power: Vec<J> = s_poly.clone();
let factorials = [1.0_f64, 1.0, 2.0, 6.0, 24.0];
for fact in factorials.iter().take(J::ORDER + 1).skip(1) {
power = conv(&power, &neg_dq);
for (m, coeff) in power.iter().enumerate() {
let term = coeff.scale(1.0 / fact);
if m < s_poly.len() {
s_poly[m] = s_poly[m].add(&term);
} else {
s_poly.push(term);
}
}
}
std::array::from_fn(|n| {
let mut acc = const_jet_like(&c[0], 0.0);
for (m, s_m) in s_poly.iter().enumerate() {
let m_npm = numeric_moments[n + m];
if m_npm != 0.0 {
acc = acc.add(&s_m.scale(m_npm));
}
}
if let Some(sr) = edge_sliver_jet(n, c, z_right, right_finite) {
acc = acc.add(&sr);
}
if let Some(sl) = edge_sliver_jet(n, c, z_left, left_finite) {
acc = acc.sub(&sl);
}
acc
})
}
fn edge_sliver_jet<J: FlexJet>(n: usize, c: &[J; 4], z_e: &J, finite: bool) -> Option<J> {
if !finite {
return None;
}
let z0 = z_e.value();
let zc = const_jet_like(z_e, z0); let eta = c[3]
.mul(&zc)
.add(&c[2])
.mul(&zc)
.add(&c[1])
.mul(&zc)
.add(&c[0]);
let z_pow = {
let mut zk = const_jet_like(z_e, 1.0);
for _ in 0..n {
zk = zk.mul(&zc);
}
zk
};
let q = zc.mul(&zc).add(&eta.mul(&eta)).scale(0.5);
let w = exp_jet(&q.scale(-1.0));
let g = z_pow.mul(&w);
let delta = tangent_jet(z_e);
let mut sliver = g.mul(&delta);
if J::ORDER == 1 {
return Some(sliver);
}
let eta_z = c[2]
.scale(2.0)
.add(&c[3].scale(3.0).mul(&zc))
.mul(&zc)
.add(&c[1]); let q_z = zc.add(&eta.mul(&eta_z));
let nz = |power: i32| -> J {
if n == 0 || z0 == 0.0 {
const_jet_like(z_e, 0.0)
} else {
const_jet_like(z_e, n as f64 / z0.powi(power))
}
};
let a1 = nz(1).sub(&q_z);
let g_z = a1.mul(&g);
let d2 = delta.mul(&delta);
sliver = sliver.add(&g_z.mul(&d2).scale(0.5));
if J::ORDER == 2 {
return Some(sliver);
}
let eta_zz = c[2].scale(2.0).add(&c[3].scale(6.0).mul(&zc)); let q_zz = add_const(&eta_z.mul(&eta_z).add(&eta.mul(&eta_zz)), 1.0);
let a1p = nz(2).scale(-1.0).sub(&q_zz);
let b2 = a1p.add(&a1.mul(&a1));
let g_zz = b2.mul(&g);
let d3 = d2.mul(&delta);
sliver = sliver.add(&g_zz.mul(&d3).scale(1.0 / 6.0));
if J::ORDER == 3 {
return Some(sliver);
}
assert_eq!(
J::ORDER,
4,
"edge sliver supports derivative orders 1 through 4"
);
let eta_zzz = c[3].scale(6.0); let q_zzz = eta_z.scale(3.0).mul(&eta_zz).add(&eta.mul(&eta_zzz));
let a1pp = nz(3).scale(2.0).sub(&q_zzz);
let b2p = a1pp.add(&a1.mul(&a1p).scale(2.0));
let g_zzz = b2p.add(&a1.mul(&b2)).mul(&g);
let d4 = d3.mul(&delta);
Some(sliver.add(&g_zzz.mul(&d4).scale(1.0 / 24.0)))
}
fn flex_timepoint_inputs_generic<J: FlexJet + MomentTerm>(
template: &J,
b_jet: &J,
du: &[J],
a0: f64,
d_check: f64,
primary_g: usize,
infl: Option<usize>,
q_jet: &J,
scale_ratio_jet: &J,
z_obs: f64,
o_infl: f64,
obs_coeff: [f64; 4],
obs_fixed: &DenestedCellPrimaryFixedPartials,
cells: &[CalibrationCellJetInputs<'_>],
) -> Result<(J, J, J), String> {
let residual =
|a: &J| calibration_residual_jet(a, b_jet, primary_g, du, q_jet, scale_ratio_jet, cells);
let a_jet = lift_intercept_flex(template, a0, 1.0 / d_check, J::ORDER, residual);
let da = tangent_jet(&a_jet);
let eta_coeff_base = cell_coeff_jets(&a_jet, obs_coeff, obs_fixed, primary_g, &da, du);
let chi_coeff_base = cell_chi_poly_jets(&a_jet, obs_fixed, primary_g, &da, du);
let eta_coeff =
std::array::from_fn(|coefficient| eta_coeff_base[coefficient].mul(scale_ratio_jet));
let chi_coeff =
std::array::from_fn(|coefficient| chi_coeff_base[coefficient].mul(scale_ratio_jet));
let mut eta = add_const(&eval_coeff_jet_at(&eta_coeff, z_obs), o_infl);
if let Some(infl_axis) = infl {
eta = eta.add(&du[infl_axis]);
}
let chi = eval_coeff_jet_at(&chi_coeff, z_obs);
let mut d = const_jet_like(template, 0.0);
for cell in cells {
let c_pos_base =
cell_coeff_jets(&a_jet, cell.base_pos_coeffs, cell.fixed, primary_g, &da, du);
let chi_jets_base = cell_chi_poly_jets(&a_jet, cell.fixed, primary_g, &da, du);
let c_pos = std::array::from_fn(|coefficient| c_pos_base[coefficient].mul(scale_ratio_jet));
let chi_jets =
std::array::from_fn(|coefficient| chi_jets_base[coefficient].mul(scale_ratio_jet));
let edge_l = cell_edge_jet(&a_jet, b_jet, cell.left_edge, cell.cell_left);
let edge_r = cell_edge_jet(&a_jet, b_jet, cell.right_edge, cell.cell_right);
d = d.add(&flex_timepoint_d_cell(
template,
&c_pos,
&chi_jets,
&edge_l,
cell.cell_left.is_finite(),
&edge_r,
cell.cell_right.is_finite(),
cell.numeric_moments,
));
}
Ok((eta, chi, d))
}
#[inline]
fn tangent_jet<J: FlexJet>(x: &J) -> J {
add_const(x, -x.value())
}
#[inline]
fn const_jet_like<J: FlexJet>(template: &J, v: f64) -> J {
add_const(&template.scale(0.0), v)
}
fn lift_intercept_flex<J: FlexJet>(
template: &J,
a0: f64,
inv_fa: f64,
iters: usize,
residual: impl Fn(&J) -> J,
) -> J {
let mut a = const_jet_like(template, a0);
for _ in 0..iters {
let r = residual(&a);
a = a.sub(&r.scale(inv_fa));
}
a
}
fn calibration_residual_jet<J: FlexJet + MomentTerm>(
a_jet: &J,
b_jet: &J,
g_axis: usize,
du: &[J],
q_jet: &J,
scale_ratio_jet: &J,
cells: &[CalibrationCellJetInputs<'_>],
) -> J {
let da = tangent_jet(a_jet);
let inv_two_pi = std::f64::consts::TAU.recip();
let mut r = const_jet_like(a_jet, 0.0);
for cell in cells {
let c_pos_base = cell_coeff_jets(a_jet, cell.base_pos_coeffs, cell.fixed, g_axis, &da, du);
let c_pos = std::array::from_fn(|coefficient| c_pos_base[coefficient].mul(scale_ratio_jet));
let edge_l = cell_edge_jet(a_jet, b_jet, cell.left_edge, cell.cell_left);
let edge_r = cell_edge_jet(a_jet, b_jet, cell.right_edge, cell.cell_right);
let m = base_moment_jets(
&c_pos,
&edge_l,
cell.cell_left.is_finite(),
&edge_r,
cell.cell_right.is_finite(),
cell.numeric_moments,
);
let mut cell_r = const_jet_like(a_jet, 0.0);
for k in 0..4 {
cell_r = cell_r.add(&c_pos[k].moment_term(&m[k]));
}
r = r.add(&cell_r.scale(inv_two_pi));
}
let q = q_jet.value();
let phi_q = crate::probability::normal_pdf(q);
let g0 = crate::probability::normal_cdf(-q);
let g1 = -phi_q;
let g2 = q * phi_q;
let g3 = (1.0 - q * q) * phi_q;
let g4 = (q * q * q - 3.0 * q) * phi_q;
let q_self = add_const(&q_jet.compose_unary([g0, g1, g2, g3, g4]), -g0);
r = r.add(&q_self);
r
}
struct CalibrationCellJetInputs<'a> {
base_pos_coeffs: [f64; 4],
fixed: &'a DenestedCellPrimaryFixedPartials,
cell_left: f64,
cell_right: f64,
left_edge: crate::cubic_cell_kernel::PartitionEdge,
right_edge: crate::cubic_cell_kernel::PartitionEdge,
numeric_moments: &'a [f64],
}
fn cell_edge_jet<J: FlexJet>(
a_jet: &J,
b_jet: &J,
edge: crate::cubic_cell_kernel::PartitionEdge,
z_value: f64,
) -> J {
match edge {
crate::cubic_cell_kernel::PartitionEdge::Crossing { tau } => {
const_jet_like(a_jet, tau).sub(a_jet).mul(&recip(b_jet))
}
crate::cubic_cell_kernel::PartitionEdge::Fixed(_) => const_jet_like(a_jet, z_value),
}
}
fn cell_coeff_jets<J: FlexJet>(
template: &J,
base_c: [f64; 4],
fixed: &DenestedCellPrimaryFixedPartials,
g_axis: usize,
da: &J,
du: &[J],
) -> [J; 4] {
let p = du.len();
let dada = da.mul(da);
let dadada = dada.mul(da);
let db = &du[g_axis];
let dadb = da.mul(db);
let dbdb = db.mul(db);
std::array::from_fn(|k| {
let mut c = const_jet_like(template, base_c[k]);
c = c
.add(&da.scale(fixed.dc_da[k]))
.add(&dada.scale(0.5 * fixed.dc_daa[k]))
.add(&dadada.scale(fixed.dc_daaa[k] / 6.0));
for u in 0..p {
if u == g_axis {
continue;
}
let duu = &du[u];
let mut chain = duu.scale(fixed.coeff_u[u][k]);
chain = chain
.add(&da.mul(duu).scale(fixed.coeff_au[u][k]))
.add(&dada.mul(duu).scale(0.5 * fixed.coeff_aau[u][k]));
chain = chain
.add(&db.mul(duu).scale(fixed.coeff_bu[u][k]))
.add(&dadb.mul(duu).scale(fixed.coeff_abu[u][k]))
.add(&dbdb.mul(duu).scale(0.5 * fixed.coeff_bbu[u][k]));
chain = chain
.add(&dadada.mul(duu).scale(fixed.coeff_aaau[u][k] / 6.0))
.add(&dada.mul(db).mul(duu).scale(0.5 * fixed.coeff_aabu[u][k]))
.add(&dadb.mul(db).mul(duu).scale(0.5 * fixed.coeff_abbu[u][k]))
.add(&dbdb.mul(db).mul(duu).scale(fixed.coeff_bbbu[u][k] / 6.0));
c = c.add(&chain);
}
c = c
.add(&db.scale(fixed.coeff_u[g_axis][k]))
.add(&dadb.scale(fixed.coeff_au[g_axis][k]))
.add(&dada.mul(db).scale(0.5 * fixed.coeff_aau[g_axis][k]))
.add(&dbdb.scale(0.5 * fixed.coeff_bu[g_axis][k]))
.add(&dadb.mul(db).scale(0.5 * fixed.coeff_abu[g_axis][k]))
.add(&dbdb.mul(db).scale(fixed.coeff_bbu[g_axis][k] / 6.0));
c
})
}
fn cell_chi_poly_jets<J: FlexJet>(
template: &J,
fixed: &DenestedCellPrimaryFixedPartials,
g_axis: usize,
da: &J,
du: &[J],
) -> [J; 4] {
let p = du.len();
let dada = da.mul(da);
let db = &du[g_axis];
std::array::from_fn(|k| {
let mut c = const_jet_like(template, fixed.dc_da[k]);
c = c
.add(&da.scale(fixed.dc_daa[k]))
.add(&dada.scale(0.5 * fixed.dc_daaa[k]));
let dbdb = db.mul(db);
let dadb = da.mul(db);
for u in 0..p {
if u == g_axis {
continue;
}
let duu = &du[u];
let chain = duu
.scale(fixed.coeff_au[u][k])
.add(&da.mul(duu).scale(fixed.coeff_aau[u][k]))
.add(&db.mul(duu).scale(fixed.coeff_abu[u][k]))
.add(&dada.mul(duu).scale(0.5 * fixed.coeff_aaau[u][k]))
.add(&dadb.mul(duu).scale(fixed.coeff_aabu[u][k]))
.add(&dbdb.mul(duu).scale(0.5 * fixed.coeff_abbu[u][k]));
c = c.add(&chain);
}
c = c
.add(&db.scale(fixed.coeff_au[g_axis][k]))
.add(&da.mul(db).scale(fixed.coeff_aau[g_axis][k]))
.add(&dbdb.scale(0.5 * fixed.coeff_abu[g_axis][k]));
c
})
}
fn flex_timepoint_d_cell<J: FlexJet>(
template: &J,
c_jets: &[J; 4],
chi_jets: &[J; 4],
edge_l: &J,
left_finite: bool,
edge_r: &J,
right_finite: bool,
numeric_moments: &[f64],
) -> J {
let m = base_moment_jets(
c_jets,
edge_l,
left_finite,
edge_r,
right_finite,
numeric_moments,
);
let mut acc = const_jet_like(template, 0.0);
for (k, chi_k) in chi_jets.iter().enumerate() {
acc = acc.add(&chi_k.mul(&m[k]));
}
acc.scale(std::f64::consts::TAU.recip())
}
#[inline]
fn eval_coeff_jet_at<J: FlexJet>(coeff_jet: &[J; 4], z: f64) -> J {
let mut zk = 1.0;
let mut acc = const_jet_like(&coeff_jet[0], 0.0);
for c in coeff_jet.iter() {
acc = acc.add(&c.scale(zk));
zk *= z;
}
acc
}
fn cells_from_cached(cached: &CachedPartitionCells) -> Vec<CalibrationCellJetInputs<'_>> {
cached
.cells
.iter()
.map(|entry| {
let cell = entry.partition_cell.cell;
CalibrationCellJetInputs {
base_pos_coeffs: [cell.c0, cell.c1, cell.c2, cell.c3],
fixed: &entry.fixed,
cell_left: cell.left,
cell_right: cell.right,
left_edge: entry.partition_cell.left_edge,
right_edge: entry.partition_cell.right_edge,
numeric_moments: entry.state.moments.as_slice(),
}
})
.collect()
}
fn observed_fixed_for(
family: &SurvivalMarginalSlopeFamily,
primary: &FlexPrimarySlices,
row: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
) -> Result<([f64; 4], DenestedCellPrimaryFixedPartials), String> {
let r = primary.total;
let scale = family.probit_frailty_scale();
let z_obs = family.observed_score_projection(row);
let u_obs = a + b * z_obs;
let obs = family.observed_denested_cell_partials(row, a, b, beta_h, beta_w)?;
let mut coeff_u = vec![[0.0; 4]; r];
let mut coeff_au = vec![[0.0; 4]; r];
let mut coeff_bu = vec![[0.0; 4]; r];
let mut coeff_aau = vec![[0.0; 4]; r];
let mut coeff_abu = vec![[0.0; 4]; r];
let mut coeff_bbu = vec![[0.0; 4]; r];
let mut coeff_aaau = vec![[0.0; 4]; r];
let mut coeff_aabu = vec![[0.0; 4]; r];
let mut coeff_abbu = vec![[0.0; 4]; r];
let mut coeff_bbbu = vec![[0.0; 4]; r];
coeff_u[primary.g] = obs.dc_db;
coeff_au[primary.g] = obs.dc_dab;
coeff_bu[primary.g] = obs.dc_dbb;
coeff_aau[primary.g] = obs.dc_daab;
coeff_abu[primary.g] = obs.dc_dabb;
coeff_bbu[primary.g] = obs.dc_dbbb;
if let Some(h_range) = primary.h.as_ref().filter(|_| family.score_warp.is_some()) {
for local_idx in 0..h_range.len() {
let idx = h_range.start + local_idx;
coeff_u[idx] = scale_coeff4(
family.observed_score_basis_coefficients(row, local_idx, z_obs, b)?,
scale,
);
coeff_bu[idx] = scale_coeff4(
family.observed_score_basis_coefficients(row, local_idx, z_obs, 1.0)?,
scale,
);
}
}
if let (Some(w_range), Some(runtime)) = (primary.w.as_ref(), family.link_dev.as_ref()) {
for local_idx in 0..w_range.len() {
let span = runtime.basis_cubic_at(local_idx, u_obs)?;
let idx = w_range.start + local_idx;
coeff_u[idx] = scale_coeff4(
exact_kernel::link_basis_cell_coefficients(span, a, b),
scale,
);
let (dc_aw, dc_bw) = exact_kernel::link_basis_cell_coefficient_partials(span, a, b);
let (dc_aaw, dc_abw, dc_bbw) =
exact_kernel::link_basis_cell_second_partials(span, a, b);
let (dc_aaaw, dc_aabw, dc_abbw, dc_bbbw) =
exact_kernel::link_basis_cell_third_partials(span);
coeff_au[idx] = scale_coeff4(dc_aw, scale);
coeff_bu[idx] = scale_coeff4(dc_bw, scale);
coeff_aau[idx] = scale_coeff4(dc_aaw, scale);
coeff_abu[idx] = scale_coeff4(dc_abw, scale);
coeff_bbu[idx] = scale_coeff4(dc_bbw, scale);
coeff_aaau[idx] = scale_coeff4(dc_aaaw, scale);
coeff_aabu[idx] = scale_coeff4(dc_aabw, scale);
coeff_abbu[idx] = scale_coeff4(dc_abbw, scale);
coeff_bbbu[idx] = scale_coeff4(dc_bbbw, scale);
}
}
let fixed = DenestedCellPrimaryFixedPartials {
dc_da: obs.dc_da,
dc_daa: obs.dc_daa,
dc_daaa: obs.dc_daaa,
coeff_u,
coeff_au,
coeff_bu,
coeff_aau,
coeff_abu,
coeff_bbu,
coeff_aaau,
coeff_aabu,
coeff_abbu,
coeff_bbbu,
};
Ok((obs.coeff, fixed))
}
impl SurvivalMarginalSlopeFamily {
pub(crate) fn compute_survival_timepoint_exact_jet(
&self,
row: usize,
primary: &FlexPrimarySlices,
q: f64,
q_index: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
o_infl: f64,
) -> Result<SurvivalFlexTimepointExact, String> {
let cached = self.build_cached_partition(primary, a, b, beta_h, beta_w)?;
self.compute_survival_timepoint_exact_jet_from_cached(
row, primary, q, q_index, a, b, beta_h, beta_w, o_infl, &cached,
)
}
pub(crate) fn compute_survival_timepoint_exact_jet_from_cached(
&self,
row: usize,
primary: &FlexPrimarySlices,
q: f64,
q_index: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
o_infl: f64,
cached: &CachedPartitionCells,
) -> Result<SurvivalFlexTimepointExact, String> {
let p = primary.total;
let d_check = self.evaluate_survival_denom_d(a, b, beta_h, beta_w)?;
let z_obs = self.observed_score_projection(row);
let (obs_coeff, obs_fixed) = observed_fixed_for(self, primary, row, a, b, beta_h, beta_w)?;
let cells = cells_from_cached(cached);
let template = Jet2::primary(0.0, usize::MAX, p);
let b_jet = Jet2::primary(b, primary.g, p);
let du: Vec<Jet2> = (0..p).map(|u| Jet2::primary(0.0, u, p)).collect();
let q_jet = add_const(&du[q_index], q);
let (eta, chi, d) = flex_timepoint_inputs_generic(
&template,
&b_jet,
&du,
a,
d_check,
primary.g,
primary.infl,
&q_jet,
&const_jet_like(&template, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)?;
let to_g = |j: &Jet2| Array1::from(j.g.clone());
let to_h = |j: &Jet2| -> Result<Array2<f64>, String> {
Array2::from_shape_vec((p, p), j.h.clone()).map_err(|e| e.to_string())
};
Ok(SurvivalFlexTimepointExact {
eta: eta.value(),
chi: chi.value(),
d: d.value(),
eta_u: to_g(&eta),
eta_uv: to_h(&eta)?,
chi_u: to_g(&chi),
chi_uv: to_h(&chi)?,
d_u: to_g(&d),
d_uv: to_h(&d)?,
})
}
pub(crate) fn compute_survival_timepoint_first_order_exact(
&self,
row: usize,
primary: &FlexPrimarySlices,
q: f64,
q_index: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
o_infl: f64,
) -> Result<SurvivalFlexTimepointFirstOrderExact, String> {
let cached = self.build_cached_partition(primary, a, b, beta_h, beta_w)?;
let p = primary.total;
let d_check = self.evaluate_survival_denom_d(a, b, beta_h, beta_w)?;
let z_obs = self.observed_score_projection(row);
let (obs_coeff, obs_fixed) = observed_fixed_for(self, primary, row, a, b, beta_h, beta_w)?;
let cells = cells_from_cached(&cached);
let template = Jet1::primary(0.0, usize::MAX, p);
let b_jet = Jet1::primary(b, primary.g, p);
let du: Vec<Jet1> = (0..p).map(|u| Jet1::primary(0.0, u, p)).collect();
let q_jet = add_const(&du[q_index], q);
let (eta, chi, d) = flex_timepoint_inputs_generic(
&template,
&b_jet,
&du,
a,
d_check,
primary.g,
primary.infl,
&q_jet,
&const_jet_like(&template, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)?;
let to_g = |j: &Jet1| Array1::from(j.g.clone());
Ok(SurvivalFlexTimepointFirstOrderExact {
eta: eta.value(),
chi: chi.value(),
d: d.value(),
eta_u: to_g(&eta),
chi_u: to_g(&chi),
d_u: to_g(&d),
})
}
}
impl SurvivalMarginalSlopeFamily {
pub(crate) fn compute_survival_timepoint_directional_jet_from_cached(
&self,
row: usize,
primary: &FlexPrimarySlices,
q: f64,
q_index: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
o_infl: f64,
cached: &CachedPartitionCells,
dir: &Array1<f64>,
arena: &DynamicJetArena,
) -> Result<(FlexTimepointBasePack, FlexTimepointDirectionalPack), String> {
let p = primary.total;
let d_check = self.evaluate_survival_denom_d(a, b, beta_h, beta_w)?;
let z_obs = self.observed_score_projection(row);
let (obs_coeff, obs_fixed) = observed_fixed_for(self, primary, row, a, b, beta_h, beta_w)?;
let cells = cells_from_cached(cached);
macro_rules! evaluate {
($template:expr, $b_jet:expr, $du:expr) => {{
let template = $template;
let b_jet = $b_jet;
let du = $du;
let q_jet = add_const(&du[q_index], q);
let (eta, chi, d) = flex_timepoint_inputs_generic(
&template,
&b_jet,
&du,
a,
d_check,
primary.g,
primary.infl,
&q_jet,
&const_jet_like(&template, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)?;
Ok(FlexThirdOutput::pack_timepoint_outputs(&eta, &chi, &d))
}};
}
match p {
8 => evaluate!(
FixedJet3::<8>::primary(0.0, usize::MAX, p, 0.0),
FixedJet3::<8>::primary(b, primary.g, p, dir[primary.g]),
(0..p)
.map(|axis| FixedJet3::<8>::primary(0.0, axis, p, dir[axis]))
.collect::<Vec<_>>()
),
_ => evaluate!(
ArenaJet3::primary(0.0, usize::MAX, p, 0.0, arena),
ArenaJet3::primary(b, primary.g, p, dir[primary.g], arena),
(0..p)
.map(|axis| ArenaJet3::primary(0.0, axis, p, dir[axis], arena))
.collect::<Vec<_>>()
),
}
}
pub(crate) fn compute_survival_timepoint_bidirectional_jet_from_cached(
&self,
row: usize,
primary: &FlexPrimarySlices,
q: f64,
q_index: usize,
a: f64,
b: f64,
beta_h: Option<&Array1<f64>>,
beta_w: Option<&Array1<f64>>,
cached: &CachedPartitionCells,
dir1: &Array1<f64>,
dir2: &Array1<f64>,
) -> Result<FlexTimepointBidirectionalPack, String> {
let p = primary.total;
let d_check = self.evaluate_survival_denom_d(a, b, beta_h, beta_w)?;
let z_obs = self.observed_score_projection(row);
let (obs_coeff, obs_fixed) = observed_fixed_for(self, primary, row, a, b, beta_h, beta_w)?;
let cells = cells_from_cached(cached);
let template = Jet4::primary(0.0, usize::MAX, p, 0.0, 0.0);
let b_jet = Jet4::primary(b, primary.g, p, dir1[primary.g], dir2[primary.g]);
let du: Vec<Jet4> = (0..p)
.map(|u| Jet4::primary(0.0, u, p, dir1[u], dir2[u]))
.collect();
let q_jet = add_const(&du[q_index], q);
let (eta, chi, d) = flex_timepoint_inputs_generic(
&template,
&b_jet,
&du,
a,
d_check,
primary.g,
primary.infl,
&q_jet,
&const_jet_like(&template, 1.0),
z_obs,
0.0,
obs_coeff,
&obs_fixed,
&cells,
)?;
Ok(FlexTimepointBidirectionalPack {
eta_uv_uv: eta.eps_del.h.clone(),
chi_uv_uv: chi.eps_del.h.clone(),
d_uv_uv: d.eps_del.h.clone(),
})
}
}
struct FlexFamilyCoefficientJets<J: FlexCoefficientJet> {
template: Dual2<J>,
q0: Dual2<J>,
q1: Dual2<J>,
qd1: Dual2<J>,
g: Dual2<J>,
scale_ratio: Dual2<J>,
du: Vec<Dual2<J>>,
o_infl: f64,
}
enum FlexCoefficientRowDirection<'a> {
Beta(&'a Array1<f64>),
Design {
block: usize,
derivative_row: &'a Array1<f64>,
beta: &'a Array1<f64>,
coefficient_range: std::ops::Range<usize>,
},
}
fn add_coefficient_row(
gradient: &mut [f64],
range: &std::ops::Range<usize>,
row: ndarray::ArrayView1<'_, f64>,
channel: &str,
) -> Result<(), String> {
if range.len() != row.len() {
return Err(format!(
"survival marginal-slope FLEX family {channel} coefficient row width {} != flattened range width {}",
row.len(),
range.len(),
));
}
for (axis, value) in range.clone().zip(row.iter().copied()) {
gradient[axis] += value;
}
Ok(())
}
fn lifted_coefficient_affine<J: FlexCoefficientJet>(
value: f64,
gradient: Vec<f64>,
coefficient_direction: Option<&FlexCoefficientRowDirection<'_>>,
design_affected: bool,
family_first: f64,
family_second: f64,
) -> Dual2<J> {
let dimension = gradient.len();
let mut directional_gradient = vec![0.0; dimension];
let directional_value = match coefficient_direction {
None => 0.0,
Some(FlexCoefficientRowDirection::Beta(direction)) => gradient
.iter()
.zip(direction.iter())
.map(|(&coefficient, &step)| coefficient * step)
.sum(),
Some(FlexCoefficientRowDirection::Design {
derivative_row,
beta,
coefficient_range,
..
}) if design_affected => {
for (axis, value) in coefficient_range
.clone()
.zip(derivative_row.iter().copied())
{
directional_gradient[axis] = value;
}
derivative_row.dot(*beta)
}
Some(FlexCoefficientRowDirection::Design { .. }) => 0.0,
};
Dual2 {
v: J::affine(value, gradient, directional_value, directional_gradient),
g: J::constant(family_first, dimension),
h: J::constant(family_second, dimension),
}
}
fn one_hot_coefficient_gradient(dimension: usize, axis: usize) -> Vec<f64> {
let mut gradient = vec![0.0; dimension];
gradient[axis] = 1.0;
gradient
}
impl SurvivalMarginalSlopeFamily {
fn validate_flex_family_coefficient_state(
&self,
row: usize,
block_states: &[ParameterBlockState],
slices: &BlockSlices,
first: FlexFamilyRowDirection,
second: FlexFamilyRowDirection,
coefficient_direction: Option<&FlexCoefficientRowDirection<'_>>,
) -> Result<(), String> {
if row >= self.n {
return Err(format!(
"survival marginal-slope FLEX family row {row} is out of range for n={}",
self.n
));
}
let directions = [
first.entry,
first.exit,
first.derivative_exit,
first.probit_scale,
second.entry,
second.exit,
second.derivative_exit,
second.probit_scale,
];
if directions.iter().any(|value| !value.is_finite()) {
return Err(
"survival marginal-slope FLEX family row directions must be finite".to_string(),
);
}
match coefficient_direction {
Some(FlexCoefficientRowDirection::Beta(direction)) => {
if direction.len() != slices.total {
return Err(format!(
"survival marginal-slope FLEX family beta direction length {} != flattened coefficient width {}",
direction.len(),
slices.total,
));
}
if direction.iter().any(|value| !value.is_finite()) {
return Err(
"survival marginal-slope FLEX family beta direction must be finite"
.to_string(),
);
}
}
Some(FlexCoefficientRowDirection::Design {
block,
derivative_row,
beta,
coefficient_range,
}) => {
if !matches!(*block, 1 | 2) {
return Err(format!(
"survival marginal-slope FLEX family design direction supports marginal/logslope blocks 1 or 2, got block {block}"
));
}
let expected_range = if *block == 1 {
&slices.marginal
} else {
&slices.logslope
};
if coefficient_range != expected_range {
return Err(format!(
"survival marginal-slope FLEX family design direction block {block} range {:?} != canonical {:?}",
coefficient_range, expected_range,
));
}
if derivative_row.len() != coefficient_range.len()
|| beta.len() != coefficient_range.len()
{
return Err(format!(
"survival marginal-slope FLEX family design direction block {block} widths disagree: derivative={}, beta={}, range={}",
derivative_row.len(),
beta.len(),
coefficient_range.len(),
));
}
if coefficient_range.end > slices.total {
return Err(format!(
"survival marginal-slope FLEX family design direction range {:?} exceeds flattened width {}",
coefficient_range, slices.total,
));
}
if derivative_row.iter().any(|value| !value.is_finite()) {
return Err(
"survival marginal-slope FLEX family design derivative row must be finite"
.to_string(),
);
}
}
None => {}
}
let mut expected = vec![
("time", slices.time.clone()),
("marginal", slices.marginal.clone()),
("logslope", slices.logslope.clone()),
];
if let Some(range) = slices.score_warp.as_ref() {
expected.push(("score warp", range.clone()));
}
if let Some(range) = slices.link_dev.as_ref() {
expected.push(("link deviation", range.clone()));
}
if let Some(range) = slices.influence.as_ref() {
expected.push(("influence", range.clone()));
}
if block_states.len() != expected.len() {
return Err(format!(
"survival marginal-slope FLEX family coefficient map expects {} block states, got {}",
expected.len(),
block_states.len(),
));
}
for (block_index, (block_state, (name, range))) in
block_states.iter().zip(expected.iter()).enumerate()
{
if block_state.beta.len() != range.len() {
return Err(format!(
"survival marginal-slope FLEX family {name} block {block_index} beta width {} != layout width {}",
block_state.beta.len(),
range.len(),
));
}
}
for (block, name) in [(0usize, "time"), (1, "marginal"), (2, "logslope")] {
if block_states[block].eta.len() <= row {
return Err(format!(
"survival marginal-slope FLEX family {name} block eta length {} does not contain row {row}",
block_states[block].eta.len(),
));
}
}
Ok(())
}
fn build_flex_family_coefficient_jets<J: FlexCoefficientJet>(
&self,
row: usize,
block_states: &[ParameterBlockState],
primary: &FlexPrimarySlices,
slices: &BlockSlices,
first: FlexFamilyRowDirection,
second: FlexFamilyRowDirection,
coefficient_direction: Option<&FlexCoefficientRowDirection<'_>>,
) -> Result<FlexFamilyCoefficientJets<J>, String> {
let dimension = slices.total;
let design_targets_marginal = matches!(
coefficient_direction,
Some(FlexCoefficientRowDirection::Design { block: 1, .. })
);
let design_targets_logslope = matches!(
coefficient_direction,
Some(FlexCoefficientRowDirection::Design { block: 2, .. })
);
let q_values = self.row_dynamic_q_values(row, block_states)?;
let time_entry_chunk = self
.design_entry
.try_row_chunk(row..row + 1)
.map_err(|error| format!("FLEX family design_entry row: {error}"))?;
let time_exit_chunk = self
.design_exit
.try_row_chunk(row..row + 1)
.map_err(|error| format!("FLEX family design_exit row: {error}"))?;
let time_derivative_chunk = self
.design_derivative_exit
.try_row_chunk(row..row + 1)
.map_err(|error| format!("FLEX family design_derivative_exit row: {error}"))?;
let marginal_chunk = self
.marginal_design
.try_row_chunk(row..row + 1)
.map_err(|error| format!("FLEX family marginal_design row: {error}"))?;
let logslope_chunk = self
.logslope_layout
.coefficient_design()
.try_row_chunk(row..row + 1)
.map_err(|error| format!("FLEX family logslope design row: {error}"))?;
let entry_row = time_entry_chunk.row(0);
let exit_row = time_exit_chunk.row(0);
let derivative_row = time_derivative_chunk.row(0);
let marginal_row = marginal_chunk.row(0);
let logslope_row = logslope_chunk.row(0);
let mut q0_gradient = vec![0.0; dimension];
let mut q1_gradient = vec![0.0; dimension];
let mut qd1_gradient = vec![0.0; dimension];
let (q0, q1, qd1) = if self.flex_timewiggle_active() {
let time_tail = self.time_wiggle_range();
let base_width = time_tail.start;
let base_range = slices.time.start..slices.time.start + base_width;
add_coefficient_row(
&mut q0_gradient,
&base_range,
entry_row.slice(s![..base_width]),
"timewiggle entry base",
)?;
add_coefficient_row(
&mut q1_gradient,
&base_range,
exit_row.slice(s![..base_width]),
"timewiggle exit base",
)?;
add_coefficient_row(
&mut qd1_gradient,
&base_range,
derivative_row.slice(s![..base_width]),
"timewiggle derivative base",
)?;
add_coefficient_row(
&mut q0_gradient,
&slices.marginal,
marginal_row,
"timewiggle entry marginal",
)?;
add_coefficient_row(
&mut q1_gradient,
&slices.marginal,
marginal_chunk.row(0),
"timewiggle exit marginal",
)?;
let beta_time = &block_states[0].beta;
let beta_time_base = beta_time.slice(s![..base_width]);
let beta_time_wiggle = beta_time.slice(s![time_tail.clone()]);
let h0 = entry_row.slice(s![..base_width]).dot(&beta_time_base)
+ self.offset_entry[row]
+ block_states[1].eta[row];
let h1 = exit_row.slice(s![..base_width]).dot(&beta_time_base)
+ self.offset_exit[row]
+ block_states[1].eta[row];
let d_raw = derivative_row.slice(s![..base_width]).dot(&beta_time_base)
+ self.derivative_offset_exit[row];
let h0_jet = lifted_coefficient_affine::<J>(
h0,
q0_gradient,
coefficient_direction,
design_targets_marginal,
first.entry,
second.entry,
);
let h1_jet = lifted_coefficient_affine::<J>(
h1,
q1_gradient,
coefficient_direction,
design_targets_marginal,
first.exit,
second.exit,
);
let d_raw_jet = lifted_coefficient_affine::<J>(
d_raw,
qd1_gradient,
coefficient_direction,
false,
first.derivative_exit,
second.derivative_exit,
);
let beta_wiggle_jets: Vec<Dual2<J>> = time_tail
.clone()
.map(|local_axis| {
lifted_coefficient_affine::<J>(
beta_time[local_axis],
one_hot_coefficient_gradient(dimension, slices.time.start + local_axis),
coefficient_direction,
false,
0.0,
0.0,
)
})
.collect();
let (entry_geometry, entry_basis_d5) = self
.time_wiggle_geometry_with_basis_d5(
Array1::from_vec(vec![h0]).view(),
beta_time_wiggle,
)?
.ok_or_else(|| {
"FLEX family timewiggle entry geometry is unavailable".to_string()
})?;
let (exit_geometry, exit_basis_d5) = self
.time_wiggle_geometry_with_basis_d5(
Array1::from_vec(vec![h1]).view(),
beta_time.slice(s![time_tail]),
)?
.ok_or_else(|| "FLEX family timewiggle exit geometry is unavailable".to_string())?;
let entry_basis =
TimewiggleBasisDerivativeRows::from_geometry(&entry_geometry, &entry_basis_d5, 0);
let exit_basis =
TimewiggleBasisDerivativeRows::from_geometry(&exit_geometry, &exit_basis_d5, 0);
let q = timewiggle_q_from_basis_derivative_rows(
&h0_jet,
&h1_jet,
&d_raw_jet,
&beta_wiggle_jets,
&entry_basis,
&exit_basis,
TimewiggleQBaseValues {
q0: q_values.q0,
q1: q_values.q1,
dq1_dh1: exit_geometry.dq_dq0[0],
},
)?;
if q.q0.value().to_bits() != q_values.q0.to_bits()
|| q.q1.value().to_bits() != q_values.q1.to_bits()
|| q.qd1.value().to_bits() != q_values.qd1.to_bits()
{
return Err(
"FLEX family generic timewiggle q values drifted from the current f64 row"
.to_string(),
);
}
(q.q0, q.q1, q.qd1)
} else {
add_coefficient_row(&mut q0_gradient, &slices.time, entry_row, "entry time")?;
add_coefficient_row(&mut q1_gradient, &slices.time, exit_row, "exit time")?;
add_coefficient_row(
&mut qd1_gradient,
&slices.time,
derivative_row,
"derivative time",
)?;
add_coefficient_row(
&mut q0_gradient,
&slices.marginal,
marginal_row,
"entry marginal",
)?;
add_coefficient_row(
&mut q1_gradient,
&slices.marginal,
marginal_chunk.row(0),
"exit marginal",
)?;
(
lifted_coefficient_affine::<J>(
q_values.q0,
q0_gradient,
coefficient_direction,
design_targets_marginal,
first.entry,
second.entry,
),
lifted_coefficient_affine::<J>(
q_values.q1,
q1_gradient,
coefficient_direction,
design_targets_marginal,
first.exit,
second.exit,
),
lifted_coefficient_affine::<J>(
q_values.qd1,
qd1_gradient,
coefficient_direction,
false,
first.derivative_exit,
second.derivative_exit,
),
)
};
let mut g_gradient = vec![0.0; dimension];
add_coefficient_row(&mut g_gradient, &slices.logslope, logslope_row, "logslope")?;
let g = lifted_coefficient_affine::<J>(
block_states[2].eta[row],
g_gradient,
coefficient_direction,
design_targets_logslope,
0.0,
0.0,
);
let zero =
|| lifted_coefficient_affine::<J>(0.0, vec![0.0; dimension], None, false, 0.0, 0.0);
let template = zero();
let mut du = vec![template.clone(); primary.total];
du[primary.q0] = tangent_jet(&q0);
du[primary.q1] = tangent_jet(&q1);
du[primary.qd1] = tangent_jet(&qd1);
du[primary.g] = tangent_jet(&g);
if let (Some(primary_range), Some(coefficient_range), Some(beta_h)) = (
primary.h.as_ref(),
slices.score_warp.as_ref(),
self.flex_score_beta(block_states)?,
) {
if primary_range.len() != coefficient_range.len() {
return Err("FLEX family score-warp primary/coefficient widths differ".to_string());
}
for local_axis in 0..primary_range.len() {
let coefficient_axis = coefficient_range.start + local_axis;
let coefficient = lifted_coefficient_affine::<J>(
beta_h[local_axis],
one_hot_coefficient_gradient(dimension, coefficient_axis),
coefficient_direction,
false,
0.0,
0.0,
);
du[primary_range.start + local_axis] = tangent_jet(&coefficient);
}
}
if let (Some(primary_range), Some(coefficient_range), Some(beta_w)) = (
primary.w.as_ref(),
slices.link_dev.as_ref(),
self.flex_link_beta(block_states)?,
) {
if primary_range.len() != coefficient_range.len() {
return Err(
"FLEX family link-deviation primary/coefficient widths differ".to_string(),
);
}
for local_axis in 0..primary_range.len() {
let coefficient_axis = coefficient_range.start + local_axis;
let coefficient = lifted_coefficient_affine::<J>(
beta_w[local_axis],
one_hot_coefficient_gradient(dimension, coefficient_axis),
coefficient_direction,
false,
0.0,
0.0,
);
du[primary_range.start + local_axis] = tangent_jet(&coefficient);
}
}
let o_infl = self.influence_index_offset(row, block_states)?;
if let (Some(primary_axis), Some(coefficient_range), Some(influence)) = (
primary.infl,
slices.influence.as_ref(),
self.influence_absorber.as_ref(),
) {
let mut influence_gradient = vec![0.0; dimension];
add_coefficient_row(
&mut influence_gradient,
coefficient_range,
influence.row(row),
"influence",
)?;
let influence_jet = lifted_coefficient_affine::<J>(
o_infl,
influence_gradient,
coefficient_direction,
false,
0.0,
0.0,
);
du[primary_axis] = tangent_jet(&influence_jet);
}
let probit_scale = self.probit_frailty_scale();
if !probit_scale.is_finite() || probit_scale <= 0.0 {
return Err(format!(
"survival marginal-slope FLEX family probit scale must be finite and positive, got {probit_scale}"
));
}
let scale_ratio = Dual2 {
v: J::constant(1.0, dimension),
g: J::constant(first.probit_scale / probit_scale, dimension),
h: J::constant(second.probit_scale / probit_scale, dimension),
};
Ok(FlexFamilyCoefficientJets {
template,
q0,
q1,
qd1,
g,
scale_ratio,
du,
o_infl,
})
}
fn flex_family_direction_row_terms_generic<J: FlexCoefficientJet>(
&self,
row: usize,
block_states: &[ParameterBlockState],
first: FlexFamilyRowDirection,
second: FlexFamilyRowDirection,
coefficient_direction: Option<&FlexCoefficientRowDirection<'_>>,
) -> Result<FlexFamilyDirectionRowTerms, String> {
self.ensure_scalar_flex_exact_score_geometry("FLEX family-direction row program")?;
let expected_blocks = 3
+ usize::from(self.score_warp.is_some())
+ usize::from(self.link_dev.is_some())
+ usize::from(self.influence_absorber.is_some());
if block_states.len() != expected_blocks {
return Err(format!(
"survival marginal-slope FLEX family coefficient map expects {expected_blocks} block states, got {}",
block_states.len(),
));
}
let primary = flex_primary_slices(self);
let slices = block_slices(self, block_states);
self.validate_flex_family_coefficient_state(
row,
block_states,
&slices,
first,
second,
coefficient_direction,
)?;
let jets = self.build_flex_family_coefficient_jets::<J>(
row,
block_states,
&primary,
&slices,
first,
second,
coefficient_direction,
)?;
if survival_derivative_guard_violated(jets.qd1.value(), self.derivative_guard) {
return Err(SurvivalMarginalSlopeError::MonotonicityViolation {
reason: format!(
"survival marginal-slope monotonicity violated at row {row}: qd1={:.3e} < guard={:.3e}",
jets.qd1.value(),
self.derivative_guard,
),
}
.into());
}
let g_value = jets.g.value();
let beta_h = self.flex_score_beta(block_states)?;
let beta_w = self.flex_link_beta(block_states)?;
let (a0, _) = self.solve_row_survival_intercept_with_slot(
jets.q0.value(),
g_value,
beta_h,
beta_w,
Some((row, SurvivalInterceptSlotKind::Entry)),
)?;
let (a1, _) = self.solve_row_survival_intercept_with_slot(
jets.q1.value(),
g_value,
beta_h,
beta_w,
Some((row, SurvivalInterceptSlotKind::Exit)),
)?;
let entry_cached = self.build_cached_partition(&primary, a0, g_value, beta_h, beta_w)?;
let exit_cached = self.build_cached_partition(&primary, a1, g_value, beta_h, beta_w)?;
let evaluate_timepoint =
|q: &Dual2<J>, a: f64, cached: &CachedPartitionCells| -> Result<_, String> {
let d_check = self.evaluate_survival_denom_d(a, g_value, beta_h, beta_w)?;
let (obs_coeff, obs_fixed) =
observed_fixed_for(self, &primary, row, a, g_value, beta_h, beta_w)?;
let cells = cells_from_cached(cached);
flex_timepoint_inputs_generic(
&jets.template,
&jets.g,
&jets.du,
a,
d_check,
primary.g,
primary.infl,
q,
&jets.scale_ratio,
self.observed_score_projection(row),
jets.o_infl,
obs_coeff,
&obs_fixed,
&cells,
)
};
let (eta0, _, _) = evaluate_timepoint(&jets.q0, a0, &entry_cached)?;
let (eta1, chi1, d1) = evaluate_timepoint(&jets.q1, a1, &exit_cached)?;
if !chi1.value().is_finite() || chi1.value() <= 0.0 {
return Err(SurvivalMarginalSlopeError::NumericalFailure {
reason: format!(
"survival marginal-slope row {row} produced non-positive observed chi1={:.3e}",
chi1.value(),
),
}
.into());
}
let output = flex_row_nll(
&eta0,
&eta1,
&chi1,
&d1,
&jets.q1,
&jets.qd1,
surv_stack(eta0.value())?,
surv_stack(eta1.value())?,
self.weights[row],
self.event[row],
);
Ok(FlexFamilyDirectionRowTerms {
first: output.g.owned_base_terms(),
second: output.h.owned_base_terms(),
directional: output.g.owned_directional_terms(),
})
}
pub(crate) fn flex_family_direction_row_terms(
&self,
row: usize,
block_states: &[ParameterBlockState],
first: FlexFamilyRowDirection,
second: FlexFamilyRowDirection,
beta_direction: Option<&Array1<f64>>,
) -> Result<FlexFamilyDirectionRowTerms, String> {
if let Some(beta_direction) = beta_direction {
let coefficient_direction = FlexCoefficientRowDirection::Beta(beta_direction);
self.flex_family_direction_row_terms_generic::<Jet3>(
row,
block_states,
first,
second,
Some(&coefficient_direction),
)
} else {
self.flex_family_direction_row_terms_generic::<Jet2>(
row,
block_states,
first,
second,
None,
)
}
}
pub(crate) fn flex_family_design_direction_row_terms(
&self,
row: usize,
block_states: &[ParameterBlockState],
first: FlexFamilyRowDirection,
second: FlexFamilyRowDirection,
block: usize,
derivative_row: &Array1<f64>,
) -> Result<FlexFamilyDirectionRowTerms, String> {
let slices = block_slices(self, block_states);
let (beta, coefficient_range) = match block {
1 => (&block_states[1].beta, slices.marginal),
2 => (&block_states[2].beta, slices.logslope),
_ => {
return Err(format!(
"survival marginal-slope FLEX family design direction supports marginal/logslope blocks 1 or 2, got block {block}"
));
}
};
let coefficient_direction = FlexCoefficientRowDirection::Design {
block,
derivative_row,
beta,
coefficient_range,
};
self.flex_family_direction_row_terms_generic::<Jet3>(
row,
block_states,
first,
second,
Some(&coefficient_direction),
)
}
}
use gam_math::nested_dual::{Dual2, JetField};
#[cfg(test)]
mod moment_engine_tests {
use super::*;
use crate::cubic_cell_kernel::{DenestedCubicCell, reduce_sextic_moments};
use crate::marginal_slope_shared::eval_coeff4_at;
use gam_math::nested_dual::Dual22;
use std::hint::black_box;
use std::time::Instant;
#[test]
fn dual2_flexjet_scales_runtime_channels_by_total_homogeneous_order() {
let factors = [2.0, 3.0, 5.0, 7.0, 11.0];
let original = Dual2 {
v: Jet2::from_parts(1.0, &[2.0, 3.0], &[4.0, 5.0, 6.0, 7.0]),
g: Jet2::from_parts(8.0, &[9.0, 10.0], &[11.0, 12.0, 13.0, 14.0]),
h: Jet2::from_parts(15.0, &[16.0, 17.0], &[18.0, 19.0, 20.0, 21.0]),
};
let scaled = original.scale_homogeneous_orders(factors);
let assert_part = |actual: &Jet2, expected: &Jet2, outer_order: usize| {
assert_eq!(actual.v, factors[outer_order] * expected.v);
for axis in 0..expected.g.len() {
assert_eq!(actual.g[axis], factors[outer_order + 1] * expected.g[axis]);
}
for axis in 0..expected.h.len() {
assert_eq!(actual.h[axis], factors[outer_order + 2] * expected.h[axis]);
}
};
assert_eq!(<Dual2<Jet2> as FlexJet>::ORDER, 4);
assert_part(&scaled.v, &original.v, 0);
assert_part(&scaled.g, &original.g, 1);
assert_part(&scaled.h, &original.h, 2);
}
#[test]
fn recursive_dual22_flexjet_scaling_matches_nested_total_order() {
let factors = [2.0, 3.0, 5.0, 7.0, 11.0];
let original_channels = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
let scaled = Dual22::from_channels(original_channels)
.scale_homogeneous_orders(factors)
.channels();
let total_orders = [0usize, 1, 1, 2, 2, 2, 3, 3, 4];
for channel in 0..scaled.len() {
assert_eq!(
scaled[channel],
factors[total_orders[channel]] * original_channels[channel]
);
}
}
#[test]
fn generic_q_jet_carries_exact_beta_family_calibration_channels() {
let q = 0.37;
let zero = Jet2::from_parts(0.0, &[0.0], &[]);
let template = Dual2 {
v: zero.clone(),
g: zero.clone(),
h: zero.clone(),
};
let q_jet = Dual2 {
v: Jet2::primary(q, 0, 1),
g: Jet2::from_parts(1.0, &[0.0], &[]),
h: zero,
};
let scale_ratio = const_jet_like(&template, 1.0);
let residual =
calibration_residual_jet(&template, &template, 0, &[], &q_jet, &scale_ratio, &[]);
let phi = crate::probability::normal_pdf(q);
let expected = [
-phi,
q * phi,
(1.0 - q * q) * phi,
(q * q * q - 3.0 * q) * phi,
];
let check = |actual: f64, wanted: f64, channel: &str| {
let tolerance = 256.0 * f64::EPSILON * (1.0 + actual.abs().max(wanted.abs()));
assert!(
(actual - wanted).abs() <= tolerance,
"{channel}: actual={actual:.17e}, wanted={wanted:.17e}, tolerance={tolerance:.3e}"
);
};
check(residual.v.v, 0.0, "value");
check(residual.v.g[0], expected[0], "beta first");
check(residual.v.h[0], expected[1], "beta second");
check(residual.g.v, expected[0], "family first");
check(residual.g.g[0], expected[1], "family-beta");
check(residual.g.h[0], expected[2], "family-beta-beta");
check(residual.h.v, expected[1], "family second");
check(residual.h.g[0], expected[2], "family-family-beta");
check(residual.h.h[0], expected[3], "family-family-beta-beta");
}
#[derive(Clone)]
struct ForcedOrder4<J>(J);
impl<J: FlexJet> JetField for ForcedOrder4<J> {
#[inline(always)]
fn value(&self) -> f64 {
self.0.value()
}
#[inline(always)]
fn add(&self, other: &Self) -> Self {
Self(self.0.add(&other.0))
}
#[inline(always)]
fn sub(&self, other: &Self) -> Self {
Self(self.0.sub(&other.0))
}
#[inline(always)]
fn mul(&self, other: &Self) -> Self {
Self(self.0.mul(&other.0))
}
#[inline(always)]
fn neg(&self) -> Self {
Self(self.0.neg())
}
#[inline(always)]
fn scale(&self, scale: f64) -> Self {
Self(self.0.scale(scale))
}
#[inline(always)]
fn compose_unary(&self, derivatives: [f64; 5]) -> Self {
Self(self.0.compose_unary(derivatives))
}
}
impl<J: FlexJet> FlexJet for ForcedOrder4<J> {
const ORDER: usize = 4;
fn scale_homogeneous_orders(&self, factors: [f64; 5]) -> Self {
Self(self.0.scale_homogeneous_orders(factors))
}
}
fn qprime_coeffs_jet<J: FlexJet>(c: &[J; 4]) -> [J; 6] {
let (c0, c1, c2, c3) = (&c[0], &c[1], &c[2], &c[3]);
let d0 = c0.mul(c1);
let d1 = add_const(&c1.mul(c1).add(&c0.mul(c2).scale(2.0)), 1.0);
let d2 = c0.mul(c3).add(&c1.mul(c2)).scale(3.0);
let d3 = c1.mul(c3).scale(4.0).add(&c2.mul(c2).scale(2.0));
let d4 = c2.mul(c3).scale(5.0);
let d5 = c3.mul(c3).scale(3.0);
[d0, d1, d2, d3, d4, d5]
}
fn cell_q_at_jet<J: FlexJet>(c: &[J; 4], z: &J) -> J {
let eta = c[3].mul(z).add(&c[2]).mul(z).add(&c[1]).mul(z).add(&c[0]);
z.mul(z).add(&eta.mul(&eta)).scale(0.5)
}
fn boundary_edge_term_jet<J: FlexJet>(
c: &[J; 4],
z: &J,
z_pow_n: &J,
finite: bool,
) -> Option<J> {
if !finite {
return None;
}
let q = cell_q_at_jet(c, z);
let w = exp_jet(&q.scale(-1.0));
Some(z_pow_n.mul(&w))
}
fn cell_moment_recurrence_jet<J: FlexJet>(
c: &[J; 4],
z_left: &J,
left_finite: bool,
z_right: &J,
right_finite: bool,
base_m0_m4: &[J; 5],
max_degree: usize,
) -> Vec<J> {
let d = qprime_coeffs_jet(c);
let inv_lead = recip(&d[5]);
let mut moments: Vec<J> = base_m0_m4.iter().cloned().collect();
if max_degree < 5 {
moments.truncate(max_degree + 1);
return moments;
}
let one_l = recip(z_left).mul(z_left);
let one_r = recip(z_right).mul(z_right);
let mut left_pow = one_l;
let mut right_pow = one_r;
for n in 0..=(max_degree - 5) {
let b_left = boundary_edge_term_jet(c, z_left, &left_pow, left_finite);
let b_right = boundary_edge_term_jet(c, z_right, &right_pow, right_finite);
let mut b_n = match (b_right, b_left) {
(Some(r), Some(l)) => r.sub(&l),
(Some(r), None) => r,
(None, Some(l)) => l.scale(-1.0),
(None, None) => moments[0].scale(0.0),
};
let mut numer = if n == 0 {
moments[0].scale(0.0)
} else {
moments[n - 1].scale(n as f64)
};
for j in 0..=4 {
numer = numer.sub(&d[j].mul(&moments[n + j]));
}
numer = numer.sub(&b_n);
moments.push(numer.mul(&inv_lead));
left_pow = if left_finite {
left_pow.mul(z_left)
} else {
b_n.scale(0.0)
};
right_pow = if right_finite {
right_pow.mul(z_right)
} else {
b_n = b_n.scale(0.0);
b_n
};
}
moments
}
fn observed_coeff_component_jet<J: FlexJet>(
template: &J,
k: usize,
coeff: [f64; 4],
dc_da: [f64; 4],
dc_db: [f64; 4],
dc_daa: [f64; 4],
dc_dab: [f64; 4],
dc_dbb: [f64; 4],
dc_daaa: [f64; 4],
dc_daab: [f64; 4],
dc_dabb: [f64; 4],
dc_dbbb: [f64; 4],
da: &J,
db: &J,
) -> J {
let dada = da.mul(da);
let dadb = da.mul(db);
let dbdb = db.mul(db);
let mut c = const_jet_like(template, coeff[k]);
c = c.add(&da.scale(dc_da[k])).add(&db.scale(dc_db[k]));
c = c
.add(&dada.scale(0.5 * dc_daa[k]))
.add(&dadb.scale(dc_dab[k]))
.add(&dbdb.scale(0.5 * dc_dbb[k]));
let inv6 = 1.0 / 6.0;
let half = 0.5;
c = c
.add(&dada.mul(da).scale(inv6 * dc_daaa[k]))
.add(&dada.mul(db).scale(half * dc_daab[k]))
.add(&dadb.mul(db).scale(half * dc_dabb[k]))
.add(&dbdb.mul(db).scale(inv6 * dc_dbbb[k]));
c
}
struct ObservedCoeffPack {
coeff: [f64; 4],
dc_da: [f64; 4],
dc_db: [f64; 4],
dc_daa: [f64; 4],
dc_dab: [f64; 4],
dc_dbb: [f64; 4],
dc_daaa: [f64; 4],
dc_daab: [f64; 4],
dc_dabb: [f64; 4],
dc_dbbb: [f64; 4],
}
fn flex_timepoint_eta_chi<J: FlexJet>(
a_jet: &J,
b_jet: &J,
z_obs: f64,
o_infl: f64,
pack: &ObservedCoeffPack,
rho_jet: &J,
tau_jet: &J,
) -> (J, J) {
let da = tangent_jet(a_jet);
let db = tangent_jet(b_jet);
let zero4 = [0.0_f64; 4];
let coeff_jets: [J; 4] = std::array::from_fn(|k| {
observed_coeff_component_jet(
a_jet,
k,
pack.coeff,
pack.dc_da,
pack.dc_db,
pack.dc_daa,
pack.dc_dab,
pack.dc_dbb,
pack.dc_daaa,
pack.dc_daab,
pack.dc_dabb,
pack.dc_dbbb,
&da,
&db,
)
});
let eta = add_const(&eval_coeff_jet_at(&coeff_jets, z_obs), o_infl).add(rho_jet);
let chi_jets: [J; 4] = std::array::from_fn(|k| {
observed_coeff_component_jet(
a_jet,
k,
pack.dc_da,
pack.dc_daa,
pack.dc_dab,
pack.dc_daaa,
pack.dc_daab,
pack.dc_dabb,
zero4,
zero4,
zero4,
zero4,
&da,
&db,
)
});
let chi = eval_coeff_jet_at(&chi_jets, z_obs).add(tau_jet);
(eta, chi)
}
#[test]
fn cell_moment_recurrence_jet_value_matches_numeric_932() {
let cell = DenestedCubicCell {
left: -1.5,
right: 2.0,
c0: 0.3,
c1: -0.4,
c2: 0.5,
c3: 0.2,
};
let base = [1.0_f64, 0.1, 0.6, -0.05, 0.4];
let max_degree = 12usize;
let reference =
reduce_sextic_moments(cell, base, max_degree).expect("numeric sextic moments");
let p = 3usize;
let konst = |x: f64| Jet2::from_parts(x, &vec![0.0; p], &[]);
let c = [
konst(cell.c0),
konst(cell.c1),
konst(cell.c2),
konst(cell.c3),
];
let zl = konst(cell.left);
let zr = konst(cell.right);
let base_jets = [
konst(base[0]),
konst(base[1]),
konst(base[2]),
konst(base[3]),
konst(base[4]),
];
let moments = cell_moment_recurrence_jet(
&c,
&zl,
cell.left.is_finite(),
&zr,
cell.right.is_finite(),
&base_jets,
max_degree,
);
assert_eq!(moments.len(), reference.len(), "moment count");
for (n, (m, r)) in moments.iter().zip(reference.iter()).enumerate() {
assert!(
(m.value() - r).abs() <= 1e-9 * (1.0 + r.abs()),
"moment {n}: jet value {} != numeric {}",
m.value(),
r
);
}
}
#[test]
fn measure_base_moment_instantiated_order_vs_forced_four_932() {
use crate::cubic_cell_kernel::evaluate_cell_moments;
let cell = DenestedCubicCell {
left: -1.2,
right: 1.7,
c0: 0.25,
c1: -0.35,
c2: 0.4,
c3: 0.15,
};
let numeric = evaluate_cell_moments(cell, 28)
.expect("numeric cell moments")
.moments
.into_vec();
let p = 6usize;
let gradient = |scale: f64| -> Vec<f64> {
(0..p)
.map(|axis| scale * (axis as f64 + 1.0) / p as f64)
.collect()
};
let hessian = |scale: f64| -> Vec<f64> {
let mut out = vec![0.0; p * p];
for i in 0..p {
for j in 0..p {
out[i * p + j] = scale * (i + j + 1) as f64 / (p * p) as f64;
}
}
out
};
let c1: [Jet1; 4] = [
Jet1 {
v: cell.c0,
g: gradient(0.13),
},
Jet1 {
v: cell.c1,
g: gradient(-0.21),
},
Jet1 {
v: cell.c2,
g: gradient(0.17),
},
Jet1 {
v: cell.c3,
g: gradient(0.09),
},
];
let left1 = Jet1 {
v: cell.left,
g: gradient(-0.23),
};
let right1 = Jet1 {
v: cell.right,
g: gradient(0.31),
};
let forced_c1 = c1.clone().map(ForcedOrder4);
let forced_left1 = ForcedOrder4(left1.clone());
let forced_right1 = ForcedOrder4(right1.clone());
let c2: [Jet2; 4] = [
Jet2::from_parts(cell.c0, &gradient(0.13), &hessian(0.017)),
Jet2::from_parts(cell.c1, &gradient(-0.21), &hessian(-0.011)),
Jet2::from_parts(cell.c2, &gradient(0.17), &hessian(0.019)),
Jet2::from_parts(cell.c3, &gradient(0.09), &hessian(-0.007)),
];
let left2 = Jet2::from_parts(cell.left, &gradient(-0.23), &hessian(0.013));
let right2 = Jet2::from_parts(cell.right, &gradient(0.31), &hessian(-0.015));
let forced_c2 = c2.clone().map(ForcedOrder4);
let forced_left2 = ForcedOrder4(left2.clone());
let forced_right2 = ForcedOrder4(right2.clone());
let native1 = base_moment_jets(&c1, &left1, true, &right1, true, &numeric);
let historical1 = base_moment_jets(
&forced_c1,
&forced_left1,
true,
&forced_right1,
true,
&numeric,
);
let native2 = base_moment_jets(&c2, &left2, true, &right2, true, &numeric);
let historical2 = base_moment_jets(
&forced_c2,
&forced_left2,
true,
&forced_right2,
true,
&numeric,
);
let close = |got: f64, want: f64, label: &str| {
let tolerance = 1e-12 * got.abs().max(want.abs()).max(1.0);
assert!((got - want).abs() <= tolerance, "{label}: {got} != {want}");
};
for n in 0..5 {
close(
native1[n].v,
historical1[n].0.v,
&format!("Jet1 M{n} value"),
);
for i in 0..p {
close(
native1[n].g[i],
historical1[n].0.g[i],
&format!("Jet1 M{n} gradient[{i}]"),
);
}
close(
native2[n].v,
historical2[n].0.v,
&format!("Jet2 M{n} value"),
);
for i in 0..p {
close(
native2[n].g[i],
historical2[n].0.g[i],
&format!("Jet2 M{n} gradient[{i}]"),
);
for j in 0..p {
close(
native2[n].h[i * p + j],
historical2[n].0.h[i * p + j],
&format!("Jet2 M{n} Hessian[{i},{j}]"),
);
}
}
}
let iterations = if cfg!(debug_assertions) { 4 } else { 5_000 };
let mut native1_best = f64::INFINITY;
let mut historical1_best = f64::INFINITY;
let mut native2_best = f64::INFINITY;
let mut historical2_best = f64::INFINITY;
for _ in 0..5 {
let start = Instant::now();
for _ in 0..iterations {
black_box(base_moment_jets(
black_box(&c1),
black_box(&left1),
true,
black_box(&right1),
true,
black_box(&numeric),
));
}
native1_best = native1_best.min(start.elapsed().as_secs_f64());
let start = Instant::now();
for _ in 0..iterations {
black_box(base_moment_jets(
black_box(&forced_c1),
black_box(&forced_left1),
true,
black_box(&forced_right1),
true,
black_box(&numeric),
));
}
historical1_best = historical1_best.min(start.elapsed().as_secs_f64());
let start = Instant::now();
for _ in 0..iterations {
black_box(base_moment_jets(
black_box(&c2),
black_box(&left2),
true,
black_box(&right2),
true,
black_box(&numeric),
));
}
native2_best = native2_best.min(start.elapsed().as_secs_f64());
let start = Instant::now();
for _ in 0..iterations {
black_box(base_moment_jets(
black_box(&forced_c2),
black_box(&forced_left2),
true,
black_box(&forced_right2),
true,
black_box(&numeric),
));
}
historical2_best = historical2_best.min(start.elapsed().as_secs_f64());
}
let to_ns = |seconds: f64| seconds * 1e9 / iterations as f64;
let native1_ns = to_ns(native1_best);
let historical1_ns = to_ns(historical1_best);
let native2_ns = to_ns(native2_best);
let historical2_ns = to_ns(historical2_best);
eprintln!(
"FLEX-MOMENT-ORDER-932 jet1={native1_ns:.2} ns historical-order4-jet1={historical1_ns:.2} ns speedup1={:.3}x jet2={native2_ns:.2} ns historical-order4-jet2={historical2_ns:.2} ns speedup2={:.3}x",
historical1_ns / native1_ns,
historical2_ns / native2_ns,
);
}
#[test]
fn base_moment_jets_first_derivative_matches_fd_932() {
use crate::cubic_cell_kernel::evaluate_cell_moments;
let c0 = [0.25_f64, -0.35, 0.4, 0.15];
let zl0 = -1.2_f64;
let zr0 = 1.7_f64;
let dc = [0.13_f64, 0.21, -0.17, 0.09];
let v_l = -0.23_f64;
let v_r = 0.31_f64;
let cell_at = |theta: f64| DenestedCubicCell {
left: zl0 + theta * v_l,
right: zr0 + theta * v_r,
c0: c0[0] + theta * dc[0],
c1: c0[1] + theta * dc[1],
c2: c0[2] + theta * dc[2],
c3: c0[3] + theta * dc[3],
};
let max_degree = 10usize;
let moments_at = |theta: f64| -> Vec<f64> {
evaluate_cell_moments(cell_at(theta), max_degree)
.expect("numeric cell moments")
.moments
.into_vec()
};
let numeric0 = moments_at(0.0);
let p = 1usize;
let seeded = |x: f64, vel: f64| {
let mut g = vec![0.0; p];
g[0] = vel;
Jet1 { v: x, g }
};
let c_jets = [
seeded(c0[0], dc[0]),
seeded(c0[1], dc[1]),
seeded(c0[2], dc[2]),
seeded(c0[3], dc[3]),
];
let zl_jet = seeded(zl0, v_l);
let zr_jet = seeded(zr0, v_r);
let m_jets = base_moment_jets(&c_jets, &zl_jet, true, &zr_jet, true, &numeric0);
let h = 1e-6_f64;
let mp = moments_at(h);
let mm = moments_at(-h);
for n in 0..5 {
let fd = (mp[n] - mm[n]) / (2.0 * h);
let jet = &m_jets[n];
assert!(
(jet.value() - numeric0[n]).abs() <= 1e-12 * (1.0 + numeric0[n].abs()),
"M_{n} value {} != numeric {}",
jet.value(),
numeric0[n]
);
assert!(
(jet.g[0] - fd).abs() <= 1e-5 * (1.0 + fd.abs()),
"M_{n} dθ analytic {} != FD {}",
jet.g[0],
fd
);
}
}
#[test]
fn base_moment_jets_second_derivative_matches_fd_932() {
use crate::cubic_cell_kernel::evaluate_cell_moments;
let c0 = [0.25_f64, -0.35, 0.4, 0.15];
let zl0 = -1.2_f64;
let zr0 = 1.7_f64;
let dc = [0.13_f64, 0.21, -0.17, 0.09];
let v_l = -0.23_f64;
let v_r = 0.31_f64;
let cell_at = |theta: f64| DenestedCubicCell {
left: zl0 + theta * v_l,
right: zr0 + theta * v_r,
c0: c0[0] + theta * dc[0],
c1: c0[1] + theta * dc[1],
c2: c0[2] + theta * dc[2],
c3: c0[3] + theta * dc[3],
};
let max_degree = 27usize;
let moments_at = |theta: f64| -> Vec<f64> {
evaluate_cell_moments(cell_at(theta), max_degree)
.expect("numeric cell moments")
.moments
.into_vec()
};
let analytic_first = |theta: f64, n: usize| -> f64 {
let numeric = moments_at(theta);
let seeded = |x: f64, vel: f64| {
let g = vec![vel];
Jet2::from_parts(x, &g, &[])
};
let cell = cell_at(theta);
let c_jets = [
seeded(cell.c0, dc[0]),
seeded(cell.c1, dc[1]),
seeded(cell.c2, dc[2]),
seeded(cell.c3, dc[3]),
];
let zl_jet = seeded(cell.left, v_l);
let zr_jet = seeded(cell.right, v_r);
let m = base_moment_jets(&c_jets, &zl_jet, true, &zr_jet, true, &numeric);
m[n].g[0]
};
let numeric0 = moments_at(0.0);
let seeded = |x: f64, vel: f64| {
let g = vec![vel];
Jet2::from_parts(x, &g, &[])
};
let c_jets = [
seeded(c0[0], dc[0]),
seeded(c0[1], dc[1]),
seeded(c0[2], dc[2]),
seeded(c0[3], dc[3]),
];
let zl_jet = seeded(zl0, v_l);
let zr_jet = seeded(zr0, v_r);
let m_jets = base_moment_jets(&c_jets, &zl_jet, true, &zr_jet, true, &numeric0);
let h = 1e-5_f64;
for n in 0..5 {
let fd2 = (analytic_first(h, n) - analytic_first(-h, n)) / (2.0 * h);
let hess = m_jets[n].h[0];
assert!(
(hess - fd2).abs() <= 2e-4 * (1.0 + fd2.abs()),
"M_{n} d²θ analytic {} != FD-of-analytic {}",
hess,
fd2
);
}
}
#[test]
fn flex_timepoint_eta_chi_value_and_grad_932() {
let z_obs = 0.7_f64;
let o_infl = 0.05_f64;
let pack = ObservedCoeffPack {
coeff: [0.2, -0.3, 0.15, 0.05],
dc_da: [1.1, 0.2, 0.03, 0.0],
dc_db: [0.4, 1.05, 0.1, 0.02],
dc_daa: [0.07, 0.02, 0.0, 0.0],
dc_dab: [0.2, 0.09, 0.01, 0.0],
dc_dbb: [0.11, 0.04, 0.005, 0.0],
dc_daaa: [0.003, 0.0, 0.0, 0.0],
dc_daab: [0.006, 0.001, 0.0, 0.0],
dc_dabb: [0.004, 0.002, 0.0, 0.0],
dc_dbbb: [0.008, 0.001, 0.0, 0.0],
};
let a0 = 0.3_f64;
let b0 = 1.2_f64;
let a_u = 0.25_f64;
let b_u = -0.4_f64;
let p = 1usize;
let a_jet = Jet2::from_parts(a0, &[a_u], &[]);
let b_jet = Jet2::from_parts(b0, &[b_u], &[]);
let zero = Jet2::from_parts(0.0, &vec![0.0; p], &[]);
let (eta, chi) = flex_timepoint_eta_chi(&a_jet, &b_jet, z_obs, o_infl, &pack, &zero, &zero);
let coeff_scalar = |da: f64, db: f64| -> [f64; 4] {
std::array::from_fn(|k| {
pack.coeff[k]
+ pack.dc_da[k] * da
+ pack.dc_db[k] * db
+ 0.5 * pack.dc_daa[k] * da * da
+ pack.dc_dab[k] * da * db
+ 0.5 * pack.dc_dbb[k] * db * db
+ pack.dc_daaa[k] * da * da * da / 6.0
+ 0.5 * pack.dc_daab[k] * da * da * db
+ 0.5 * pack.dc_dabb[k] * da * db * db
+ pack.dc_dbbb[k] * db * db * db / 6.0
})
};
let eta_scalar = |theta: f64| -> f64 {
let c = coeff_scalar(a_u * theta, b_u * theta);
eval_coeff4_scalar(&c, z_obs) + o_infl
};
let chi_scalar = |theta: f64| -> f64 {
let dc = coeff_scalar_da(&pack, a_u * theta, b_u * theta);
eval_coeff4_scalar(&dc, z_obs)
};
assert!(
(eta.value() - eta_scalar(0.0)).abs() <= 1e-12 * (1.0 + eta_scalar(0.0).abs()),
"eta value {} != {}",
eta.value(),
eta_scalar(0.0)
);
assert!(
(chi.value() - chi_scalar(0.0)).abs() <= 1e-12 * (1.0 + chi_scalar(0.0).abs()),
"chi value {} != {}",
chi.value(),
chi_scalar(0.0)
);
let h = 1e-6_f64;
let eta_fd = (eta_scalar(h) - eta_scalar(-h)) / (2.0 * h);
let chi_fd = (chi_scalar(h) - chi_scalar(-h)) / (2.0 * h);
assert!(
(eta.g[0] - eta_fd).abs() <= 1e-5 * (1.0 + eta_fd.abs()),
"eta grad {} != FD {}",
eta.g[0],
eta_fd
);
assert!(
(chi.g[0] - chi_fd).abs() <= 1e-5 * (1.0 + chi_fd.abs()),
"chi grad {} != FD {}",
chi.g[0],
chi_fd
);
}
fn eval_coeff4_scalar(c: &[f64; 4], z: f64) -> f64 {
let mut acc = 0.0;
for &ck in c.iter().rev() {
acc = acc * z + ck;
}
acc
}
fn coeff_scalar_da(pack: &ObservedCoeffPack, da: f64, db: f64) -> [f64; 4] {
std::array::from_fn(|k| {
pack.dc_da[k]
+ pack.dc_daa[k] * da
+ pack.dc_dab[k] * db
+ 0.5 * pack.dc_daaa[k] * da * da
+ pack.dc_daab[k] * da * db
+ 0.5 * pack.dc_dabb[k] * db * db
})
}
#[test]
fn cell_coeff_jets_value_and_grad_932() {
let p = 3usize;
let g_axis = 1usize;
let base_c = [0.2_f64, -0.3, 0.15, 0.05];
let mk_run = |seed: f64| -> Vec<[f64; 4]> {
(0..p)
.map(|u| std::array::from_fn(|k| seed * (1.0 + u as f64) * (1.0 + k as f64) * 0.01))
.collect()
};
let fixed = DenestedCellPrimaryFixedPartials {
dc_da: [1.1, 0.2, 0.03, 0.0],
dc_daa: [0.07, 0.02, 0.0, 0.0],
dc_daaa: [0.003, 0.0, 0.0, 0.0],
coeff_u: mk_run(0.9),
coeff_au: mk_run(0.4),
coeff_bu: mk_run(0.5),
coeff_aau: mk_run(0.12),
coeff_abu: mk_run(0.16),
coeff_bbu: mk_run(0.11),
coeff_aaau: mk_run(0.02),
coeff_aabu: mk_run(0.03),
coeff_abbu: mk_run(0.04),
coeff_bbbu: mk_run(0.05),
};
let a0 = 0.3_f64;
let a_u = 0.25_f64;
let v = [0.2_f64, -0.4, 0.33];
let seeded = |x: f64, vel: f64| {
let g = vec![vel];
Jet2::from_parts(x, &g, &[])
};
let a_jet = seeded(a0, a_u);
let da = tangent_jet(&a_jet);
let du: Vec<Jet2> = (0..p).map(|u| seeded(0.0, v[u])).collect();
let jets = cell_coeff_jets(&a_jet, base_c, &fixed, g_axis, &da, &du);
let scalar_c = |theta: f64| -> [f64; 4] {
let da = a_u * theta;
let db = v[g_axis] * theta;
std::array::from_fn(|k| {
let mut acc = base_c[k]
+ fixed.dc_da[k] * da
+ 0.5 * fixed.dc_daa[k] * da * da
+ fixed.dc_daaa[k] * da * da * da / 6.0;
for u in 0..p {
if u == g_axis {
continue;
}
let duu = v[u] * theta;
acc += fixed.coeff_u[u][k] * duu
+ fixed.coeff_au[u][k] * da * duu
+ 0.5 * fixed.coeff_aau[u][k] * da * da * duu
+ fixed.coeff_bu[u][k] * db * duu
+ fixed.coeff_abu[u][k] * da * db * duu
+ 0.5 * fixed.coeff_bbu[u][k] * db * db * duu
+ fixed.coeff_aaau[u][k] * da * da * da * duu / 6.0
+ 0.5 * fixed.coeff_aabu[u][k] * da * da * db * duu
+ 0.5 * fixed.coeff_abbu[u][k] * da * db * db * duu
+ fixed.coeff_bbbu[u][k] * db * db * db * duu / 6.0;
}
acc += fixed.coeff_u[g_axis][k] * db
+ fixed.coeff_au[g_axis][k] * da * db
+ 0.5 * fixed.coeff_aau[g_axis][k] * da * da * db
+ 0.5 * fixed.coeff_bu[g_axis][k] * db * db
+ 0.5 * fixed.coeff_abu[g_axis][k] * da * db * db
+ fixed.coeff_bbu[g_axis][k] * db * db * db / 6.0;
acc
})
};
let h = 1e-6_f64;
let c0 = scalar_c(0.0);
let cp = scalar_c(h);
let cm = scalar_c(-h);
for k in 0..4 {
assert!(
(jets[k].value() - c0[k]).abs() <= 1e-12 * (1.0 + c0[k].abs()),
"c_{k} value {} != {}",
jets[k].value(),
c0[k]
);
let fd = (cp[k] - cm[k]) / (2.0 * h);
assert!(
(jets[k].g[0] - fd).abs() <= 1e-5 * (1.0 + fd.abs()),
"c_{k} grad {} != FD {}",
jets[k].g[0],
fd
);
}
let da_iso = Jet2::from_parts(0.0, &vec![0.0; p], &[]);
let du_iso: Vec<Jet2> = (0..p).map(|u| Jet2::primary(0.0, u, p)).collect();
let jets_iso = cell_coeff_jets(&da_iso, base_c, &fixed, g_axis, &da_iso, &du_iso);
for k in 0..4 {
let hgg = jets_iso[k].h[g_axis * p + g_axis];
assert!(
(hgg - fixed.coeff_bu[g_axis][k]).abs()
<= 1e-12 * (1.0 + fixed.coeff_bu[g_axis][k].abs()),
"c_{k} Hess[g,g] {} != dc_dbb {} (2× = the pre-fix g-diagonal bug)",
hgg,
fixed.coeff_bu[g_axis][k]
);
}
}
#[test]
fn flex_timepoint_d_cell_value_and_grad_932() {
use crate::cubic_cell_kernel::evaluate_cell_moments;
let zl = -1.1_f64;
let zr = 1.6_f64;
let c_base = [0.2_f64, -0.3, 0.18, 0.06];
let dc_da = [1.05_f64, 0.22, 0.04, 0.0];
let dc_daa = [0.08_f64, 0.03, 0.0, 0.0];
let dc_daaa = [0.004_f64, 0.0, 0.0, 0.0];
let cell_at = |theta: f64| {
let c: [f64; 4] = std::array::from_fn(|k| {
c_base[k]
+ dc_da[k] * theta
+ 0.5 * dc_daa[k] * theta * theta
+ dc_daaa[k] * theta * theta * theta / 6.0
});
DenestedCubicCell {
left: zl,
right: zr,
c0: c[0],
c1: c[1],
c2: c[2],
c3: c[3],
}
};
let dc_da_at = |theta: f64| -> [f64; 4] {
std::array::from_fn(|k| dc_da[k] + dc_daa[k] * theta + 0.5 * dc_daaa[k] * theta * theta)
};
let max_degree = 10usize;
let moments_at = |theta: f64| -> Vec<f64> {
evaluate_cell_moments(cell_at(theta), max_degree)
.expect("numeric cell moments")
.moments
.into_vec()
};
let d_scalar = |theta: f64| -> f64 {
let m = moments_at(theta);
let chi = dc_da_at(theta);
let mut acc = 0.0;
for k in 0..4 {
acc += chi[k] * m[k];
}
acc * std::f64::consts::TAU.recip()
};
let seeded = |x: f64, vel: f64| {
let g = vec![vel];
Jet1 { v: x, g }
};
let cell0 = cell_at(0.0);
let c_jets = [
seeded(cell0.c0, dc_da[0]),
seeded(cell0.c1, dc_da[1]),
seeded(cell0.c2, dc_da[2]),
seeded(cell0.c3, dc_da[3]),
];
let dc_da0 = dc_da_at(0.0);
let chi_jets = [
seeded(dc_da0[0], dc_daa[0]),
seeded(dc_da0[1], dc_daa[1]),
seeded(dc_da0[2], dc_daa[2]),
seeded(dc_da0[3], dc_daa[3]),
];
let template = seeded(0.0, 0.0);
let edge_l = seeded(zl, 0.0); let edge_r = seeded(zr, 0.0);
let numeric0 = moments_at(0.0);
let d_jet = flex_timepoint_d_cell(
&template, &c_jets, &chi_jets, &edge_l, true, &edge_r, true, &numeric0,
);
assert!(
(d_jet.value() - d_scalar(0.0)).abs() <= 1e-10 * (1.0 + d_scalar(0.0).abs()),
"D value {} != {}",
d_jet.value(),
d_scalar(0.0)
);
let h = 1e-6_f64;
let fd = (d_scalar(h) - d_scalar(-h)) / (2.0 * h);
assert!(
(d_jet.g[0] - fd).abs() <= 1e-4 * (1.0 + fd.abs()),
"D grad (d_u) {} != FD {}",
d_jet.g[0],
fd
);
}
#[test]
fn cell_chi_poly_jets_value_and_grad_932() {
let p = 2usize;
let g_axis = 1usize;
let mk_run = |seed: f64| -> Vec<[f64; 4]> {
(0..p)
.map(|u| std::array::from_fn(|k| seed * (1.0 + u as f64) * (1.0 + k as f64) * 0.01))
.collect()
};
let fixed = DenestedCellPrimaryFixedPartials {
dc_da: [1.05, 0.22, 0.04, 0.0],
dc_daa: [0.08, 0.03, 0.0, 0.0],
dc_daaa: [0.004, 0.0, 0.0, 0.0],
coeff_u: mk_run(0.9),
coeff_au: mk_run(0.4),
coeff_bu: mk_run(0.5),
coeff_aau: mk_run(0.12),
coeff_abu: mk_run(0.16),
coeff_bbu: mk_run(0.11),
coeff_aaau: mk_run(0.02),
coeff_aabu: mk_run(0.03),
coeff_abbu: mk_run(0.04),
coeff_bbbu: mk_run(0.05),
};
let a_u = 0.25_f64;
let v = [0.2_f64, -0.4];
let seeded = |x: f64, vel: f64| {
let g = vec![vel];
Jet2::from_parts(x, &g, &[])
};
let a_jet = seeded(0.3, a_u);
let da = tangent_jet(&a_jet);
let du: Vec<Jet2> = (0..p).map(|u| seeded(0.0, v[u])).collect();
let chi = cell_chi_poly_jets(&a_jet, &fixed, g_axis, &da, &du);
let chi_scalar = |theta: f64| -> [f64; 4] {
let da = a_u * theta;
let db = v[g_axis] * theta;
std::array::from_fn(|k| {
let mut acc =
fixed.dc_da[k] + fixed.dc_daa[k] * da + 0.5 * fixed.dc_daaa[k] * da * da;
for u in 0..p {
let duu = v[u] * theta;
acc += fixed.coeff_au[u][k] * duu
+ fixed.coeff_aau[u][k] * da * duu
+ fixed.coeff_abu[u][k] * db * duu;
}
acc
})
};
let h = 1e-6_f64;
let c0 = chi_scalar(0.0);
let cp = chi_scalar(h);
let cm = chi_scalar(-h);
for k in 0..4 {
assert!(
(chi[k].value() - c0[k]).abs() <= 1e-12 * (1.0 + c0[k].abs()),
"chi_{k} value {} != {}",
chi[k].value(),
c0[k]
);
let fd = (cp[k] - cm[k]) / (2.0 * h);
assert!(
(chi[k].g[0] - fd).abs() <= 1e-5 * (1.0 + fd.abs()),
"chi_{k} grad {} != FD {}",
chi[k].g[0],
fd
);
}
}
fn make_g_only_flex_family(n: usize) -> SurvivalMarginalSlopeFamily {
let event: Array1<f64> =
Array1::from_iter((0..n).map(|i| if (i * 31 + 7) % 5 >= 3 { 1.0 } else { 0.0 }));
let weights: Array1<f64> =
Array1::from_iter((0..n).map(|i| 0.5 + ((i * 13 + 4) % 5) as f64 * 0.1));
let z: Array1<f64> = Array1::from_iter(
(0..n).map(|i| -1.0 + 2.0 * (((i * 17 + 5) % n) as f64 + 0.5) / (n as f64)),
);
let offset_entry: Array1<f64> = Array1::from_iter(
(0..n).map(|i| -0.4 + 0.7 * (((i * 11 + 3) % n) as f64 + 0.5) / (n as f64)),
);
let offset_exit: Array1<f64> = Array1::from_iter(
(0..n).map(|i| 0.1 + 0.6 * (((i * 19 + 7) % n) as f64 + 0.5) / (n as f64)),
);
let derivative_offset_exit: Array1<f64> =
Array1::from_iter((0..n).map(|i| 0.5 + 0.05 * ((i * 23 + 1) % 3) as f64));
let marginal_design = Array2::from_shape_fn((n, 1), |(i, _)| {
0.3 + 0.4 * (((i * 29 + 11) % n) as f64) / (n as f64)
});
let logslope_design = Array2::from_shape_fn((n, 1), |(i, _)| {
0.2 + 0.5 * (((i * 37 + 9) % n) as f64) / (n as f64)
});
SurvivalMarginalSlopeFamily {
n,
event: Arc::new(event),
weights: Arc::new(weights),
z: Arc::new(z.insert_axis(Axis(1))),
score_covariance: MarginalSlopeCovariance::diagonal(Array1::from(vec![1.0])).unwrap(),
gaussian_frailty_sd: None,
family_hyper: SurvivalMarginalSlopeFamilyHyperState::default(),
derivative_guard: 1e-6,
design_entry: DesignMatrix::from(Array2::zeros((n, 0))),
design_exit: DesignMatrix::from(Array2::zeros((n, 0))),
design_derivative_exit: DesignMatrix::from(Array2::zeros((n, 0))),
offset_entry: Arc::new(offset_entry),
offset_exit: Arc::new(offset_exit),
derivative_offset_exit: Arc::new(derivative_offset_exit),
marginal_design: DesignMatrix::from(marginal_design),
logslope_layout: DesignMatrix::from(logslope_design).into(),
score_warp: None,
link_dev: None,
influence_absorber: None,
time_linear_constraints: None,
time_wiggle_knots: None,
time_wiggle_degree: None,
time_wiggle_ncols: 0,
intercept_warm_starts: None,
auto_subsample_phase_counter: Arc::new(AtomicUsize::new(0)),
auto_subsample_last_rho: Arc::new(Mutex::new(None)),
}
}
#[test]
fn flex_third_arena_reuses_warmed_tape_932() {
let family = make_g_only_flex_family(16);
let primary = flex_primary_slices(&family);
let row = 5usize;
let g = 0.21_f64;
let q1 = family.offset_exit[row] + family.marginal_design.to_dense()[[row, 0]] * 0.15;
let a1 = family
.solve_row_survival_intercept_with_slot(
q1,
g,
None,
None,
Some((row, SurvivalInterceptSlotKind::Exit)),
)
.expect("intercept solve")
.0;
let cached = family
.build_cached_partition(&primary, a1, g, None, None)
.expect("cached partition");
let dir = Array1::from_iter((0..primary.total).map(|axis| 0.1 + 0.03 * axis as f64));
let run = || {
with_flex_third_jet_arena(|arena| {
family
.compute_survival_timepoint_directional_jet_from_cached(
row, &primary, q1, primary.q1, a1, g, None, None, 0.0, &cached, &dir, arena,
)
.expect("production third-order timepoint")
})
};
let retained_bytes = || with_flex_third_jet_arena(|arena| arena.allocated_bytes());
run();
let first = retained_bytes();
assert!(first > 0, "FLEX third arena did not retain its warm tape");
run();
assert_eq!(
retained_bytes(),
first,
"same-width FLEX third row grew its warmed arena"
);
}
#[test]
fn flex_timepoint_first_order_matches_jet2_and_fd_932() {
let n = 16usize;
let family = make_g_only_flex_family(n);
let primary = flex_primary_slices(&family);
let p = primary.total;
let row = 6usize;
let g = 0.19_f64;
let o_infl = 0.0_f64;
let m_beta = 0.15_f64;
let q1 = family.offset_exit[row] + family.marginal_design.to_dense()[[row, 0]] * m_beta;
let (a1, _d1) = family
.solve_row_survival_intercept_with_slot(
q1,
g,
None,
None,
Some((row, SurvivalInterceptSlotKind::Exit)),
)
.expect("intercept solve");
let first = family
.compute_survival_timepoint_first_order_exact(
row, &primary, q1, primary.q1, a1, g, None, None, o_infl,
)
.expect("grad-only jet1");
let full = family
.compute_survival_timepoint_exact_jet(
row, &primary, q1, primary.q1, a1, g, None, None, o_infl,
)
.expect("grad+hess jet2");
let rel = |a: f64, b: f64| (a - b).abs() <= 1e-9 * (1.0 + b.abs());
assert!(
rel(first.eta, full.eta),
"eta {} != {}",
first.eta,
full.eta
);
assert!(
rel(first.chi, full.chi),
"chi {} != {}",
first.chi,
full.chi
);
assert!(rel(first.d, full.d), "d {} != {}", first.d, full.d);
for u in 0..p {
assert!(
rel(first.eta_u[u], full.eta_u[u]),
"eta_u[{u}] {} != {}",
first.eta_u[u],
full.eta_u[u]
);
assert!(
rel(first.chi_u[u], full.chi_u[u]),
"chi_u[{u}] {} != {}",
first.chi_u[u],
full.chi_u[u]
);
assert!(
rel(first.d_u[u], full.d_u[u]),
"d_u[{u}] {} != {}",
first.d_u[u],
full.d_u[u]
);
}
let eval = |gg: f64, qq: f64| -> (f64, f64, f64) {
let (aa, _) = family
.solve_row_survival_intercept_with_slot(
qq,
gg,
None,
None,
Some((row, SurvivalInterceptSlotKind::Exit)),
)
.expect("fd intercept solve");
let tp = family
.compute_survival_timepoint_first_order_exact(
row, &primary, qq, primary.q1, aa, gg, None, None, o_infl,
)
.expect("fd grad-only");
(tp.eta, tp.chi, tp.d)
};
let h = 1e-6_f64;
let (eta_gp, chi_gp, d_gp) = eval(g + h, q1);
let (eta_gm, chi_gm, d_gm) = eval(g - h, q1);
let (eta_qp, chi_qp, d_qp) = eval(g, q1 + h);
let (eta_qm, chi_qm, d_qm) = eval(g, q1 - h);
let fd = |plus: f64, minus: f64| (plus - minus) / (2.0 * h);
let check = |label: &str, analytic: f64, numeric: f64| {
assert!(
(analytic - numeric).abs() <= 1e-5 * (1.0 + analytic.abs()),
"{label}: analytic {analytic} != fd {numeric}"
);
};
check("d eta/dg", first.eta_u[primary.g], fd(eta_gp, eta_gm));
check("d chi/dg", first.chi_u[primary.g], fd(chi_gp, chi_gm));
check("d d/dg", first.d_u[primary.g], fd(d_gp, d_gm));
check("d eta/dq", first.eta_u[primary.q1], fd(eta_qp, eta_qm));
check("d chi/dq", first.chi_u[primary.q1], fd(chi_qp, chi_qm));
check("d d/dq", first.d_u[primary.q1], fd(d_qp, d_qm));
}
#[test]
fn flex_timepoint_inputs_nested_dual_matches_jet4_contraction_932() {
let n = 16usize;
let family = make_ghw_flex_family(n);
let primary = flex_primary_slices(&family);
let p = primary.total;
let row = 6usize;
let g = 0.2_f64;
let h_len = primary.h.as_ref().map(|r| r.len()).unwrap_or(0);
let w_len = primary.w.as_ref().map(|r| r.len()).unwrap_or(0);
let beta_h = Array1::from_iter(
(0..h_len).map(|i| 0.1 + 0.05 * (i as f64) - 0.02 * ((i % 2) as f64)),
);
let beta_w = Array1::from_iter(
(0..w_len).map(|i| -0.08 + 0.04 * (i as f64) + 0.01 * ((i % 3) as f64)),
);
let bh = Some(&beta_h);
let bw = Some(&beta_w);
let m_beta = 0.15_f64;
let q1 = family.offset_exit[row] + family.marginal_design.to_dense()[[row, 0]] * m_beta;
let o_infl = 0.0_f64;
let a1 = family
.solve_row_survival_intercept_with_slot(
q1,
g,
bh,
bw,
Some((row, SurvivalInterceptSlotKind::Exit)),
)
.expect("intercept solve")
.0;
let cached = family
.build_cached_partition(&primary, a1, g, bh, bw)
.expect("cached partition");
let (obs_coeff, obs_fixed) =
observed_fixed_for(&family, &primary, row, a1, g, bh, bw).expect("obs fixed");
let cells = cells_from_cached(&cached);
let z_obs = family.observed_score_projection(row);
let d_check = family
.evaluate_survival_denom_d(a1, g, bh, bw)
.expect("denom");
for trial in 0..4usize {
let f = trial as f64;
let d1 =
Array1::from_iter((0..p).map(|c| {
0.11 + 0.03 * (c as f64) - 0.02 * (((c + trial) % 2) as f64) + 0.01 * f
}));
let d2 = Array1::from_iter((0..p).map(|c| {
-0.06 + 0.045 * (((c + trial) % 3) as f64) + 0.02 * (c as f64) - 0.015 * f
}));
let tpl4 = Jet4::primary(0.0, usize::MAX, p, 0.0, 0.0);
let b4 = Jet4::primary(g, primary.g, p, d2[primary.g], d2[primary.g]);
let du4: Vec<Jet4> = (0..p)
.map(|u| Jet4::primary(0.0, u, p, d2[u], d2[u]))
.collect();
let q4 = add_const(&du4[primary.q1], q1);
let (eta4, chi4, d4) = flex_timepoint_inputs_generic(
&tpl4,
&b4,
&du4,
a1,
d_check,
primary.g,
primary.infl,
&q4,
&const_jet_like(&tpl4, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)
.expect("generic jet4");
let contract = |m: &Vec<f64>| -> f64 {
let mut s = 0.0;
for u in 0..p {
for v in 0..p {
s += d1[u] * d1[v] * m[u * p + v];
}
}
s
};
let eta_ref = contract(&eta4.eps_del.h);
let chi_ref = contract(&chi4.eps_del.h);
let d_ref = contract(&d4.eps_del.h);
let tpl2 = Dual22::seed_directional(0.0, 0.0, 0.0);
let b2 = Dual22::seed_directional(g, d1[primary.g], d2[primary.g]);
let du2: Vec<Dual22> = (0..p)
.map(|u| Dual22::seed_directional(0.0, d1[u], d2[u]))
.collect();
let q2 = add_const(&du2[primary.q1], q1);
let (eta2, chi2, d2n) = flex_timepoint_inputs_generic(
&tpl2,
&b2,
&du2,
a1,
d_check,
primary.g,
primary.infl,
&q2,
&const_jet_like(&tpl2, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)
.expect("generic dual22");
let eta_got = eta2.channels()[8];
let chi_got = chi2.channels()[8];
let d_got = d2n.channels()[8];
let check = |label: &str, got: f64, refv: f64| {
assert!(
(got - refv).abs() <= 1e-9 * (1.0 + refv.abs()),
"trial {trial} {label}: nested-dual {got} != jet4 contraction {refv}"
);
};
check("eta_uv_uv", eta_got, eta_ref);
check("chi_uv_uv", chi_got, chi_ref);
check("d_uv_uv", d_got, d_ref);
}
}
fn flex_test_deviation_runtime() -> DeviationRuntime {
build_score_warp_deviation_block_from_seed(
&Array1::from(vec![-1.0, 0.0, 1.0]),
&DeviationBlockConfig {
degree: 3,
num_internal_knots: 1,
penalty_order: 2,
penalty_orders: vec![1, 2, 3],
double_penalty: false,
monotonicity_eps: 1e-4,
},
)
.expect("build test deviation runtime")
.runtime
}
fn make_ghw_flex_family(n: usize) -> SurvivalMarginalSlopeFamily {
let mut family = make_g_only_flex_family(n);
family.score_warp = Some(flex_test_deviation_runtime());
family.link_dev = Some(flex_test_deviation_runtime());
family
}
fn make_complete_family_map_fixture() -> (SurvivalMarginalSlopeFamily, Vec<ParameterBlockState>)
{
let mut family = make_ghw_flex_family(1);
let knots = Array1::from(vec![0.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0]);
let degree = 3usize;
let wiggle_width = time_wiggle_basis_ncols(&knots, degree).expect("timewiggle basis width");
let time_width = 1 + wiggle_width;
let mut entry_design = Array2::zeros((1, time_width));
let mut exit_design = Array2::zeros((1, time_width));
let mut derivative_design = Array2::zeros((1, time_width));
entry_design[[0, 0]] = 0.25;
exit_design[[0, 0]] = 0.45;
derivative_design[[0, 0]] = 0.15;
family.design_entry = DesignMatrix::from(entry_design);
family.design_exit = DesignMatrix::from(exit_design);
family.design_derivative_exit = DesignMatrix::from(derivative_design);
family.offset_entry = Arc::new(Array1::from(vec![0.20]));
family.offset_exit = Arc::new(Array1::from(vec![0.35]));
family.derivative_offset_exit = Arc::new(Array1::from(vec![0.80]));
family.marginal_design = DesignMatrix::from(
Array2::from_shape_vec((1, 1), vec![0.30]).expect("marginal fixture shape"),
);
family.logslope_layout = DesignMatrix::from(
Array2::from_shape_vec((1, 1), vec![0.40]).expect("logslope fixture shape"),
)
.into();
family.influence_absorber = Some(
Array2::from_shape_vec((1, 2), vec![0.20, -0.15]).expect("influence fixture shape"),
);
family.time_wiggle_knots = Some(knots);
family.time_wiggle_degree = Some(degree);
family.time_wiggle_ncols = wiggle_width;
family.event = Arc::new(Array1::from(vec![1.0]));
family.weights = Arc::new(Array1::from(vec![0.9]));
family.z =
Arc::new(Array2::from_shape_vec((1, 1), vec![0.2]).expect("score fixture shape"));
family.gaussian_frailty_sd = Some(0.6);
let mut beta_time = Array1::zeros(time_width);
beta_time[0] = 0.10;
for local in 0..wiggle_width {
beta_time[1 + local] = 0.006 + 0.002 * local as f64;
}
let beta_marginal = Array1::from(vec![0.12]);
let beta_logslope = Array1::from(vec![0.20]);
let score_width = family
.score_warp
.as_ref()
.expect("score runtime")
.basis_dim();
let link_width = family.link_dev.as_ref().expect("link runtime").basis_dim();
let beta_score =
Array1::from_iter((0..score_width).map(|axis| 0.004 * (1.0 + axis as f64)));
let beta_link = Array1::from_iter((0..link_width).map(|axis| -0.003 + 0.001 * axis as f64));
let beta_influence = Array1::from(vec![0.03, -0.02]);
let time_eta = family.design_exit.dot_row(0, &beta_time) + family.offset_exit[0];
let marginal_eta = family.marginal_design.dot_row(0, &beta_marginal);
let logslope_eta = family
.logslope_layout
.coefficient_design()
.dot_row(0, &beta_logslope);
let states = vec![
ParameterBlockState {
beta: beta_time,
eta: Array1::from(vec![time_eta]),
},
ParameterBlockState {
beta: beta_marginal,
eta: Array1::from(vec![marginal_eta]),
},
ParameterBlockState {
beta: beta_logslope,
eta: Array1::from(vec![logslope_eta]),
},
ParameterBlockState {
beta: beta_score,
eta: Array1::zeros(1),
},
ParameterBlockState {
beta: beta_link,
eta: Array1::zeros(1),
},
ParameterBlockState {
beta: beta_influence,
eta: Array1::zeros(1),
},
];
(family, states)
}
#[test]
fn complete_family_map_owns_every_coefficient_block_and_log_sigma_scale_stack() {
let (family, states) = make_complete_family_map_fixture();
let primary = flex_primary_slices(&family);
let slices = block_slices(&family, &states);
let sigma = family.gaussian_frailty_sd.expect("fixture sigma");
let sigma2 = sigma * sigma;
let alpha = sigma2 / (1.0 + sigma2);
let scale = family.probit_frailty_scale();
let first = FlexFamilyRowDirection {
entry: 0.07,
exit: -0.04,
derivative_exit: 0.03,
probit_scale: -scale * alpha,
};
let second = FlexFamilyRowDirection {
entry: -0.02,
exit: 0.05,
derivative_exit: -0.01,
probit_scale: scale * alpha * (3.0 * alpha - 2.0),
};
let jets = family
.build_flex_family_coefficient_jets::<Jet2>(
0, &states, &primary, &slices, first, second, None,
)
.expect("complete family coefficient map");
let q = family
.row_dynamic_q_geometry(0, &states)
.expect("analytic dynamic q geometry");
let check = |actual: f64, expected: f64, channel: &str| {
let tolerance = 1024.0 * f64::EPSILON * (1.0 + actual.abs().max(expected.abs()));
assert!(
(actual - expected).abs() <= tolerance,
"{channel}: actual={actual:.17e}, expected={expected:.17e}, tolerance={tolerance:.3e}"
);
};
for local in 0..slices.time.len() {
let axis = slices.time.start + local;
check(jets.q0.v.g[axis], q.dq0_time[local], "q0 time");
check(jets.q1.v.g[axis], q.dq1_time[local], "q1 time");
check(jets.qd1.v.g[axis], q.dqd1_time[local], "qd1 time");
}
for local in 0..slices.marginal.len() {
let axis = slices.marginal.start + local;
check(jets.q0.v.g[axis], q.dq0_marginal[local], "q0 marginal");
check(jets.q1.v.g[axis], q.dq1_marginal[local], "q1 marginal");
check(jets.qd1.v.g[axis], q.dqd1_marginal[local], "qd1 marginal");
}
assert_eq!(
jets.g.v.g[slices.logslope.start],
family.logslope_layout.coefficient_design().to_dense()[[0, 0]],
);
let h_primary = primary.h.as_ref().expect("score primary").start;
let h_coefficient = slices.score_warp.as_ref().expect("score slice").start;
assert_eq!(jets.du[h_primary].v.g[h_coefficient], 1.0);
let w_primary = primary.w.as_ref().expect("link primary").start;
let w_coefficient = slices.link_dev.as_ref().expect("link slice").start;
assert_eq!(jets.du[w_primary].v.g[w_coefficient], 1.0);
let influence_primary = primary.infl.expect("influence primary");
let influence_range = slices.influence.as_ref().expect("influence slice");
assert_eq!(jets.du[influence_primary].v.g[influence_range.start], 0.20);
assert_eq!(
jets.du[influence_primary].v.g[influence_range.start + 1],
-0.15
);
assert_eq!(jets.scale_ratio.v.value(), 1.0);
check(jets.scale_ratio.g.value(), -alpha, "ds/s");
check(
jets.scale_ratio.h.value(),
alpha * (3.0 * alpha - 2.0),
"d2s/s",
);
assert!(jets.scale_ratio.g.g.iter().all(|value| *value == 0.0));
assert!(jets.scale_ratio.h.h.iter().all(|value| *value == 0.0));
}
#[test]
fn complete_family_row_dual2_jet2_and_jet3_share_channels_without_fd() {
let (family, states) = make_complete_family_map_fixture();
let slices = block_slices(&family, &states);
let scale = family.probit_frailty_scale();
let first = FlexFamilyRowDirection {
entry: 0.04,
exit: -0.03,
derivative_exit: 0.02,
probit_scale: -0.08 * scale,
};
let second = FlexFamilyRowDirection {
entry: -0.01,
exit: 0.025,
derivative_exit: 0.015,
probit_scale: 0.03 * scale,
};
let order2 = family
.flex_family_direction_row_terms(0, &states, first, second, None)
.expect("Dual2<Jet2> family row");
let direction =
Array1::from_iter((0..slices.total).map(|axis| -0.035 + 0.007 * (axis % 9) as f64));
let order3 = family
.flex_family_direction_row_terms(0, &states, first, second, Some(&direction))
.expect("Dual2<Jet3> family row");
let assert_same = |left: &FlexFamilyCoefficientTerms,
right: &FlexFamilyCoefficientTerms,
channel: &str| {
assert_eq!(
left.objective.to_bits(),
right.objective.to_bits(),
"{channel} V"
);
for axis in 0..slices.total {
assert_eq!(
left.gradient[axis].to_bits(),
right.gradient[axis].to_bits(),
"{channel} g[{axis}]"
);
for other in 0..slices.total {
assert_eq!(
left.hessian[[axis, other]].to_bits(),
right.hessian[[axis, other]].to_bits(),
"{channel} H[{axis},{other}]"
);
}
}
};
assert_same(&order2.first, &order3.first, "first");
assert_same(&order2.second, &order3.second, "second");
let drift = order3.directional.expect("nonzero Jet3 beta drift");
let expected_value = order2.first.gradient.dot(&direction);
let tolerance =
|actual: f64, expected: f64| 1.0e-10 * (1.0 + actual.abs().max(expected.abs()));
assert!(
(drift.objective - expected_value).abs() <= tolerance(drift.objective, expected_value),
"family beta drift V {} != g·d {}",
drift.objective,
expected_value,
);
for axis in 0..slices.total {
let expected_gradient = order2.first.hessian.row(axis).dot(&direction);
assert!(
(drift.gradient[axis] - expected_gradient).abs()
<= tolerance(drift.gradient[axis], expected_gradient),
"family beta drift g[{axis}] {} != H[row]·d {}",
drift.gradient[axis],
expected_gradient,
);
}
assert!(drift.hessian.iter().all(|value| value.is_finite()));
assert!(drift.hessian.iter().any(|value| *value != 0.0));
let x = family.marginal_design.to_dense()[[0, 0]];
let x_psi = 0.17;
let beta = states[1].beta[0];
let design = family
.flex_family_design_direction_row_terms(
0,
&states,
first,
second,
1,
&Array1::from(vec![x_psi]),
)
.expect("Dual2<Jet3> family-by-design row");
assert_same(&order2.first, &design.first, "design first base");
assert_same(&order2.second, &design.second, "design second base");
let design_drift = design.directional.expect("nonzero design drift");
let marginal_axis = slices.marginal.start;
let h_psi = x_psi * beta;
let expected_design_objective = order2.first.gradient[marginal_axis] * h_psi / x;
assert!(
(design_drift.objective - expected_design_objective).abs()
<= tolerance(design_drift.objective, expected_design_objective),
"family-by-design objective {} != analytic {}",
design_drift.objective,
expected_design_objective,
);
let expected_design_score = x_psi * order2.first.gradient[marginal_axis] / x
+ h_psi * order2.first.hessian[[marginal_axis, marginal_axis]] / x;
assert!(
(design_drift.gradient[marginal_axis] - expected_design_score).abs()
<= tolerance(design_drift.gradient[marginal_axis], expected_design_score,),
"family-by-design marginal score {} != analytic pullback {}",
design_drift.gradient[marginal_axis],
expected_design_score,
);
assert!(design_drift.hessian.iter().all(|value| value.is_finite()));
assert!(design_drift.hessian.iter().any(|value| *value != 0.0));
}
#[test]
fn flex_timepoint_inputs_ghw_jet4_matches_scalar_fd_932() {
let n = 16usize;
let family = make_ghw_flex_family(n);
let primary = flex_primary_slices(&family);
let p = primary.total;
let row = 6usize;
let g = 0.2_f64;
let h_len = primary.h.as_ref().map(|r| r.len()).unwrap_or(0);
let w_len = primary.w.as_ref().map(|r| r.len()).unwrap_or(0);
let beta_h = Array1::from_iter(
(0..h_len).map(|i| 0.1 + 0.05 * (i as f64) - 0.02 * ((i % 2) as f64)),
);
let beta_w = Array1::from_iter(
(0..w_len).map(|i| -0.08 + 0.04 * (i as f64) + 0.01 * ((i % 3) as f64)),
);
let bh = Some(&beta_h);
let bw = Some(&beta_w);
let m_beta = 0.15_f64;
let q1 = family.offset_exit[row] + family.marginal_design.to_dense()[[row, 0]] * m_beta;
let o_infl = 0.0_f64;
let a1 = family
.solve_row_survival_intercept_with_slot(
q1,
g,
bh,
bw,
Some((row, SurvivalInterceptSlotKind::Exit)),
)
.expect("intercept solve")
.0;
let cached = family
.build_cached_partition(&primary, a1, g, bh, bw)
.expect("cached partition");
let (obs_coeff, obs_fixed) =
observed_fixed_for(&family, &primary, row, a1, g, bh, bw).expect("obs fixed");
let cells = cells_from_cached(&cached);
let z_obs = family.observed_score_projection(row);
let d_check = family
.evaluate_survival_denom_d(a1, g, bh, bw)
.expect("denom");
let (oracle_eta_uvuv, oracle_chi_uvuv, oracle_d_uvuv) = {
let dir1 = Array1::from_iter(
(0..p).map(|c| 0.12 + 0.04 * (c as f64) - 0.01 * ((c % 2) as f64)),
);
let dir2 = Array1::from_iter(
(0..p).map(|c| -0.07 + 0.05 * ((c % 3) as f64) + 0.02 * (c as f64)),
);
let scalars_of = |pert: &Array1<f64>| -> (f64, f64, f64, f64) {
let q1_pert = q1 + pert[primary.q1];
let g_pert = g + pert[primary.g];
let bh_pert: Array1<f64> = Array1::from_iter(
(0..h_len).map(|i| beta_h[i] + pert[primary.h.as_ref().unwrap().start + i]),
);
let bw_pert: Array1<f64> = Array1::from_iter(
(0..w_len).map(|i| beta_w[i] + pert[primary.w.as_ref().unwrap().start + i]),
);
let a_pert = family
.solve_row_survival_intercept_with_slot(
q1_pert,
g_pert,
Some(&bh_pert),
Some(&bw_pert),
None,
)
.expect("oracle intercept solve")
.0;
let obs = family
.observed_denested_cell_partials(
row,
a_pert,
g_pert,
Some(&bh_pert),
Some(&bw_pert),
)
.expect("oracle observed partials");
let d_pert = family
.evaluate_survival_denom_d(a_pert, g_pert, Some(&bh_pert), Some(&bw_pert))
.expect("oracle denom");
(
a_pert,
eval_coeff4_at(&obs.coeff, z_obs) + o_infl,
eval_coeff4_at(&obs.dc_da, z_obs),
d_pert,
)
};
let hq = 2.0e-3_f64;
let ht = 3.0e-3_f64;
let pert_vec =
|su: f64, u: usize, sv: f64, v: usize, t1: f64, t2: f64| -> Array1<f64> {
let mut pert = &dir1 * t1 + &dir2 * t2;
pert[u] += su;
pert[v] += sv;
pert
};
let mixed = |u: usize, v: usize| -> (f64, f64, f64, f64) {
let acc =
|w: f64, su: f64, sv: f64, t1: f64, t2: f64, out: &mut (f64, f64, f64, f64)| {
let s = scalars_of(&pert_vec(su, u, sv, v, t1, t2));
out.0 += w * s.0;
out.1 += w * s.1;
out.2 += w * s.2;
out.3 += w * s.3;
};
let hess_uv = |t1: f64, t2: f64| -> (f64, f64, f64, f64) {
let mut o = (0.0, 0.0, 0.0, 0.0);
if u == v {
acc(1.0, hq, 0.0, t1, t2, &mut o);
acc(-2.0, 0.0, 0.0, t1, t2, &mut o);
acc(1.0, -hq, 0.0, t1, t2, &mut o);
let inv = 1.0 / (hq * hq);
(o.0 * inv, o.1 * inv, o.2 * inv, o.3 * inv)
} else {
acc(1.0, hq, hq, t1, t2, &mut o);
acc(-1.0, hq, -hq, t1, t2, &mut o);
acc(-1.0, -hq, hq, t1, t2, &mut o);
acc(1.0, -hq, -hq, t1, t2, &mut o);
let inv = 1.0 / (4.0 * hq * hq);
(o.0 * inv, o.1 * inv, o.2 * inv, o.3 * inv)
}
};
let a = hess_uv(ht, ht);
let b = hess_uv(ht, -ht);
let c = hess_uv(-ht, ht);
let d = hess_uv(-ht, -ht);
let inv = 1.0 / (4.0 * ht * ht);
(
(a.0 - b.0 - c.0 + d.0) * inv,
(a.1 - b.1 - c.1 + d.1) * inv,
(a.2 - b.2 - c.2 + d.2) * inv,
(a.3 - b.3 - c.3 + d.3) * inv,
)
};
let template4 = Jet4::primary(0.0, usize::MAX, p, 0.0, 0.0);
let b_jet4 = Jet4::primary(g, primary.g, p, dir1[primary.g], dir2[primary.g]);
let du4: Vec<Jet4> = (0..p)
.map(|u| Jet4::primary(0.0, u, p, dir1[u], dir2[u]))
.collect();
let q_jet4 = add_const(&du4[primary.q1], q1);
let scale_ratio4 = const_jet_like(&template4, 1.0);
let residual_probe = |a: &Jet4| {
calibration_residual_jet(
a,
&b_jet4,
primary.g,
&du4,
&q_jet4,
&scale_ratio4,
&cells,
)
};
let a_jet_probe = lift_intercept_flex(&template4, a1, 1.0 / d_check, 4, residual_probe);
let jet_a_uvuv = a_jet_probe.eps_del.h[primary.q1 * p + primary.q1];
let mut o_eta = Array2::<f64>::zeros((p, p));
let mut o_chi = Array2::<f64>::zeros((p, p));
let mut o_d = Array2::<f64>::zeros((p, p));
let mut ref_a_uvuv = 0.0_f64;
for u in 0..p {
for v in u..p {
let (a_uvuv, eta_uvuv, chi_uvuv, d_uvuv) = mixed(u, v);
o_eta[[u, v]] = eta_uvuv;
o_eta[[v, u]] = eta_uvuv;
o_chi[[u, v]] = chi_uvuv;
o_chi[[v, u]] = chi_uvuv;
o_d[[u, v]] = d_uvuv;
o_d[[v, u]] = d_uvuv;
if u == primary.q1 && v == primary.q1 {
ref_a_uvuv = a_uvuv;
}
}
}
assert!(
(jet_a_uvuv - ref_a_uvuv).abs() <= 1e-3 * (1.0 + ref_a_uvuv.abs()),
"#932 PROBE a_uv_uv[q1,q1]: jet {jet_a_uvuv} != scalar-FD {ref_a_uvuv} \
(diff {})",
jet_a_uvuv - ref_a_uvuv,
);
(o_eta, o_chi, o_d)
};
let dir1 =
Array1::from_iter((0..p).map(|c| 0.12 + 0.04 * (c as f64) - 0.01 * ((c % 2) as f64)));
let dir2 =
Array1::from_iter((0..p).map(|c| -0.07 + 0.05 * ((c % 3) as f64) + 0.02 * (c as f64)));
let template4 = Jet4::primary(0.0, usize::MAX, p, 0.0, 0.0);
let b_jet4 = Jet4::primary(g, primary.g, p, dir1[primary.g], dir2[primary.g]);
let du4: Vec<Jet4> = (0..p)
.map(|u| Jet4::primary(0.0, u, p, dir1[u], dir2[u]))
.collect();
let q_jet4 = add_const(&du4[primary.q1], q1);
let (eta4, chi4, d4) = flex_timepoint_inputs_generic(
&template4,
&b_jet4,
&du4,
a1,
d_check,
primary.g,
primary.infl,
&q_jet4,
&const_jet_like(&template4, 1.0),
z_obs,
o_infl,
obs_coeff,
&obs_fixed,
&cells,
)
.expect("generic jet4");
let cmp_mat_oracle = |label: &str, jet: &Vec<f64>, oracle: &Array2<f64>| {
let mut fails: Vec<String> = Vec::new();
for u in 0..p {
for v in 0..p {
let o = oracle[[u, v]];
let j = jet[u * p + v];
if (j - o).abs() > 1e-3 * (1.0 + o.abs()) {
let rel = (j - o).abs() / (1.0 + o.abs());
fails.push(format!("[{u},{v}] jet {j:.6} oracle {o:.6} rel {rel:.2e}"));
}
}
}
assert!(
fails.is_empty(),
"{label} jet != scalar-FD oracle at {} entr{}: {}",
fails.len(),
if fails.len() == 1 { "y" } else { "ies" },
fails.join("; "),
);
};
cmp_mat_oracle("eta_uv_uv", &eta4.eps_del.h, &oracle_eta_uvuv);
cmp_mat_oracle("chi_uv_uv", &chi4.eps_del.h, &oracle_chi_uvuv);
cmp_mat_oracle("d_uv_uv", &d4.eps_del.h, &oracle_d_uvuv);
}
}
#[cfg(test)]
mod compiled_order2_oracle_tests {
use super::*;
fn xorshift(state: &mut u64) -> f64 {
let mut x = *state;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
*state = x;
let u = (x >> 11) as f64 / ((1u64 << 53) as f64);
2.0 * u - 1.0
}
fn rand_dense(p: usize, st: &mut u64) -> (f64, Vec<f64>, Vec<f64>) {
let v = xorshift(st);
let g: Vec<f64> = (0..p).map(|_| xorshift(st)).collect();
let mut h = vec![0.0; p * p];
for i in 0..p {
for j in i..p {
let x = xorshift(st);
h[i * p + j] = x;
h[j * p + i] = x;
}
}
(v, g, h)
}
#[test]
fn compiled_order2_row_nll_matches_generic_plan() {
for &p in &[3usize, 6, 9, 12, 20] {
let mut st = 0x5DEE_CE66_D9C7_F123u64 ^ (p as u64).wrapping_mul(0x9E3779B97F4A7C15);
for _ in 0..2200 {
let wi = (xorshift(&mut st) + 1.5).abs() + 0.1;
let di = if xorshift(&mut st) > 0.0 { 1.0 } else { 0.0 };
let surv0: [f64; 5] = std::array::from_fn(|_| xorshift(&mut st));
let surv1: [f64; 5] = std::array::from_fn(|_| xorshift(&mut st));
let (e0v, e0g, e0h) = rand_dense(p, &mut st);
let (e1v, e1g, e1h) = rand_dense(p, &mut st);
let (mut cv, cg, ch) = rand_dense(p, &mut st);
cv = (cv + 2.0).abs() + 0.3;
let (mut dv, dg, dh) = rand_dense(p, &mut st);
dv = (dv + 2.0).abs() + 0.3;
let q1v = (xorshift(&mut st) + 2.0).abs() + 0.2;
let qd1v = (xorshift(&mut st) + 2.0).abs() + 0.2;
let qax = p - 2;
let qdax = p - 1;
let g_out = flex_row_nll(
&Jet2::from_parts(e0v, &e0g, &e0h),
&Jet2::from_parts(e1v, &e1g, &e1h),
&Jet2::from_parts(cv, &cg, &ch),
&Jet2::from_parts(dv, &dg, &dh),
&Jet2::primary(q1v, qax, p),
&Jet2::primary(qd1v, qdax, p),
surv0,
surv1,
wi,
di,
);
let plan = FlexOuterPlan::new(cv, dv, qd1v, surv0, surv1, wi, di);
let (f_value, f_gradient, f_hessian) = lower_flex_outer_plan_order2(
&plan,
FlexOrder2Inputs {
eta0: FlexOrder2View {
value: e0v,
gradient: ndarray::ArrayView1::from(&e0g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e0h).unwrap(),
},
eta1: FlexOrder2View {
value: e1v,
gradient: ndarray::ArrayView1::from(&e1g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e1h).unwrap(),
},
q1: (q1v, qax),
chi1: FlexOrder2View {
value: cv,
gradient: ndarray::ArrayView1::from(&cg),
hessian: ndarray::ArrayView2::from_shape((p, p), &ch).unwrap(),
},
d1: FlexOrder2View {
value: dv,
gradient: ndarray::ArrayView1::from(&dg),
hessian: ndarray::ArrayView2::from_shape((p, p), &dh).unwrap(),
},
qd1: (qd1v, qdax),
},
p,
);
let close = |left: f64, right: f64, channel: &str| {
let tolerance = 1e-12 * left.abs().max(right.abs()).max(1.0);
assert!(
(left - right).abs() <= tolerance,
"{channel} p={p}: generic={left:+.16e} compiled={right:+.16e}"
);
};
close(g_out.v, f_value, "value");
for i in 0..p {
close(g_out.g[i], f_gradient[i], &format!("grad[{i}]"));
}
for k in 0..p * p {
close(g_out.h[k], f_hessian[k], &format!("hess[{k}]"));
}
}
}
}
#[test]
fn release_measure_flex_compiled_order2_vs_generic_plan_932() {
use std::time::Instant;
let p = 12usize;
let qax = p - 2;
let qdax = p - 1;
fn best_ns<F: FnMut(f64) -> f64>(iterations: usize, base: f64, mut evaluate: F) -> f64 {
let mut best = f64::INFINITY;
for _ in 0..5 {
let mut checksum = 0.0_f64;
let started = Instant::now();
for _ in 0..iterations {
checksum += evaluate(base + checksum * 1e-18);
}
assert!(
checksum.is_finite(),
"flex order-2 release-measure checksum must stay finite"
);
best = best.min(started.elapsed().as_secs_f64());
}
best * 1e9 / iterations as f64
}
let iterations = 20_000usize;
for di in [0.0_f64, 1.0] {
let mut st = 0xA5F0_3C11_9D2E_7B41u64 ^ di.to_bits();
let wi = (xorshift(&mut st) + 1.5).abs() + 0.1;
let surv0: [f64; 5] = std::array::from_fn(|_| xorshift(&mut st));
let surv1: [f64; 5] = std::array::from_fn(|_| xorshift(&mut st));
let (e0v, e0g, e0h) = rand_dense(p, &mut st);
let (e1v, e1g, e1h) = rand_dense(p, &mut st);
let (mut cv, cg, ch) = rand_dense(p, &mut st);
cv = (cv + 2.0).abs() + 0.3;
let (mut dv, dg, dh) = rand_dense(p, &mut st);
dv = (dv + 2.0).abs() + 0.3;
let q1v = (xorshift(&mut st) + 2.0).abs() + 0.2;
let qd1v = (xorshift(&mut st) + 2.0).abs() + 0.2;
let plan = FlexOuterPlan::new(cv, dv, qd1v, surv0, surv1, wi, di);
let generic = flex_row_nll(
&Jet2::from_parts(e0v, &e0g, &e0h),
&Jet2::from_parts(e1v, &e1g, &e1h),
&Jet2::from_parts(cv, &cg, &ch),
&Jet2::from_parts(dv, &dg, &dh),
&Jet2::primary(q1v, qax, p),
&Jet2::primary(qd1v, qdax, p),
surv0,
surv1,
wi,
di,
);
let (compiled_value, ..) = lower_flex_outer_plan_order2(
&plan,
FlexOrder2Inputs {
eta0: FlexOrder2View {
value: e0v,
gradient: ndarray::ArrayView1::from(&e0g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e0h).unwrap(),
},
eta1: FlexOrder2View {
value: e1v,
gradient: ndarray::ArrayView1::from(&e1g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e1h).unwrap(),
},
q1: (q1v, qax),
chi1: FlexOrder2View {
value: cv,
gradient: ndarray::ArrayView1::from(&cg),
hessian: ndarray::ArrayView2::from_shape((p, p), &ch).unwrap(),
},
d1: FlexOrder2View {
value: dv,
gradient: ndarray::ArrayView1::from(&dg),
hessian: ndarray::ArrayView2::from_shape((p, p), &dh).unwrap(),
},
qd1: (qd1v, qdax),
},
p,
);
let tolerance = 1e-12 * generic.v.abs().max(compiled_value.abs()).max(1.0);
assert!(
(generic.v - compiled_value).abs() <= tolerance,
"di={di:.0} value: generic={:+.16e} compiled={compiled_value:+.16e}",
generic.v,
);
let production_ns = best_ns(iterations, e0v, |perturbed_e0v| {
let (value, gradient, hessian) = lower_flex_outer_plan_order2(
&plan,
FlexOrder2Inputs {
eta0: FlexOrder2View {
value: perturbed_e0v,
gradient: ndarray::ArrayView1::from(&e0g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e0h).unwrap(),
},
eta1: FlexOrder2View {
value: e1v,
gradient: ndarray::ArrayView1::from(&e1g),
hessian: ndarray::ArrayView2::from_shape((p, p), &e1h).unwrap(),
},
q1: (q1v, qax),
chi1: FlexOrder2View {
value: cv,
gradient: ndarray::ArrayView1::from(&cg),
hessian: ndarray::ArrayView2::from_shape((p, p), &ch).unwrap(),
},
d1: FlexOrder2View {
value: dv,
gradient: ndarray::ArrayView1::from(&dg),
hessian: ndarray::ArrayView2::from_shape((p, p), &dh).unwrap(),
},
qd1: (qd1v, qdax),
},
p,
);
value + gradient[0] + hessian[0]
});
let generic_ns = best_ns(iterations, e0v, |perturbed_e0v| {
let out = flex_row_nll(
&Jet2::from_parts(perturbed_e0v, &e0g, &e0h),
&Jet2::from_parts(e1v, &e1g, &e1h),
&Jet2::from_parts(cv, &cg, &ch),
&Jet2::from_parts(dv, &dg, &dh),
&Jet2::primary(q1v, qax, p),
&Jet2::primary(qd1v, qdax, p),
surv0,
surv1,
wi,
di,
);
out.v + out.g[0] + out.h[0]
});
eprintln!(
"FLEX-ORDER2-932 p={p} event={di:.0} production={production_ns:.2} ns/row \
generic_plan={generic_ns:.2} ns/row hand_over_production={:.6}",
generic_ns / production_ns,
);
}
}
}