Skip to content

Commit d76617f

Browse files
committed
Add covariance convergence test
1 parent 24fb64f commit d76617f

1 file changed

Lines changed: 58 additions & 7 deletions

File tree

tests/dist.rs

Lines changed: 58 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -184,22 +184,73 @@ fn test_wishart_empirical_convergence() {
184184
let scale = matrix(c!(1.0, 0.5, 0.5, 1.0), 2, 2, Row);
185185
let w = MatDist::Wishart(5.0, scale);
186186

187-
let n_samples = 10_000;
187+
let n_samples = 20_000;
188188
let samples = w.sample_with_rng(&mut rng, n_samples);
189189

190-
let mut emp_mean = zeros(2, 2);
190+
let p = 2;
191+
let p_sq = p * p;
192+
let n_f64 = n_samples as f64;
193+
194+
let mut emp_mean = zeros(p, p);
191195
for s in &samples {
192196
emp_mean = &emp_mean + s;
193197
}
194198

195-
let n_f64 = n_samples as f64;
199+
for i in 0..p {
200+
for j in 0..p {
201+
emp_mean[(i, j)] /= n_f64;
202+
}
203+
}
204+
205+
let mut emp_cov = zeros(p_sq, p_sq);
206+
for s in &samples {
207+
for i in 0..p {
208+
for j in 0..p {
209+
for k in 0..p {
210+
for l in 0..p {
211+
let row_idx = i * p + j;
212+
let col_idx = k * p + l;
213+
214+
let diff1 = s[(i, j)] - emp_mean[(i, j)];
215+
let diff2 = s[(k, l)] - emp_mean[(k, l)];
216+
217+
emp_cov[(row_idx, col_idx)] += diff1 * diff2;
218+
}
219+
}
220+
}
221+
}
222+
}
223+
224+
for i in 0..p_sq {
225+
for j in 0..p_sq {
226+
emp_cov[(i, j)] /= n_f64 - 1.0;
227+
}
228+
}
229+
196230
let theo_mean = w.mean();
231+
let theo_cov = w.cov();
197232

198233
// Ensure empirical mean is within a small epsilon of theoretical mean
199-
for i in 0..2 {
200-
for j in 0..2 {
201-
emp_mean[(i, j)] /= n_f64;
202-
assert!((emp_mean[(i, j)] - theo_mean[(i, j)]).abs() < 0.1);
234+
for i in 0..p {
235+
for j in 0..p {
236+
assert!(
237+
(emp_mean[(i, j)] - theo_mean[(i, j)]).abs() < 0.1,
238+
"Mean failed to converge at ({}, {})",
239+
i,
240+
j
241+
);
242+
}
243+
}
244+
245+
// We have a slightly higher epsilon as covariance uses higher moments
246+
for i in 0..p_sq {
247+
for j in 0..p_sq {
248+
assert!(
249+
(emp_cov[(i, j)] - theo_cov[(i, j)]).abs() < 0.5,
250+
"Covariance failed to converge at ({}, {})",
251+
i,
252+
j
253+
);
203254
}
204255
}
205256
}

0 commit comments

Comments
 (0)