1use std::fmt;
12
13use crate::stats::population_variance;
14
15#[derive(Debug, Clone, PartialEq)]
19pub enum TransformError {
20 NonPositiveData,
22 InsufficientData,
24 InvalidInverse,
26}
27
28impl fmt::Display for TransformError {
29 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
30 match self {
31 TransformError::NonPositiveData => {
32 write!(f, "Box-Cox requires all y > 0")
33 }
34 TransformError::InsufficientData => {
35 write!(f, "need at least 2 data points")
36 }
37 TransformError::InvalidInverse => {
38 write!(f, "inverse transformation produced non-finite values")
39 }
40 }
41 }
42}
43
44impl std::error::Error for TransformError {}
45
46fn validate_positive_slice(y: &[f64]) -> Result<(), TransformError> {
49 if y.len() < 2 {
50 return Err(TransformError::InsufficientData);
51 }
52 if y.iter().any(|&v| v <= 0.0) {
53 return Err(TransformError::NonPositiveData);
54 }
55 Ok(())
56}
57
58pub fn box_cox(y: &[f64], lambda: f64) -> Result<Vec<f64>, TransformError> {
83 validate_positive_slice(y)?;
84 let result = if lambda.abs() < 1e-10 {
85 y.iter().map(|&v| v.ln()).collect()
86 } else {
87 y.iter().map(|&v| (v.powf(lambda) - 1.0) / lambda).collect()
88 };
89 Ok(result)
90}
91
92pub fn inverse_box_cox(y_t: &[f64], lambda: f64) -> Result<Vec<f64>, TransformError> {
119 let result: Vec<f64> = if lambda.abs() < 1e-10 {
120 y_t.iter().map(|&v| v.exp()).collect()
121 } else {
122 y_t.iter()
123 .map(|&v| (v * lambda + 1.0).powf(1.0 / lambda))
124 .collect()
125 };
126 if result.iter().any(|v| !v.is_finite()) {
127 return Err(TransformError::InvalidInverse);
128 }
129 Ok(result)
130}
131
132pub fn estimate_lambda(y: &[f64], lambda_min: f64, lambda_max: f64) -> Result<f64, TransformError> {
157 if lambda_min >= lambda_max {
158 return Err(TransformError::InsufficientData);
159 }
160 validate_positive_slice(y)?;
161
162 let n = y.len() as f64;
163 let log_sum: f64 = y.iter().map(|&v| v.ln()).sum::<f64>();
164
165 let profile_ll = |lambda: f64| -> f64 {
167 let y_t: Vec<f64> = if lambda.abs() < 1e-10 {
168 y.iter().map(|&v| v.ln()).collect()
169 } else {
170 y.iter().map(|&v| (v.powf(lambda) - 1.0) / lambda).collect()
171 };
172 let var = population_variance(&y_t).expect("slice has >= 2 elements — variance is defined");
173 if var <= 0.0 {
174 return f64::NEG_INFINITY;
175 }
176 -(n / 2.0) * var.ln() + (lambda - 1.0) * log_sum
177 };
178
179 const PHI: f64 = 0.618_033_988_749_895; let mut a = lambda_min;
182 let mut b = lambda_max;
183
184 let mut x1 = b - PHI * (b - a);
185 let mut x2 = a + PHI * (b - a);
186 let mut f1 = profile_ll(x1);
187 let mut f2 = profile_ll(x2);
188
189 for _ in 0..100 {
190 if (b - a).abs() < 1e-6 {
191 break;
192 }
193 if f1 < f2 {
194 a = x1;
195 x1 = x2;
196 f1 = f2;
197 x2 = a + PHI * (b - a);
198 f2 = profile_ll(x2);
199 } else {
200 b = x2;
201 x2 = x1;
202 f2 = f1;
203 x1 = b - PHI * (b - a);
204 f1 = profile_ll(x1);
205 }
206 }
207
208 Ok((a + b) / 2.0)
209}
210
211#[cfg(test)]
214mod tests {
215 use super::*;
216
217 #[test]
218 fn box_cox_log_transform() {
219 let y = vec![1.0, std::f64::consts::E, std::f64::consts::E.powi(2)];
221 let y_t = box_cox(&y, 0.0).unwrap();
222 assert!((y_t[0] - 0.0).abs() < 1e-10);
223 assert!((y_t[1] - 1.0).abs() < 1e-9);
224 assert!((y_t[2] - 2.0).abs() < 1e-9);
225 }
226
227 #[test]
228 fn box_cox_identity_lambda_1() {
229 let y = vec![2.0, 5.0, 10.0];
231 let y_t = box_cox(&y, 1.0).unwrap();
232 assert!((y_t[0] - 1.0).abs() < 1e-10);
233 assert!((y_t[1] - 4.0).abs() < 1e-10);
234 }
235
236 #[test]
237 fn box_cox_sqrt_lambda_half() {
238 let y = vec![4.0, 9.0];
240 let y_t = box_cox(&y, 0.5).unwrap();
241 assert!((y_t[0] - 2.0).abs() < 1e-10); assert!((y_t[1] - 4.0).abs() < 1e-10); }
244
245 #[test]
246 fn inverse_roundtrip_multiple_lambdas() {
247 let y = vec![1.5, 2.3, 4.7, 8.1, 15.2];
248 for &lambda in &[-2.0_f64, -1.0, -0.5, 0.0, 0.5, 1.0, 2.0] {
249 let y_t = box_cox(&y, lambda).unwrap();
250 let y_rec = inverse_box_cox(&y_t, lambda).unwrap();
251 for (orig, rec) in y.iter().zip(y_rec.iter()) {
252 assert!(
253 (orig - rec).abs() < 1e-9,
254 "lambda={lambda} orig={orig} rec={rec}"
255 );
256 }
257 }
258 }
259
260 #[test]
261 fn estimate_lambda_near_zero_for_exponential() {
262 let y: Vec<f64> = (1..=30).map(|i| (i as f64 * 0.2).exp()).collect();
264 let lambda = estimate_lambda(&y, -2.0, 2.0).unwrap();
265 assert!(lambda.abs() < 0.3, "Expected lambda near 0, got {lambda}");
266 }
267
268 #[test]
269 fn estimate_lambda_near_half_for_quadratic() {
270 let y: Vec<f64> = (1..=20).map(|i| (i as f64).powi(2)).collect();
272 let lambda = estimate_lambda(&y, -2.0, 2.0).unwrap();
273 assert!(
274 lambda > 0.2 && lambda < 0.8,
275 "Expected lambda ~0.5, got {lambda}"
276 );
277 }
278
279 #[test]
280 fn non_positive_returns_error() {
281 assert!(box_cox(&[1.0, -1.0, 2.0], 0.5).is_err());
282 assert!(box_cox(&[0.0, 1.0, 2.0], 0.5).is_err());
283 }
284
285 #[test]
286 fn insufficient_data_returns_error() {
287 assert!(box_cox(&[1.0], 0.5).is_err());
288 assert!(estimate_lambda(&[1.0], -2.0, 2.0).is_err());
289 }
290
291 #[test]
292 fn inverse_invalid_returns_error() {
293 let y_t = vec![-1.0, -0.8];
296 assert!(inverse_box_cox(&y_t, 2.0).is_err());
297 }
298
299 #[test]
300 fn estimate_lambda_invalid_range() {
301 let y = vec![1.0, 2.0, 3.0, 4.0];
302 assert!(estimate_lambda(&y, 1.0, 0.0).is_err()); assert!(estimate_lambda(&y, 0.5, 0.5).is_err()); }
305
306 #[test]
307 fn box_cox_negative_lambda() {
308 let y = vec![2.0, 4.0];
310 let y_t = box_cox(&y, -1.0).unwrap();
311 assert!((y_t[0] - 0.5).abs() < 1e-10);
313 assert!((y_t[1] - 0.75).abs() < 1e-10);
315 }
316}