gam_models/transformation_normal/
alo_replay.rs1use super::{TRANSFORMATION_MONOTONICITY_EPS, log_normal_cdf_diff_derivatives};
2use ndarray::{Array1, Array2};
3
4pub 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#[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
70pub 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}