1use super::TestResult;
10use crate::error::FdarError;
11use crate::function_on_scalar::integrated_f_statistic;
12use crate::helpers::simpsons_weights;
13use crate::matrix::FdMatrix;
14use rand::rngs::StdRng;
15use rand::SeedableRng;
16
17pub const DEFAULT_N_PERM: usize = 999;
19
20fn validate_two_samples(
25 data_a: &FdMatrix,
26 data_b: &FdMatrix,
27 argvals: &[f64],
28) -> Result<(usize, usize, usize), FdarError> {
29 let (n_a, m_a) = data_a.shape();
30 let (n_b, m_b) = data_b.shape();
31 if m_a == 0 || m_b == 0 {
32 return Err(FdarError::InvalidDimension {
33 parameter: "data",
34 expected: "at least 1 column (grid points)".to_string(),
35 actual: format!("data_a has {m_a} columns, data_b has {m_b} columns"),
36 });
37 }
38 if m_a != m_b {
39 return Err(FdarError::InvalidDimension {
40 parameter: "data_b",
41 expected: format!("{m_a} columns (matching data_a)"),
42 actual: format!("{m_b} columns"),
43 });
44 }
45 if argvals.len() != m_a {
46 return Err(FdarError::InvalidDimension {
47 parameter: "argvals",
48 expected: format!("{m_a} elements (matching data columns)"),
49 actual: format!("{} elements", argvals.len()),
50 });
51 }
52 if n_a < 2 || n_b < 2 {
53 return Err(FdarError::InvalidDimension {
54 parameter: "data",
55 expected: "at least 2 rows per sample".to_string(),
56 actual: format!("data_a has {n_a} rows, data_b has {n_b} rows"),
57 });
58 }
59 Ok((n_a, n_b, m_a))
60}
61
62fn pool_two_samples(
65 data_a: &FdMatrix,
66 data_b: &FdMatrix,
67 n_a: usize,
68 n_b: usize,
69 m: usize,
70) -> FdMatrix {
71 let mut pooled = FdMatrix::zeros(n_a + n_b, m);
72 for j in 0..m {
73 for i in 0..n_a {
74 pooled[(i, j)] = data_a[(i, j)];
75 }
76 for i in 0..n_b {
77 pooled[(n_a + i, j)] = data_b[(i, j)];
78 }
79 }
80 pooled
81}
82
83fn integrated_l2_mean_diff(
87 pooled: &FdMatrix,
88 labels: &[usize],
89 n_a: usize,
90 m: usize,
91 weights: &[f64],
92) -> f64 {
93 let mut mean_a = vec![0.0; m];
94 let mut mean_b = vec![0.0; m];
95 let n_b = labels.len() - n_a;
96 for (i, &lab) in labels.iter().enumerate() {
97 if lab == 0 {
98 for j in 0..m {
99 mean_a[j] += pooled[(i, j)];
100 }
101 } else {
102 for j in 0..m {
103 mean_b[j] += pooled[(i, j)];
104 }
105 }
106 }
107 for j in 0..m {
108 mean_a[j] /= n_a as f64;
109 mean_b[j] /= n_b as f64;
110 }
111 let mut acc = 0.0;
112 for j in 0..m {
113 let d = mean_a[j] - mean_b[j];
114 acc += d * d * weights[j];
115 }
116 acc.sqrt()
117}
118
119fn shuffle_labels(v: &mut [usize], rng: &mut StdRng) {
121 use rand::Rng;
122 let n = v.len();
123 for i in (1..n).rev() {
124 let j = rng.gen_range(0..=i);
125 v.swap(i, j);
126 }
127}
128
129pub fn t_perm_test(
153 data_a: &FdMatrix,
154 data_b: &FdMatrix,
155 argvals: &[f64],
156 n_perm: usize,
157 seed: u64,
158) -> Result<TestResult, FdarError> {
159 let (n_a, n_b, m) = validate_two_samples(data_a, data_b, argvals)?;
160 if n_perm == 0 {
161 return Err(FdarError::InvalidParameter {
162 parameter: "n_perm",
163 message: "must be >= 1".to_string(),
164 });
165 }
166
167 let weights = simpsons_weights(argvals);
168 let pooled = pool_two_samples(data_a, data_b, n_a, n_b, m);
169
170 let mut labels: Vec<usize> = (0..(n_a + n_b)).map(|i| usize::from(i >= n_a)).collect();
171 let observed = integrated_l2_mean_diff(&pooled, &labels, n_a, m, &weights);
172
173 let mut rng = StdRng::seed_from_u64(seed);
176 let mut n_ge = 0usize;
177 for _ in 0..n_perm {
178 shuffle_labels(&mut labels, &mut rng);
179 let perm_stat = integrated_l2_mean_diff(&pooled, &labels, n_a, m, &weights);
180 if perm_stat >= observed {
181 n_ge += 1;
182 }
183 }
184
185 let p_value = (n_ge as f64 + 1.0) / (n_perm as f64 + 1.0);
186 Ok(TestResult {
187 statistic: observed,
188 p_value,
189 n_perm,
190 })
191}
192
193pub fn f_perm_test(
217 data_a: &FdMatrix,
218 data_b: &FdMatrix,
219 argvals: &[f64],
220 n_perm: usize,
221 seed: u64,
222) -> Result<TestResult, FdarError> {
223 let (n_a, n_b, m) = validate_two_samples(data_a, data_b, argvals)?;
224 if n_perm == 0 {
225 return Err(FdarError::InvalidParameter {
226 parameter: "n_perm",
227 message: "must be >= 1".to_string(),
228 });
229 }
230
231 let pooled = pool_two_samples(data_a, data_b, n_a, n_b, m);
232 let labels_dedup = [0usize, 1usize];
233
234 let mut groups: Vec<usize> = (0..(n_a + n_b)).map(|i| usize::from(i >= n_a)).collect();
236 let observed = integrated_f_statistic(&pooled, &groups, &labels_dedup);
237
238 let mut rng = StdRng::seed_from_u64(seed);
241 let mut n_ge = 0usize;
242 for _ in 0..n_perm {
243 shuffle_labels(&mut groups, &mut rng);
244 let perm_stat = integrated_f_statistic(&pooled, &groups, &labels_dedup);
245 if perm_stat >= observed {
246 n_ge += 1;
247 }
248 }
249
250 let p_value = (n_ge as f64 + 1.0) / (n_perm as f64 + 1.0);
251 Ok(TestResult {
252 statistic: observed,
253 p_value,
254 n_perm,
255 })
256}
257
258#[cfg(test)]
259mod tests {
260 use super::*;
261 use crate::test_helpers::uniform_grid;
262
263 fn make_sample(n: usize, argvals: &[f64], shift: f64, seed: u64) -> FdMatrix {
267 let m = argvals.len();
268 let mut mat = FdMatrix::zeros(n, m);
269 let mut state = seed.wrapping_mul(2_654_435_761).wrapping_add(1);
271 for i in 0..n {
272 for (j, &t) in argvals.iter().enumerate() {
273 state = state
274 .wrapping_mul(6_364_136_223_846_793_005)
275 .wrapping_add(1_442_695_040_888_963_407);
276 let noise = ((state >> 33) as f64 / (1u64 << 31) as f64) - 1.0; mat[(i, j)] = (2.0 * std::f64::consts::PI * t).sin() + 0.1 * noise + shift;
278 }
279 }
280 mat
281 }
282
283 #[test]
284 fn t_perm_separated_small_p() {
285 let argvals = uniform_grid(25);
286 let a = make_sample(15, &argvals, 0.0, 1);
287 let b = make_sample(15, &argvals, 5.0, 2); let res = t_perm_test(&a, &b, &argvals, 199, 42).unwrap();
289 assert!(
290 res.p_value < 0.05,
291 "separated samples should give small p, got {}",
292 res.p_value
293 );
294 }
295
296 #[test]
297 fn t_perm_null_large_p() {
298 let argvals = uniform_grid(25);
299 let a = make_sample(15, &argvals, 0.0, 10);
300 let b = make_sample(15, &argvals, 0.0, 20); let res = t_perm_test(&a, &b, &argvals, 199, 7).unwrap();
302 assert!(
303 res.p_value > 0.1,
304 "null samples should give large p, got {}",
305 res.p_value
306 );
307 }
308
309 #[test]
310 fn t_perm_deterministic() {
311 let argvals = uniform_grid(20);
312 let a = make_sample(10, &argvals, 0.0, 3);
313 let b = make_sample(12, &argvals, 1.0, 4);
314 let r1 = t_perm_test(&a, &b, &argvals, 99, 123).unwrap();
315 let r2 = t_perm_test(&a, &b, &argvals, 99, 123).unwrap();
316 assert_eq!(r1, r2, "same seed must give bit-identical result");
317 }
318
319 #[test]
320 fn t_perm_invalid_input() {
321 let argvals = uniform_grid(20);
322 let a = make_sample(10, &argvals, 0.0, 5);
323 let argvals_b = uniform_grid(15);
325 let b = make_sample(10, &argvals_b, 0.0, 6);
326 assert!(matches!(
327 t_perm_test(&a, &b, &argvals, 99, 1),
328 Err(FdarError::InvalidDimension { .. })
329 ));
330 let b2 = make_sample(10, &argvals, 0.0, 7);
332 assert!(matches!(
333 t_perm_test(&a, &b2, &argvals, 0, 1),
334 Err(FdarError::InvalidParameter { .. })
335 ));
336 let a_small = make_sample(1, &argvals, 0.0, 8);
338 assert!(matches!(
339 t_perm_test(&a_small, &b2, &argvals, 99, 1),
340 Err(FdarError::InvalidDimension { .. })
341 ));
342 }
343
344 #[test]
345 fn f_perm_separated_small_p() {
346 let argvals = uniform_grid(25);
347 let a = make_sample(15, &argvals, 0.0, 11);
348 let b = make_sample(15, &argvals, 5.0, 12);
349 let res = f_perm_test(&a, &b, &argvals, 199, 42).unwrap();
350 assert!(
351 res.p_value < 0.05,
352 "separated samples should give small p, got {}",
353 res.p_value
354 );
355 }
356
357 #[test]
358 fn f_perm_null_large_p() {
359 let argvals = uniform_grid(25);
360 let a = make_sample(15, &argvals, 0.0, 30);
361 let b = make_sample(15, &argvals, 0.0, 40);
362 let res = f_perm_test(&a, &b, &argvals, 199, 7).unwrap();
363 assert!(
364 res.p_value > 0.1,
365 "null samples should give large p, got {}",
366 res.p_value
367 );
368 }
369
370 #[test]
371 fn f_perm_deterministic() {
372 let argvals = uniform_grid(20);
373 let a = make_sample(10, &argvals, 0.0, 3);
374 let b = make_sample(12, &argvals, 1.0, 4);
375 let r1 = f_perm_test(&a, &b, &argvals, 99, 555).unwrap();
376 let r2 = f_perm_test(&a, &b, &argvals, 99, 555).unwrap();
377 assert_eq!(r1, r2);
378 }
379
380 #[allow(deprecated)]
382 #[test]
383 fn f_perm_agrees_with_fanova_decision() {
384 use crate::function_on_scalar::fanova;
385 let argvals = uniform_grid(25);
386 let a = make_sample(15, &argvals, 0.0, 111);
387 let b = make_sample(15, &argvals, 5.0, 112);
388 let n_a = 15usize;
390 let n_b = 15usize;
391 let m = argvals.len();
392 let mut pooled = FdMatrix::zeros(n_a + n_b, m);
393 for j in 0..m {
394 for i in 0..n_a {
395 pooled[(i, j)] = a[(i, j)];
396 }
397 for i in 0..n_b {
398 pooled[(n_a + i, j)] = b[(i, j)];
399 }
400 }
401 let groups: Vec<usize> = (0..(n_a + n_b)).map(|i| usize::from(i >= n_a)).collect();
402 let fa = fanova(&pooled, &groups, 199).unwrap();
403 let fp = f_perm_test(&a, &b, &argvals, 199, 42).unwrap();
404 assert!(fa.p_value < 0.05);
406 assert!(fp.p_value < 0.05);
407 }
408
409 #[test]
410 fn f_perm_invalid_input() {
411 let argvals = uniform_grid(20);
412 let a = make_sample(10, &argvals, 0.0, 5);
413 let b2 = make_sample(10, &argvals, 0.0, 7);
414 assert!(matches!(
415 f_perm_test(&a, &b2, &argvals, 0, 1),
416 Err(FdarError::InvalidParameter { .. })
417 ));
418 let a_small = make_sample(1, &argvals, 0.0, 8);
419 assert!(matches!(
420 f_perm_test(&a_small, &b2, &argvals, 99, 1),
421 Err(FdarError::InvalidDimension { .. })
422 ));
423 }
424}