@@ -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