use super::*;
use gam_math::jet_scalar::JetScalar;
pub(crate) const RIGID_LINEAR_MASK: u32 = (1 << 0) | (1 << 1) | (1 << 2);
#[inline(always)]
const fn axis_is_linear(mask: u32, a: usize) -> bool {
(mask >> a) & 1 == 1
}
#[inline(always)]
const fn h_block_is_zero(mask: u32, i: usize, j: usize) -> bool {
axis_is_linear(mask, i) && axis_is_linear(mask, j)
}
#[inline(always)]
const fn t3_block_is_zero(mask: u32, i: usize, j: usize, k: usize) -> bool {
(axis_is_linear(mask, i) as u32
+ axis_is_linear(mask, j) as u32
+ axis_is_linear(mask, k) as u32)
>= 2
}
#[inline(always)]
const fn t4_block_is_zero(mask: u32, i: usize, j: usize, k: usize, l: usize) -> bool {
(axis_is_linear(mask, i) as u32
+ axis_is_linear(mask, j) as u32
+ axis_is_linear(mask, k) as u32
+ axis_is_linear(mask, l) as u32)
>= 2
}
#[derive(Clone, Copy)]
pub(crate) struct SparseTower3<const LIN: u32> {
pub(crate) v: f64,
pub(crate) g: [f64; 4],
pub(crate) h: [[f64; 4]; 4],
pub(crate) t3: [[[f64; 4]; 4]; 4],
}
impl<const LIN: u32> SparseTower3<LIN> {
#[inline(always)]
fn check_contract(&self) {
for i in 0..4 {
for j in 0..4 {
if h_block_is_zero(LIN, i, j) {
assert!(
self.h[i][j] == 0.0,
"static-sparsity contract violated: h[{i}][{j}]={} != 0",
self.h[i][j]
);
}
for k in 0..4 {
if t3_block_is_zero(LIN, i, j, k) {
assert!(
self.t3[i][j][k] == 0.0,
"static-sparsity contract violated: t3[{i}][{j}][{k}]={} != 0",
self.t3[i][j][k]
);
}
}
}
}
}
}
impl<const LIN: u32> JetScalar<4> for SparseTower3<LIN> {
fn constant(c: f64) -> Self {
Self {
v: c,
g: [0.0; 4],
h: [[0.0; 4]; 4],
t3: [[[0.0; 4]; 4]; 4],
}
}
fn variable(x: f64, axis: usize) -> Self {
let mut out = Self::constant(x);
out.g[axis] = 1.0;
out
}
}
impl<const LIN: u32> gam_math::nested_dual::JetField for SparseTower3<LIN> {
fn value(&self) -> f64 {
self.v
}
fn add(&self, o: &Self) -> Self {
let mut r = *self;
r.v += o.v;
for i in 0..4 {
r.g[i] += o.g[i];
for j in 0..4 {
r.h[i][j] += o.h[i][j];
for k in 0..4 {
r.t3[i][j][k] += o.t3[i][j][k];
}
}
}
r
}
fn sub(&self, o: &Self) -> Self {
self.add(&o.neg())
}
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn scale(&self, s: f64) -> Self {
let mut o = *self;
o.v *= s;
for i in 0..4 {
o.g[i] *= s;
for j in 0..4 {
o.h[i][j] *= s;
for k in 0..4 {
o.t3[i][j][k] *= s;
}
}
}
o
}
fn mul(&self, o: &Self) -> Self {
let (a, b) = (self, o);
a.check_contract();
b.check_contract();
let mut out = Self::constant(a.v * b.v);
for i in 0..4 {
let mut s = 0.0;
s += a.v * b.g[i];
s += a.g[i] * b.v;
out.g[i] = s;
}
for i in 0..4 {
for j in 0..4 {
let mut s = 0.0;
if !h_block_is_zero(LIN, i, j) {
s += a.v * b.h[i][j];
}
s += a.g[i] * b.g[j];
s += a.g[j] * b.g[i];
if !h_block_is_zero(LIN, i, j) {
s += a.h[i][j] * b.v;
}
out.h[i][j] = s;
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
let mut s = 0.0;
if !t3_block_is_zero(LIN, i, j, k) {
s += a.v * b.t3[i][j][k];
}
if !h_block_is_zero(LIN, j, k) {
s += a.g[i] * b.h[j][k];
}
if !h_block_is_zero(LIN, i, k) {
s += a.g[j] * b.h[i][k];
}
if !h_block_is_zero(LIN, i, j) {
s += a.h[i][j] * b.g[k];
}
if !h_block_is_zero(LIN, i, j) {
s += a.g[k] * b.h[i][j];
}
if !h_block_is_zero(LIN, i, k) {
s += a.h[i][k] * b.g[j];
}
if !h_block_is_zero(LIN, j, k) {
s += a.h[j][k] * b.g[i];
}
if !t3_block_is_zero(LIN, i, j, k) {
s += a.t3[i][j][k] * b.v;
}
out.t3[i][j][k] = s;
}
}
}
out
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
self.check_contract();
let mut out = Self::constant(d[0]);
for i in 0..4 {
let mut s = 0.0;
s += d[1] * self.g[i];
out.g[i] = s;
}
for i in 0..4 {
for j in 0..4 {
let mut s = 0.0;
if !h_block_is_zero(LIN, i, j) {
s += d[1] * self.h[i][j];
}
s += d[2] * self.g[i] * self.g[j];
out.h[i][j] = s;
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
let mut s = 0.0;
if !t3_block_is_zero(LIN, i, j, k) {
s += d[1] * self.t3[i][j][k];
}
if !h_block_is_zero(LIN, i, j) {
s += d[2] * self.h[i][j] * self.g[k];
}
if !h_block_is_zero(LIN, i, k) {
s += d[2] * self.h[i][k] * self.g[j];
}
if !h_block_is_zero(LIN, j, k) {
s += d[2] * self.g[i] * self.h[j][k];
}
s += d[3] * self.g[i] * self.g[j] * self.g[k];
out.t3[i][j][k] = s;
}
}
}
out
}
}
#[derive(Clone, Copy)]
pub(crate) struct SparseTower4<const LIN: u32> {
pub(crate) v: f64,
pub(crate) g: [f64; 4],
pub(crate) h: [[f64; 4]; 4],
pub(crate) t3: [[[f64; 4]; 4]; 4],
pub(crate) t4: [[[[f64; 4]; 4]; 4]; 4],
}
impl<const LIN: u32> SparseTower4<LIN> {
#[inline(always)]
fn check_contract(&self) {
for i in 0..4 {
for j in 0..4 {
if h_block_is_zero(LIN, i, j) {
assert!(
self.h[i][j] == 0.0,
"static-sparsity contract violated: h[{i}][{j}]={} != 0",
self.h[i][j]
);
}
for k in 0..4 {
if t3_block_is_zero(LIN, i, j, k) {
assert!(
self.t3[i][j][k] == 0.0,
"static-sparsity contract violated: t3[{i}][{j}][{k}]={} != 0",
self.t3[i][j][k]
);
}
for l in 0..4 {
if t4_block_is_zero(LIN, i, j, k, l) {
assert!(
self.t4[i][j][k][l] == 0.0,
"static-sparsity contract violated: t4[{i}][{j}][{k}][{l}]={} != 0",
self.t4[i][j][k][l]
);
}
}
}
}
}
}
#[inline]
pub(crate) fn fourth_contracted(&self, u: &[f64; 4], w: &[f64; 4]) -> [[f64; 4]; 4] {
let mut out = [[0.0; 4]; 4];
for i in 0..4 {
for j in 0..4 {
let mut acc = 0.0;
for k in 0..4 {
for l in 0..4 {
acc += self.t4[i][j][k][l] * u[k] * w[l];
}
}
out[i][j] = acc;
}
}
out
}
}
impl<const LIN: u32> JetScalar<4> for SparseTower4<LIN> {
fn constant(c: f64) -> Self {
Self {
v: c,
g: [0.0; 4],
h: [[0.0; 4]; 4],
t3: [[[0.0; 4]; 4]; 4],
t4: [[[[0.0; 4]; 4]; 4]; 4],
}
}
fn variable(x: f64, axis: usize) -> Self {
let mut out = Self::constant(x);
out.g[axis] = 1.0;
out
}
}
impl<const LIN: u32> gam_math::nested_dual::JetField for SparseTower4<LIN> {
fn value(&self) -> f64 {
self.v
}
fn add(&self, o: &Self) -> Self {
let mut r = *self;
r.v += o.v;
for i in 0..4 {
r.g[i] += o.g[i];
for j in 0..4 {
r.h[i][j] += o.h[i][j];
for k in 0..4 {
r.t3[i][j][k] += o.t3[i][j][k];
for l in 0..4 {
r.t4[i][j][k][l] += o.t4[i][j][k][l];
}
}
}
}
r
}
fn sub(&self, o: &Self) -> Self {
self.add(&o.neg())
}
fn neg(&self) -> Self {
self.scale(-1.0)
}
fn scale(&self, s: f64) -> Self {
let mut o = *self;
o.v *= s;
for i in 0..4 {
o.g[i] *= s;
for j in 0..4 {
o.h[i][j] *= s;
for k in 0..4 {
o.t3[i][j][k] *= s;
for l in 0..4 {
o.t4[i][j][k][l] *= s;
}
}
}
}
o
}
fn mul(&self, o: &Self) -> Self {
let (a, b) = (self, o);
a.check_contract();
b.check_contract();
let mut out = Self::constant(a.v * b.v);
for i in 0..4 {
let mut s = 0.0;
s += a.v * b.g[i];
s += a.g[i] * b.v;
out.g[i] = s;
}
for i in 0..4 {
for j in 0..4 {
let mut s = 0.0;
if !h_block_is_zero(LIN, i, j) {
s += a.v * b.h[i][j];
}
s += a.g[i] * b.g[j];
s += a.g[j] * b.g[i];
if !h_block_is_zero(LIN, i, j) {
s += a.h[i][j] * b.v;
}
out.h[i][j] = s;
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
let mut s = 0.0;
if !t3_block_is_zero(LIN, i, j, k) {
s += a.v * b.t3[i][j][k];
}
if !h_block_is_zero(LIN, j, k) {
s += a.g[i] * b.h[j][k];
}
if !h_block_is_zero(LIN, i, k) {
s += a.g[j] * b.h[i][k];
}
if !h_block_is_zero(LIN, i, j) {
s += a.h[i][j] * b.g[k];
}
if !h_block_is_zero(LIN, i, j) {
s += a.g[k] * b.h[i][j];
}
if !h_block_is_zero(LIN, i, k) {
s += a.h[i][k] * b.g[j];
}
if !h_block_is_zero(LIN, j, k) {
s += a.h[j][k] * b.g[i];
}
if !t3_block_is_zero(LIN, i, j, k) {
s += a.t3[i][j][k] * b.v;
}
out.t3[i][j][k] = s;
}
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
for l in 0..4 {
let mut s = 0.0;
if !t4_block_is_zero(LIN, i, j, k, l) {
s += a.v * b.t4[i][j][k][l];
}
if !t3_block_is_zero(LIN, j, k, l) {
s += a.g[i] * b.t3[j][k][l];
}
if !t3_block_is_zero(LIN, i, k, l) {
s += a.g[j] * b.t3[i][k][l];
}
if !(h_block_is_zero(LIN, i, j) || h_block_is_zero(LIN, k, l)) {
s += a.h[i][j] * b.h[k][l];
}
if !t3_block_is_zero(LIN, i, j, l) {
s += a.g[k] * b.t3[i][j][l];
}
if !(h_block_is_zero(LIN, i, k) || h_block_is_zero(LIN, j, l)) {
s += a.h[i][k] * b.h[j][l];
}
if !(h_block_is_zero(LIN, j, k) || h_block_is_zero(LIN, i, l)) {
s += a.h[j][k] * b.h[i][l];
}
if !t3_block_is_zero(LIN, i, j, k) {
s += a.t3[i][j][k] * b.g[l];
}
if !t3_block_is_zero(LIN, i, j, k) {
s += a.g[l] * b.t3[i][j][k];
}
if !(h_block_is_zero(LIN, i, l) || h_block_is_zero(LIN, j, k)) {
s += a.h[i][l] * b.h[j][k];
}
if !(h_block_is_zero(LIN, j, l) || h_block_is_zero(LIN, i, k)) {
s += a.h[j][l] * b.h[i][k];
}
if !t3_block_is_zero(LIN, i, j, l) {
s += a.t3[i][j][l] * b.g[k];
}
if !(h_block_is_zero(LIN, k, l) || h_block_is_zero(LIN, i, j)) {
s += a.h[k][l] * b.h[i][j];
}
if !t3_block_is_zero(LIN, i, k, l) {
s += a.t3[i][k][l] * b.g[j];
}
if !t3_block_is_zero(LIN, j, k, l) {
s += a.t3[j][k][l] * b.g[i];
}
if !t4_block_is_zero(LIN, i, j, k, l) {
s += a.t4[i][j][k][l] * b.v;
}
out.t4[i][j][k][l] = s;
}
}
}
}
out
}
fn compose_unary(&self, d: [f64; 5]) -> Self {
self.check_contract();
let mut out = Self::constant(d[0]);
for i in 0..4 {
let mut s = 0.0;
s += d[1] * self.g[i];
out.g[i] = s;
}
for i in 0..4 {
for j in 0..4 {
let mut s = 0.0;
if !h_block_is_zero(LIN, i, j) {
s += d[1] * self.h[i][j];
}
s += d[2] * self.g[i] * self.g[j];
out.h[i][j] = s;
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
let mut s = 0.0;
if !t3_block_is_zero(LIN, i, j, k) {
s += d[1] * self.t3[i][j][k];
}
if !h_block_is_zero(LIN, i, j) {
s += d[2] * self.h[i][j] * self.g[k];
}
if !h_block_is_zero(LIN, i, k) {
s += d[2] * self.h[i][k] * self.g[j];
}
if !h_block_is_zero(LIN, j, k) {
s += d[2] * self.g[i] * self.h[j][k];
}
s += d[3] * self.g[i] * self.g[j] * self.g[k];
out.t3[i][j][k] = s;
}
}
}
for i in 0..4 {
for j in 0..4 {
for k in 0..4 {
for l in 0..4 {
let mut s = 0.0;
if !t4_block_is_zero(LIN, i, j, k, l) {
s += d[1] * self.t4[i][j][k][l];
}
if !t3_block_is_zero(LIN, i, j, k) {
s += d[2] * self.t3[i][j][k] * self.g[l];
}
if !t3_block_is_zero(LIN, i, j, l) {
s += d[2] * self.t3[i][j][l] * self.g[k];
}
if !(h_block_is_zero(LIN, i, j) || h_block_is_zero(LIN, k, l)) {
s += d[2] * self.h[i][j] * self.h[k][l];
}
if !h_block_is_zero(LIN, i, j) {
s += d[3] * self.h[i][j] * self.g[k] * self.g[l];
}
if !t3_block_is_zero(LIN, i, k, l) {
s += d[2] * self.t3[i][k][l] * self.g[j];
}
if !(h_block_is_zero(LIN, i, k) || h_block_is_zero(LIN, j, l)) {
s += d[2] * self.h[i][k] * self.h[j][l];
}
if !h_block_is_zero(LIN, i, k) {
s += d[3] * self.h[i][k] * self.g[j] * self.g[l];
}
if !(h_block_is_zero(LIN, i, l) || h_block_is_zero(LIN, j, k)) {
s += d[2] * self.h[i][l] * self.h[j][k];
}
if !t3_block_is_zero(LIN, j, k, l) {
s += d[2] * self.g[i] * self.t3[j][k][l];
}
if !h_block_is_zero(LIN, j, k) {
s += d[3] * self.g[i] * self.h[j][k] * self.g[l];
}
if !h_block_is_zero(LIN, i, l) {
s += d[3] * self.h[i][l] * self.g[j] * self.g[k];
}
if !h_block_is_zero(LIN, j, l) {
s += d[3] * self.g[i] * self.h[j][l] * self.g[k];
}
if !h_block_is_zero(LIN, k, l) {
s += d[3] * self.g[i] * self.g[j] * self.h[k][l];
}
s += d[4] * self.g[i] * self.g[j] * self.g[k] * self.g[l];
out.t4[i][j][k][l] = s;
}
}
}
}
out
}
}
#[inline]
pub(crate) fn tower3_third_contracted(t3: &[[[f64; 4]; 4]; 4], dir: &[f64; 4]) -> [[f64; 4]; 4] {
let mut out = [[0.0; 4]; 4];
for a in 0..4 {
for b in 0..4 {
let mut acc = 0.0;
for c in 0..4 {
acc += t3[a][b][c] * dir[c];
}
out[a][b] = acc;
}
}
out
}
pub(crate) struct SurvivalMarginalSlopeRowKernel {
pub(crate) family: SurvivalMarginalSlopeFamily,
pub(crate) block_states: Vec<ParameterBlockState>,
pub(crate) slices: BlockSlices,
}
impl SurvivalMarginalSlopeRowKernel {
pub(crate) fn new(
family: SurvivalMarginalSlopeFamily,
block_states: Vec<ParameterBlockState>,
) -> Self {
let slices = block_slices(&family, &block_states);
Self {
family,
block_states,
slices,
}
}
}
#[cfg(all(test, target_os = "linux"))]
mod rigid_row_admission_tests {
use super::*;
fn inputs(wi: f64, di: f64) -> RigidRowInputs {
RigidRowInputs {
row: 7,
wi,
di,
z_sum: 0.0,
covariance_ones: 1.0,
probit_scale: 1.0,
qd1_lower: 0.0,
}
}
fn admit(primaries: [f64; 4], inputs: &RigidRowInputs) -> Result<(), String> {
let [neg_eta0, neg_eta1, adjusted_derivative] =
rigid_row_admission_witnesses(&primaries, inputs);
validate_rigid_row_admission(
primaries[2],
inputs,
neg_eta0,
neg_eta1,
adjusted_derivative,
)
}
#[test]
fn scalar_gpu_admission_witnesses_match_cpu_signed_margin_domain() {
for primaries in [[f64::NAN, 0.0, 1.0, 0.0], [0.0, f64::INFINITY, 1.0, 0.0]] {
let error = admit(primaries, &inputs(1.0, 0.0))
.expect_err("active non-finite signed margin must be rejected");
assert!(error.contains("non-finite signed margin"));
}
admit(
[f64::NEG_INFINITY, f64::NEG_INFINITY, 1.0, 0.0],
&inputs(1.0, 0.0),
)
.expect("positive-infinity signed margins are the admitted saturated tail");
admit([f64::NAN, f64::NAN, 1.0, 0.0], &inputs(0.0, 0.0))
.expect("zero-weight margins do not contribute to the row");
}
}
pub(crate) fn rigid_row_kernel_primaries(
family: &SurvivalMarginalSlopeFamily,
block_states: &[ParameterBlockState],
row: usize,
) -> Result<[f64; 4], String> {
let q_geom = family.row_dynamic_q_values(row, block_states)?;
Ok([q_geom.q0, q_geom.q1, q_geom.qd1, block_states[2].eta[row]])
}
pub(crate) struct RigidRowInputs {
pub(crate) row: usize,
pub(crate) wi: f64,
pub(crate) di: f64,
pub(crate) z_sum: f64,
pub(crate) covariance_ones: f64,
pub(crate) probit_scale: f64,
pub(crate) qd1_lower: f64,
}
pub(crate) fn rigid_row_inputs(
family: &SurvivalMarginalSlopeFamily,
block_states: &[ParameterBlockState],
row: usize,
context: &str,
) -> Result<RigidRowInputs, String> {
let (z_sum, covariance_ones) = family.exact_shared_score_summary(row, block_states, context)?;
Ok(RigidRowInputs {
row,
wi: family.weights[row],
di: family.event[row],
z_sum,
covariance_ones,
probit_scale: family.probit_frailty_scale(),
qd1_lower: family.time_derivative_lower_bound(),
})
}
#[inline(always)]
fn rigid_row_feature_jets<S: JetScalar<4>>(
vars: &[S; 4],
inputs: &RigidRowInputs,
) -> [S; RIGID_FEATURE_DIMENSION] {
let observed_g = vars[3].scale(inputs.probit_scale);
let linear = observed_g.scale(inputs.z_sum);
let variance = vars[3].mul(&vars[3]).scale(inputs.covariance_ones);
[vars[0], vars[1], vars[2], linear, variance]
}
#[inline(always)]
fn rigid_row_feature_values(
primaries: &[f64; 4],
inputs: &RigidRowInputs,
) -> [f64; RIGID_FEATURE_DIMENSION] {
let [q0, q1, qd1, g] = *primaries;
let observed_g = inputs.probit_scale * g;
[
q0,
q1,
qd1,
observed_g * inputs.z_sum,
(g * g) * inputs.covariance_ones,
]
}
#[inline(always)]
#[cfg(target_os = "linux")]
pub(crate) fn rigid_row_admission_witnesses(
primaries: &[f64; 4],
inputs: &RigidRowInputs,
) -> [f64; 3] {
let [q0, q1, qd1, linear, variance] = rigid_row_feature_values(primaries, inputs);
rigid_feature_program_witnesses(q0, q1, qd1, linear, variance, inputs.probit_scale)
}
pub(crate) fn rigid_row_nll<S: JetScalar<4>>(
vars: &[S; 4],
inputs: &RigidRowInputs,
) -> Result<S, String> {
let features = rigid_row_feature_jets(vars, inputs);
let (nll, [neg_eta0, neg_eta1, adjusted_derivative]) = rigid_feature_program::<4, S>(
&features[FEATURE_Q0],
&features[FEATURE_Q1],
&features[FEATURE_QD1],
&features[FEATURE_LINEAR],
&features[FEATURE_VARIANCE],
inputs.wi,
inputs.di,
inputs.probit_scale,
);
validate_rigid_row_admission(
vars[2].value(),
inputs,
neg_eta0,
neg_eta1,
adjusted_derivative,
)?;
Ok(nll)
}
#[inline(always)]
pub(crate) fn rigid_row_order2(
primaries: &[f64; 4],
inputs: &RigidRowInputs,
) -> Result<(f64, [f64; 4], [[f64; 4]; 4]), String> {
let [q0, q1, qd1, linear, variance] = rigid_row_feature_values(primaries, inputs);
let (value, feature_gradient, feature_hessian, [neg_eta0, neg_eta1, adjusted_derivative]) =
rigid_feature_program_order2(
q0,
q1,
qd1,
linear,
variance,
inputs.wi,
inputs.di,
inputs.probit_scale,
);
validate_rigid_row_admission(
primaries[FEATURE_QD1],
inputs,
neg_eta0,
neg_eta1,
adjusted_derivative,
)?;
const DIMENSION: usize = 4;
let mut jacobian = [0.0; RIGID_FEATURE_DIMENSION * DIMENSION];
jacobian[FEATURE_Q0 * DIMENSION + FEATURE_Q0] = 1.0;
jacobian[FEATURE_Q1 * DIMENSION + FEATURE_Q1] = 1.0;
jacobian[FEATURE_QD1 * DIMENSION + FEATURE_QD1] = 1.0;
jacobian[FEATURE_LINEAR * DIMENSION + 3] = inputs.probit_scale * inputs.z_sum;
jacobian[FEATURE_VARIANCE * DIMENSION + 3] = 2.0 * primaries[3] * inputs.covariance_ones;
let mut gradient = [0.0; DIMENSION];
let mut flat_hessian = [0.0; DIMENSION * DIMENSION];
order2_feature_pullback_into(
&feature_gradient,
&feature_hessian,
&jacobian,
|axis| if axis < 3 { 1 } else { 2 },
|axis, slot| {
if axis < 3 {
axis
} else {
FEATURE_LINEAR + slot
}
},
DIMENSION,
&mut gradient,
&mut flat_hessian,
|gradient, hessian| {
hessian[3 * DIMENSION + 3] += gradient[FEATURE_VARIANCE] * 2.0 * inputs.covariance_ones;
},
);
let hessian = [
[
flat_hessian[0],
flat_hessian[1],
flat_hessian[2],
flat_hessian[3],
],
[
flat_hessian[4],
flat_hessian[5],
flat_hessian[6],
flat_hessian[7],
],
[
flat_hessian[8],
flat_hessian[9],
flat_hessian[10],
flat_hessian[11],
],
[
flat_hessian[12],
flat_hessian[13],
flat_hessian[14],
flat_hessian[15],
],
];
Ok((value, gradient, hessian))
}
pub(crate) fn validate_rigid_row_admission(
qd1: f64,
inputs: &RigidRowInputs,
neg_eta0: f64,
neg_eta1: f64,
adjusted_derivative: f64,
) -> Result<(), String> {
let RigidRowInputs {
row,
wi,
di,
qd1_lower,
..
} = *inputs;
if survival_derivative_guard_violated(qd1, qd1_lower) {
return Err(SurvivalMarginalSlopeError::MonotonicityViolation {
reason: format!(
"survival marginal-slope monotonicity violated at row {row}: raw time derivative={:.3e} must be at least derivative_guard={:.3e}; transformed time derivative={:.3e}",
qd1, qd1_lower, adjusted_derivative
),
}
.into());
}
let reject_nonfinite_margin = |margin: f64, weight: f64| -> Result<(), String> {
if weight != 0.0 && margin != f64::INFINITY && !margin.is_finite() {
Err(SurvivalMarginalSlopeError::NumericalFailure {
reason: format!(
"non-finite signed margin in rigid survival marginal-slope row tower at row {row}: {margin}"
),
}
.into())
} else {
Ok(())
}
};
reject_nonfinite_margin(neg_eta0, wi)?;
reject_nonfinite_margin(neg_eta1, wi * (1.0 - di))?;
Ok(())
}
impl gam_math::jet_tower::RowProgram<4> for SurvivalMarginalSlopeRowKernel {
fn n_rows(&self) -> usize {
self.family.n
}
fn primaries(&self, row: usize) -> Result<[f64; 4], String> {
rigid_row_kernel_primaries(&self.family, &self.block_states, row)
}
fn eval<S: JetScalar<4>>(&self, row: usize, p: &[S; 4]) -> Result<S, String> {
let inputs = rigid_row_inputs(
&self.family,
&self.block_states,
row,
"survival marginal-slope rigid row program",
)?;
rigid_row_nll(p, &inputs)
}
}
impl RowKernel<4> for SurvivalMarginalSlopeRowKernel {
fn n_coefficients(&self) -> usize {
self.slices.total
}
fn row_kernel(&self, row: usize) -> Result<(f64, [f64; 4], [[f64; 4]; 4]), String> {
let inputs = rigid_row_inputs(
&self.family,
&self.block_states,
row,
"survival marginal-slope rigid row kernel",
)?;
let p = rigid_row_kernel_primaries(&self.family, &self.block_states, row)?;
rigid_row_order2(&p, &inputs)
}
fn batched_value_grad_hess_all(
&self,
) -> Option<Result<(Vec<f64>, Vec<[f64; 4]>, Vec<[[f64; 4]; 4]>), String>> {
use crate::gpu_kernels::survival_rowjet::survival_rigid_row_vgh_device_selected;
let n = self.family.n;
match survival_rigid_row_vgh_device_selected(n) {
Ok(true) => {}
Ok(false) => return None,
Err(error) => return Some(Err(error)),
}
#[cfg(target_os = "linux")]
{
use crate::gpu_kernels::survival_rowjet::{SurvivalRowInputs, survival_rigid_row_vgh};
let probit_scale = self.family.probit_frailty_scale();
let gather: Result<Vec<SurvivalRowInputs>, String> = (0..n)
.into_par_iter()
.map(|row| {
let p = rigid_row_kernel_primaries(&self.family, &self.block_states, row)?;
let inputs = rigid_row_inputs(
&self.family,
&self.block_states,
row,
"survival marginal-slope rigid row kernel (batched)",
)?;
let [neg_eta0, neg_eta1, adjusted_derivative] =
rigid_row_admission_witnesses(&p, &inputs);
validate_rigid_row_admission(
p[2],
&inputs,
neg_eta0,
neg_eta1,
adjusted_derivative,
)?;
Ok(SurvivalRowInputs {
primaries: p,
wi: inputs.wi,
di: inputs.di,
z_sum: inputs.z_sum,
cov_ones: inputs.covariance_ones,
})
})
.collect();
let rows = match gather {
Ok(rows) => rows,
Err(error) => return Some(Err(error)),
};
let ch = match survival_rigid_row_vgh(&rows, probit_scale) {
Ok(channels) => channels,
Err(error) => return Some(Err(error)),
};
let mut grads = vec![[0.0_f64; 4]; n];
let mut hesss = vec![[[0.0_f64; 4]; 4]; n];
for row in 0..n {
for a in 0..4 {
grads[row][a] = ch.grad[row * 4 + a];
for b in 0..4 {
hesss[row][a][b] = ch.hess[row * 16 + a * 4 + b];
}
}
}
Some(Ok((ch.value, grads, hesss)))
}
#[cfg(not(target_os = "linux"))]
None
}
fn jacobian_action(&self, row: usize, d_beta: &[f64]) -> [f64; 4] {
let d_beta = ndarray::ArrayView1::from(d_beta);
let d_time = d_beta.slice(s![self.slices.time.clone()]);
let d_marginal = d_beta.slice(s![self.slices.marginal.clone()]);
let d_logslope = d_beta.slice(s![self.slices.logslope.clone()]);
[
self.family.design_entry.dot_row_view(row, d_time)
+ self.family.marginal_design.dot_row_view(row, d_marginal),
self.family.design_exit.dot_row_view(row, d_time)
+ self.family.marginal_design.dot_row_view(row, d_marginal),
self.family.design_derivative_exit.dot_row_view(row, d_time),
self.family
.logslope_layout
.coefficient_design()
.dot_row_view(row, d_logslope),
]
}
fn jacobian_action_matrix(&self, factor: ArrayView2<'_, f64>) -> Option<Array2<f64>> {
if factor.nrows() != self.slices.total {
return None;
}
let n_rows = self.family.n;
Some(self.assemble_jf(factor, n_rows, |design, factor_block| {
crate::row_kernel::row_kernel_design_jf(design, factor_block, n_rows)
}))
}
fn jacobian_action_matrix_rows(
&self,
factor: ArrayView2<'_, f64>,
start: usize,
end: usize,
) -> Array2<f64> {
assert_eq!(
factor.nrows(),
self.slices.total,
"survival marginal-slope tiled Jacobian factor width must match coefficients",
);
let b = end.saturating_sub(start);
self.assemble_jf(factor, b, |design, factor_block| {
crate::row_kernel::row_kernel_design_jf_rows(design, factor_block, start, end)
})
}
fn jacobian_transpose_action(&self, row: usize, v: &[f64; 4], out: &mut [f64]) {
{
let mut time = ndarray::ArrayViewMut1::from(&mut out[self.slices.time.clone()]);
self.family
.design_entry
.axpy_row_into(row, v[0], &mut time)
.expect("time entry axpy dim mismatch");
self.family
.design_exit
.axpy_row_into(row, v[1], &mut time)
.expect("time exit axpy dim mismatch");
self.family
.design_derivative_exit
.axpy_row_into(row, v[2], &mut time)
.expect("time deriv axpy dim mismatch");
}
{
let mut marginal = ndarray::ArrayViewMut1::from(&mut out[self.slices.marginal.clone()]);
self.family
.marginal_design
.axpy_row_into(row, v[0] + v[1], &mut marginal)
.expect("marginal axpy dim mismatch");
}
{
let mut logslope = ndarray::ArrayViewMut1::from(&mut out[self.slices.logslope.clone()]);
self.family
.logslope_layout
.coefficient_design()
.axpy_row_into(row, v[3], &mut logslope)
.expect("logslope axpy dim mismatch");
}
}
fn add_pullback_hessian(&self, row: usize, h: &[[f64; 4]; 4], target: &mut Array2<f64>) {
let mut h_arr = Array2::<f64>::zeros((4, 4));
for a in 0..4 {
for b in 0..4 {
h_arr[[a, b]] = h[a][b];
}
}
self.family
.add_pullback_primary_hessian(target, row, &self.slices, &h_arr);
}
fn add_diagonal_quadratic(&self, row: usize, h: &[[f64; 4]; 4], diag: &mut [f64]) {
let designs: [(usize, &DesignMatrix); 3] = [
(0, &self.family.design_entry),
(1, &self.family.design_exit),
(2, &self.family.design_derivative_exit),
];
for &(pi, des) in &designs {
{
let mut td = ndarray::ArrayViewMut1::from(&mut diag[self.slices.time.clone()]);
des.squared_axpy_row_into(row, h[pi][pi], &mut td)
.expect("time squared_axpy dim mismatch");
}
for &(pj, des_j) in &designs {
if pj <= pi {
continue;
}
let mut td = ndarray::ArrayViewMut1::from(&mut diag[self.slices.time.clone()]);
des.crossdiag_axpy_row_into(row, des_j, 2.0 * h[pi][pj], &mut td)
.expect("time crossdiag dim mismatch");
}
}
{
let alpha = h[0][0] + 2.0 * h[0][1] + h[1][1];
let mut md = ndarray::ArrayViewMut1::from(&mut diag[self.slices.marginal.clone()]);
self.family
.marginal_design
.squared_axpy_row_into(row, alpha, &mut md)
.expect("marginal squared_axpy dim mismatch");
}
{
let mut gd = ndarray::ArrayViewMut1::from(&mut diag[self.slices.logslope.clone()]);
self.family
.logslope_layout
.coefficient_design()
.squared_axpy_row_into(row, h[3][3], &mut gd)
.expect("logslope squared_axpy dim mismatch");
}
}
fn directional_derivative_all_axes_dense_override(
&self,
rows: &crate::row_kernel::RowSet,
p: usize,
) -> Option<Result<Vec<Array2<f64>>, String>> {
if p != self.n_coefficients() {
return Some(Err(format!(
"survival marginal-slope directional_derivative_all_axes_dense_override: \
axis count {p} disagrees with n_coefficients() {}",
self.n_coefficients(),
)));
}
if !matches!(rows, crate::row_kernel::RowSet::All) {
return None;
}
Some(self.directional_derivative_all_axes_build_once(p))
}
fn second_directional_derivative_all_axes_dense_override(
&self,
rows: &crate::row_kernel::RowSet,
d_beta_u: &[f64],
) -> Option<Result<Vec<Array2<f64>>, String>> {
if d_beta_u.len() != self.n_coefficients() {
return Some(Err(format!(
"survival marginal-slope second_directional_derivative_all_axes_dense_override: \
fixed direction has {} entries, expected {}",
d_beta_u.len(),
self.n_coefficients(),
)));
}
if !matches!(rows, crate::row_kernel::RowSet::All) {
return None;
}
Some(self.second_directional_derivative_all_axes_build_once(d_beta_u))
}
}
impl SurvivalMarginalSlopeRowKernel {
pub(crate) fn assemble_jf<F>(
&self,
factor: ArrayView2<'_, f64>,
n_out: usize,
axis: F,
) -> Array2<f64>
where
F: Fn(&DesignMatrix, ArrayView2<'_, f64>) -> Array2<f64>,
{
let rank = factor.ncols();
if rank == 0 {
return Array2::<f64>::zeros((n_out, 0));
}
let f_time = factor.slice(s![self.slices.time.clone(), ..]);
let f_marginal = factor.slice(s![self.slices.marginal.clone(), ..]);
let f_logslope = factor.slice(s![self.slices.logslope.clone(), ..]);
let jf_marginal = axis(&self.family.marginal_design, f_marginal);
let mut axis0 = axis(&self.family.design_entry, f_time);
axis0 += &jf_marginal;
let mut axis1 = axis(&self.family.design_exit, f_time);
axis1 += &jf_marginal;
let axis2 = axis(&self.family.design_derivative_exit, f_time);
let axis3 = axis(self.family.logslope_layout.coefficient_design(), f_logslope);
crate::row_kernel::row_kernel_pack_jf_axes::<4>(
n_out,
rank,
[(0, axis0), (1, axis1), (2, axis2), (3, axis3)],
)
}
}
impl SurvivalMarginalSlopeRowKernel {
fn build_row_towers(&self) -> Result<Vec<SparseTower4<RIGID_LINEAR_MASK>>, String> {
let n = gam_math::jet_tower::RowProgram::n_rows(self);
(0..n)
.into_par_iter()
.map(|row| {
let inputs = rigid_row_inputs(
&self.family,
&self.block_states,
row,
"survival marginal-slope rigid row fourth tower (build-once)",
)?;
let p = rigid_row_kernel_primaries(&self.family, &self.block_states, row)?;
let vars: [SparseTower4<RIGID_LINEAR_MASK>; 4] =
std::array::from_fn(|a| SparseTower4::variable(p[a], a));
rigid_row_nll(&vars, &inputs)
})
.collect()
}
fn build_row_third_towers(&self) -> Result<Vec<SparseTower3<RIGID_LINEAR_MASK>>, String> {
let n = gam_math::jet_tower::RowProgram::n_rows(self);
(0..n)
.into_par_iter()
.map(|row| {
let inputs = rigid_row_inputs(
&self.family,
&self.block_states,
row,
"survival marginal-slope rigid row third tower (build-once)",
)?;
let p = rigid_row_kernel_primaries(&self.family, &self.block_states, row)?;
let vars: [SparseTower3<RIGID_LINEAR_MASK>; 4] =
std::array::from_fn(|a| SparseTower3::variable(p[a], a));
rigid_row_nll(&vars, &inputs)
})
.collect()
}
fn chunked_pullback_reduce<F>(&self, p: usize, per_row: F) -> Result<Array2<f64>, String>
where
F: Fn(usize, &mut Array2<f64>) -> Result<(), String> + Sync,
{
let n = gam_math::jet_tower::RowProgram::n_rows(self);
let chunk = crate::outer_subsample::ARROW_ROW_CHUNK;
let n_chunks = crate::outer_subsample::arrow_row_chunk_count(n);
let chunk_accumulators: Vec<Result<Array2<f64>, String>> = (0..n_chunks)
.into_par_iter()
.map(|chunk_idx| {
let start = chunk_idx * chunk;
let end = (start + chunk).min(n);
let mut acc = Array2::<f64>::zeros((p, p));
for row in start..end {
per_row(row, &mut acc)?;
}
Ok(acc)
})
.collect();
let mut total = Array2::<f64>::zeros((p, p));
for acc in chunk_accumulators {
total += &acc?;
}
Ok(total)
}
fn directional_derivative_all_axes_build_once(
&self,
p: usize,
) -> Result<Vec<Array2<f64>>, String> {
let towers = self.build_row_third_towers()?;
(0..p)
.into_par_iter()
.map(|a| {
let mut axis = vec![0.0_f64; p];
axis[a] = 1.0;
gam_problem::with_nested_parallel(|| {
self.chunked_pullback_reduce(p, |row, acc| {
let dir = self.jacobian_action(row, &axis);
let third = tower3_third_contracted(&towers[row].t3, &dir);
self.add_pullback_hessian(row, &third, acc);
Ok(())
})
})
})
.collect()
}
fn second_directional_derivative_all_axes_build_once(
&self,
d_beta_u: &[f64],
) -> Result<Vec<Array2<f64>>, String> {
let p = self.n_coefficients();
let towers = self.build_row_towers()?;
(0..p)
.into_par_iter()
.map(|a| {
let mut axis = vec![0.0_f64; p];
axis[a] = 1.0;
gam_problem::with_nested_parallel(|| {
self.chunked_pullback_reduce(p, |row, acc| {
let dir_u = self.jacobian_action(row, d_beta_u);
let dir_v = self.jacobian_action(row, &axis);
let fourth = towers[row].fourth_contracted(&dir_u, &dir_v);
self.add_pullback_hessian(row, &fourth, acc);
Ok(())
})
})
})
.collect()
}
pub(crate) fn contracted_trace_hessian(
&self,
weight: &Array2<f64>,
) -> Result<Array2<f64>, String> {
let p = self.n_coefficients();
if weight.dim() != (p, p) {
return Err(format!(
"SurvivalMarginalSlopeRowKernel::contracted_trace_hessian: weight shape {:?} != ({p}, {p})",
weight.dim()
));
}
let towers = self.build_row_towers()?;
self.chunked_pullback_reduce(p, |row, acc| -> Result<(), String> {
let w_row = self.primary_trace_weight(row, weight)?;
let t4 = &towers[row].t4;
let mut coeff = [[0.0_f64; 4]; 4];
for c in 0..4 {
for d in 0..4 {
let mut s = 0.0;
for a in 0..4 {
for b in 0..4 {
s += w_row[a][b] * t4[a][b][c][d];
}
}
coeff[c][d] = s;
}
}
self.add_pullback_hessian(row, &coeff, acc);
Ok(())
})
}
fn primary_trace_weight(
&self,
row: usize,
weight: &Array2<f64>,
) -> Result<[[f64; 4]; 4], String> {
let xt_e = self
.family
.design_entry
.try_row_chunk(row..row + 1)
.map_err(|e| format!("primary_trace_weight: design_entry row chunk failed: {e}"))?;
let xt_x = self
.family
.design_exit
.try_row_chunk(row..row + 1)
.map_err(|e| format!("primary_trace_weight: design_exit row chunk failed: {e}"))?;
let xt_d = self
.family
.design_derivative_exit
.try_row_chunk(row..row + 1)
.map_err(|e| {
format!("primary_trace_weight: design_derivative_exit row chunk failed: {e}")
})?;
let xm = self
.family
.marginal_design
.try_row_chunk(row..row + 1)
.map_err(|e| format!("primary_trace_weight: marginal_design row chunk failed: {e}"))?;
let xg = self
.family
.logslope_layout
.coefficient_design()
.try_row_chunk(row..row + 1)
.map_err(|e| format!("primary_trace_weight: logslope_design row chunk failed: {e}"))?;
struct Component<'a> {
vec: ArrayView1<'a, f64>,
range: std::ops::Range<usize>,
}
let components: [Vec<Component<'_>>; 4] = [
vec![
Component {
vec: xt_e.row(0),
range: self.slices.time.clone(),
},
Component {
vec: xm.row(0),
range: self.slices.marginal.clone(),
},
],
vec![
Component {
vec: xt_x.row(0),
range: self.slices.time.clone(),
},
Component {
vec: xm.row(0),
range: self.slices.marginal.clone(),
},
],
vec![Component {
vec: xt_d.row(0),
range: self.slices.time.clone(),
}],
vec![Component {
vec: xg.row(0),
range: self.slices.logslope.clone(),
}],
];
let mut w_row = [[0.0_f64; 4]; 4];
for a in 0..4 {
for b in 0..4 {
let mut acc = 0.0;
for ca in &components[a] {
for cb in &components[b] {
let wblk = weight.slice(s![ca.range.clone(), cb.range.clone()]);
acc += ca.vec.dot(&wblk.dot(&cb.vec));
}
}
w_row[a][b] = acc;
}
}
Ok(w_row)
}
}