Skip to main content

gam_models/transformation_normal/
alo_replay.rs

1use super::{TRANSFORMATION_MONOTONICITY_EPS, log_normal_cdf_diff_derivatives};
2use ndarray::{Array1, Array2};
3
4/// Complete local state for one saved transformation-normal likelihood row.
5pub struct TransformationNormalAloRowInput<'a> {
6    pub response_value_basis: &'a [f64],
7    pub response_derivative_basis: &'a [f64],
8    pub response_lower_basis: &'a [f64],
9    pub response_upper_basis: &'a [f64],
10    pub alpha: &'a [f64],
11    pub additive_offset: f64,
12    pub response_floor_offset: f64,
13    pub response_lower_floor_offset: f64,
14    pub response_upper_floor_offset: f64,
15    pub prior_weight: f64,
16}
17
18/// Exact negative-log-likelihood derivatives in the affine local coordinates
19/// `alpha_k(x) = covariate_row(x) beta_k`.
20#[derive(Clone, Debug, PartialEq)]
21pub struct TransformationNormalAloRowGeometry {
22    pub negative_log_likelihood: f64,
23    pub nll_score: Array1<f64>,
24    pub observed_hessian: Array2<f64>,
25}
26
27fn validate_row(input: &TransformationNormalAloRowInput<'_>) -> Result<usize, String> {
28    let dimension = input.alpha.len();
29    if dimension == 0
30        || input.response_value_basis.len() != dimension
31        || input.response_derivative_basis.len() != dimension
32        || input.response_lower_basis.len() != dimension
33        || input.response_upper_basis.len() != dimension
34    {
35        return Err(format!(
36            "transformation-normal ALO row dimension mismatch: alpha={dimension}, value={}, derivative={}, lower={}, upper={}",
37            input.response_value_basis.len(),
38            input.response_derivative_basis.len(),
39            input.response_lower_basis.len(),
40            input.response_upper_basis.len(),
41        ));
42    }
43    if !input.prior_weight.is_finite() || input.prior_weight < 0.0 {
44        return Err(format!(
45            "transformation-normal ALO prior weight must be finite and non-negative, got {}",
46            input.prior_weight
47        ));
48    }
49    if input
50        .response_value_basis
51        .iter()
52        .chain(input.response_derivative_basis)
53        .chain(input.response_lower_basis)
54        .chain(input.response_upper_basis)
55        .chain(input.alpha)
56        .copied()
57        .chain([
58            input.additive_offset,
59            input.response_floor_offset,
60            input.response_lower_floor_offset,
61            input.response_upper_floor_offset,
62        ])
63        .any(|value| !value.is_finite())
64    {
65        return Err("transformation-normal ALO row state must be finite".to_string());
66    }
67    Ok(dimension)
68}
69
70/// Replay one row of the fitted finite-support SCOP likelihood.
71///
72/// This is the row factorization of the same score and negative Hessian used by
73/// `TransformationNormalFamily`: every component is affine in direct-alpha
74/// coordinates, the monotonicity derivative floor is exact, and both
75/// transformed support endpoints contribute through the normalized Gaussian
76/// mass. Feasibility of the shape coordinates is owned by the fitted model's
77/// Khatri-Rao cone before this row replay is called.
78pub fn transformation_normal_alo_row_geometry(
79    input: TransformationNormalAloRowInput<'_>,
80) -> Result<TransformationNormalAloRowGeometry, String> {
81    let dimension = validate_row(&input)?;
82    if input.prior_weight == 0.0 {
83        return Ok(TransformationNormalAloRowGeometry {
84            negative_log_likelihood: 0.0,
85            nll_score: Array1::zeros(dimension),
86            observed_hessian: Array2::zeros((dimension, dimension)),
87        });
88    }
89
90    let alpha0 = input.alpha[0];
91    let mut h = input.response_value_basis[0] * alpha0
92        + input.additive_offset
93        + input.response_floor_offset;
94    let mut h_prime = input.response_derivative_basis[0] * alpha0 + TRANSFORMATION_MONOTONICITY_EPS;
95    let mut lower = input.response_lower_basis[0] * alpha0
96        + input.additive_offset
97        + input.response_lower_floor_offset;
98    let mut upper = input.response_upper_basis[0] * alpha0
99        + input.additive_offset
100        + input.response_upper_floor_offset;
101    for component in 1..dimension {
102        let alpha = input.alpha[component];
103        h += input.response_value_basis[component] * alpha;
104        h_prime += input.response_derivative_basis[component] * alpha;
105        lower += input.response_lower_basis[component] * alpha;
106        upper += input.response_upper_basis[component] * alpha;
107    }
108    if !(h.is_finite() && h_prime.is_finite() && lower.is_finite() && upper.is_finite()) {
109        return Err(format!(
110            "transformation-normal ALO row transform is non-finite: h={h}, h_prime={h_prime}, lower={lower}, upper={upper}"
111        ));
112    }
113    if h_prime <= 0.0 {
114        return Err(format!(
115            "transformation-normal ALO row derivative must be positive, got {h_prime}"
116        ));
117    }
118    let endpoint = log_normal_cdf_diff_derivatives(upper, lower)?;
119    let weight = input.prior_weight;
120    let negative_log_likelihood = weight
121        * (0.5 * h * h + 0.5 * (2.0 * std::f64::consts::PI).ln() - h_prime.ln() + endpoint.log_z);
122
123    let mut dh = vec![0.0; dimension];
124    let mut dh_prime = vec![0.0; dimension];
125    let mut dlower = vec![0.0; dimension];
126    let mut dupper = vec![0.0; dimension];
127    for component in 0..dimension {
128        dh[component] = input.response_value_basis[component];
129        dh_prime[component] = input.response_derivative_basis[component];
130        dlower[component] = input.response_lower_basis[component];
131        dupper[component] = input.response_upper_basis[component];
132    }
133
134    let inverse_h_prime = 1.0 / h_prime;
135    let inverse_h_prime_squared = inverse_h_prime * inverse_h_prime;
136    let mut nll_score = Array1::<f64>::zeros(dimension);
137    let mut observed_hessian = Array2::<f64>::zeros((dimension, dimension));
138    for left in 0..dimension {
139        let endpoint_first = endpoint.first[0] * dupper[left] + endpoint.first[1] * dlower[left];
140        nll_score[left] =
141            weight * (h * dh[left] - dh_prime[left] * inverse_h_prime + endpoint_first);
142        for right in 0..dimension {
143            let endpoint_second = endpoint.second[0][0] * dupper[left] * dupper[right]
144                + endpoint.second[0][1] * dupper[left] * dlower[right]
145                + endpoint.second[1][0] * dlower[left] * dupper[right]
146                + endpoint.second[1][1] * dlower[left] * dlower[right];
147            observed_hessian[[left, right]] = weight
148                * (dh[left] * dh[right]
149                    + dh_prime[left] * dh_prime[right] * inverse_h_prime_squared
150                    + endpoint_second);
151        }
152    }
153    if !negative_log_likelihood.is_finite()
154        || nll_score.iter().any(|value| !value.is_finite())
155        || observed_hessian.iter().any(|value| !value.is_finite())
156    {
157        return Err("transformation-normal ALO row geometry is non-finite".to_string());
158    }
159    Ok(TransformationNormalAloRowGeometry {
160        negative_log_likelihood,
161        nll_score,
162        observed_hessian,
163    })
164}
165
166#[cfg(test)]
167mod tests {
168    use super::*;
169    use crate::transformation_normal::log_normal_cdf_diff;
170
171    fn scalar_nll(alpha: [f64; 2]) -> f64 {
172        let value = [1.0, 0.4];
173        let derivative = [0.0, 0.7];
174        let lower_basis = [1.0, 0.1];
175        let upper_basis = [1.0, 0.9];
176        let offset = -0.15;
177        let floor = 0.02;
178        let lower_floor = -0.04;
179        let upper_floor = 0.06;
180        let weight = 1.3;
181        let h = value[0] * alpha[0] + value[1] * alpha[1] + offset + floor;
182        let h_prime = TRANSFORMATION_MONOTONICITY_EPS
183            + derivative[0] * alpha[0]
184            + derivative[1] * alpha[1];
185        let lower =
186            lower_basis[0] * alpha[0] + lower_basis[1] * alpha[1] + offset + lower_floor;
187        let upper =
188            upper_basis[0] * alpha[0] + upper_basis[1] * alpha[1] + offset + upper_floor;
189        weight
190            * (0.5 * h * h + 0.5 * (2.0 * std::f64::consts::PI).ln() - h_prime.ln()
191                + log_normal_cdf_diff(upper, lower).expect("finite endpoint mass"))
192    }
193
194    #[test]
195    fn saved_transformation_row_geometry_matches_independent_scalar_finite_difference() {
196        let alpha: [f64; 2] = [0.25, 0.8];
197        let geometry = transformation_normal_alo_row_geometry(TransformationNormalAloRowInput {
198            response_value_basis: &[1.0, 0.4],
199            response_derivative_basis: &[0.0, 0.7],
200            response_lower_basis: &[1.0, 0.1],
201            response_upper_basis: &[1.0, 0.9],
202            alpha: &alpha,
203            additive_offset: -0.15,
204            response_floor_offset: 0.02,
205            response_lower_floor_offset: -0.04,
206            response_upper_floor_offset: 0.06,
207            prior_weight: 1.3,
208        })
209        .expect("saved transformation-normal row must replay");
210        let step = 2.0e-5;
211        let base = scalar_nll(alpha);
212        assert!((geometry.negative_log_likelihood - base).abs() <= 2.0e-13);
213        for axis in 0..2 {
214            let mut plus = alpha;
215            let mut minus = alpha;
216            plus[axis] += step;
217            minus[axis] -= step;
218            let gradient_fd = (scalar_nll(plus) - scalar_nll(minus)) / (2.0 * step);
219            assert!(
220                (geometry.nll_score[axis] - gradient_fd).abs() <= 2.0e-8,
221                "score[{axis}] analytic={} fd={gradient_fd}",
222                geometry.nll_score[axis]
223            );
224            for other in 0..2 {
225                let mut pp = alpha;
226                let mut pm = alpha;
227                let mut mp = alpha;
228                let mut mm = alpha;
229                pp[axis] += step;
230                pp[other] += step;
231                pm[axis] += step;
232                pm[other] -= step;
233                mp[axis] -= step;
234                mp[other] += step;
235                mm[axis] -= step;
236                mm[other] -= step;
237                let hessian_fd = (scalar_nll(pp) - scalar_nll(pm) - scalar_nll(mp)
238                    + scalar_nll(mm))
239                    / (4.0 * step * step);
240                assert!(
241                    (geometry.observed_hessian[[axis, other]] - hessian_fd).abs() <= 3.0e-6,
242                    "hessian[{axis},{other}] analytic={} fd={hessian_fd}",
243                    geometry.observed_hessian[[axis, other]]
244                );
245            }
246        }
247        assert!(
248            (geometry.observed_hessian[[0, 0]] - geometry.nll_score[0].powi(2)).abs() > 1.0e-3,
249            "observed curvature must remain distinct from score covariance"
250        );
251    }
252}