Skip to main content

torsh_tensor/
stats.rs

1//! Comprehensive tensor statistical operations
2//!
3//! This module provides a wide range of statistical functions for tensors,
4//! including descriptive statistics, percentiles, histograms, correlations,
5//! and probability distributions.
6
7use crate::{FloatElement, Tensor, TensorElement};
8use torsh_core::error::{Result, TorshError};
9
10/// Statistical computation modes
11#[derive(Debug, Clone, Copy, PartialEq)]
12pub enum StatMode {
13    /// Population statistics (divide by N)
14    Population,
15    /// Sample statistics (divide by N-1)
16    Sample,
17}
18
19/// Histogram configuration
20#[derive(Debug, Clone)]
21pub struct HistogramConfig {
22    /// Number of bins
23    pub bins: usize,
24    /// Minimum value (auto-computed if None)
25    pub min_val: Option<f64>,
26    /// Maximum value (auto-computed if None)
27    pub max_val: Option<f64>,
28    /// Include values outside range in first/last bins
29    pub include_outliers: bool,
30}
31
32impl Default for HistogramConfig {
33    fn default() -> Self {
34        Self {
35            bins: 50,
36            min_val: None,
37            max_val: None,
38            include_outliers: true,
39        }
40    }
41}
42
43/// Histogram result
44#[derive(Debug, Clone)]
45pub struct Histogram {
46    /// Bin counts
47    pub counts: Vec<usize>,
48    /// Bin edges (length = counts.len() + 1)
49    pub edges: Vec<f64>,
50    /// Total number of values
51    pub total_count: usize,
52}
53
54/// Correlation methods
55#[derive(Debug, Clone, Copy, PartialEq)]
56pub enum CorrelationMethod {
57    /// Pearson correlation coefficient
58    Pearson,
59    /// Spearman rank correlation
60    Spearman,
61    /// Kendall tau correlation
62    Kendall,
63}
64
65/// Statistical summary
66#[derive(Debug, Clone)]
67pub struct StatSummary {
68    pub count: usize,
69    pub mean: f64,
70    pub std: f64,
71    pub min: f64,
72    pub max: f64,
73    pub q25: f64, // 25th percentile
74    pub q50: f64, // 50th percentile (median)
75    pub q75: f64, // 75th percentile
76}
77
78/// Statistical operations for tensors
79impl<
80        T: TensorElement
81            + FloatElement
82            + Copy
83            + Default
84            + std::ops::Add<Output = T>
85            + std::ops::AddAssign
86            + std::ops::Sub<Output = T>
87            + std::ops::Mul<Output = T>
88            + std::ops::MulAssign
89            + std::ops::Div<Output = T>
90            + PartialOrd
91            + num_traits::FromPrimitive
92            + std::iter::Sum,
93    > Tensor<T>
94{
95    /// Compute mean along specified dimensions (legacy stats implementation)
96    pub fn mean_stats(&self, dims: Option<&[usize]>, keepdim: bool) -> Result<Self> {
97        let sum = if let Some(dims) = dims {
98            self.sum_dim(&dims.iter().map(|&d| d as i32).collect::<Vec<_>>(), keepdim)?
99        } else {
100            self.sum()?
101        };
102        let count = if let Some(dims) = dims {
103            dims.iter()
104                .map(|&d| self.shape().dims()[d])
105                .product::<usize>() as f64
106        } else {
107            self.numel() as f64
108        };
109
110        sum.div_scalar(
111            <T as num_traits::FromPrimitive>::from_f64(count)
112                .unwrap_or_else(|| <T as num_traits::One>::one()),
113        )
114    }
115
116    /// Compute variance along specified dimensions
117    ///
118    /// With `dims = Some(..)` the mean is computed per reduced slice (keepdim)
119    /// and broadcast back over the input, so multi-element reductions work; with
120    /// `dims = None` the whole tensor is reduced. `StatMode::Sample` applies
121    /// Bessel's correction and reports an error when the correction would leave
122    /// zero degrees of freedom.
123    pub fn var(&self, dims: Option<&[usize]>, keepdim: bool, mode: StatMode) -> Result<Self> {
124        let shape_binding = self.shape();
125        let input_shape = shape_binding.dims().to_vec();
126        let ndim = input_shape.len();
127
128        // Validate and normalise the requested dimensions up front.
129        let reduce_dims: Option<Vec<usize>> = match dims {
130            Some(requested) => {
131                for &dim in requested {
132                    if dim >= ndim {
133                        return Err(TorshError::InvalidArgument(format!(
134                            "Dimension {} out of range for {}-dimensional tensor",
135                            dim, ndim
136                        )));
137                    }
138                }
139                let mut normalized = requested.to_vec();
140                normalized.sort_unstable();
141                normalized.dedup();
142                Some(normalized)
143            }
144            None => None,
145        };
146
147        // Center the data against a mean of matching rank, broadcasting the
148        // dimension-wise mean instead of extracting a scalar with `item()`.
149        let diff = match &reduce_dims {
150            Some(reduce) => {
151                let mean = self.mean(Some(reduce), true)?;
152                let expanded = mean.expand(&input_shape)?;
153                self.sub(&expanded)?
154            }
155            None => {
156                let mean_value = self.mean(None, false)?.item()?;
157                self.sub_scalar(mean_value)?
158            }
159        };
160
161        let squared_diff = diff.mul_op(&diff)?;
162        let sum_sq = match &reduce_dims {
163            Some(reduce) => squared_diff.sum_dim(
164                &reduce.iter().map(|&d| d as i32).collect::<Vec<_>>(),
165                keepdim,
166            )?,
167            None => squared_diff.sum()?,
168        };
169
170        let count = match &reduce_dims {
171            Some(reduce) => reduce.iter().map(|&d| input_shape[d]).product::<usize>(),
172            None => self.numel(),
173        };
174
175        let divisor = match mode {
176            StatMode::Population => count,
177            StatMode::Sample => count.checked_sub(1).unwrap_or(0),
178        };
179
180        if divisor == 0 {
181            return Err(TorshError::InvalidArgument(
182                "Cannot compute variance with zero degrees of freedom".to_string(),
183            ));
184        }
185
186        sum_sq.div_scalar(
187            <T as num_traits::FromPrimitive>::from_usize(divisor)
188                .unwrap_or_else(|| <T as num_traits::One>::one()),
189        )
190    }
191
192    /// Compute standard deviation along specified dimensions
193    pub fn std(&self, dims: Option<&[usize]>, keepdim: bool, mode: StatMode) -> Result<Self> {
194        let variance = self.var(dims, keepdim, mode)?;
195        variance.sqrt()
196    }
197
198    /// Compute percentile along the last dimension
199    ///
200    /// `keepdim` keeps the reduced dimension with extent `1`, like PyTorch.
201    pub fn percentile(&self, q: f64, dim: Option<usize>, keepdim: bool) -> Result<Self> {
202        if !(0.0..=100.0).contains(&q) {
203            return Err(TorshError::InvalidArgument(format!(
204                "Percentile must be between 0 and 100, got {q}"
205            )));
206        }
207
208        let dim = dim.unwrap_or(self.shape().ndim() - 1);
209        if dim >= self.shape().ndim() {
210            return Err(TorshError::dimension_error(
211                &format!(
212                    "Dimension {} out of bounds for tensor with {} dimensions",
213                    dim,
214                    self.shape().ndim()
215                ),
216                "tensor statistics operation",
217            ));
218        }
219
220        // For now, implement a simple linear interpolation method
221        // In a full implementation, this would be optimized
222        let (sorted, _indices) = self.sort(Some(dim as i32), false)?; // Sort in ascending order
223        let size = self.shape().dims()[dim];
224
225        // Calculate the position in the sorted array
226        let pos = q / 100.0 * (size - 1) as f64;
227        let lower_idx = pos.floor() as usize;
228        let upper_idx = (pos.ceil() as usize).min(size - 1);
229        let weight = pos - pos.floor();
230
231        let reduced = if lower_idx == upper_idx {
232            // Exact index
233            sorted.select(dim as i32, lower_idx as i64)?
234        } else {
235            // Interpolate between two values
236            let lower = sorted.select(dim as i32, lower_idx as i64)?;
237            let upper = sorted.select(dim as i32, upper_idx as i64)?;
238            let diff = upper.sub(&lower)?;
239            let weight_scalar = <T as TensorElement>::from_f64(weight)
240                .unwrap_or_else(|| <T as TensorElement>::from_f64(0.0).unwrap_or_default());
241            let weighted_diff = diff.mul_scalar(weight_scalar)?;
242            lower.add_op(&weighted_diff)?
243        };
244
245        if keepdim {
246            reduced.unsqueeze(dim as i32)
247        } else {
248            Ok(reduced)
249        }
250    }
251
252    /// Compute median (50th percentile)
253    pub fn median(&self, dim: Option<usize>, keepdim: bool) -> Result<Self> {
254        self.percentile(50.0, dim, keepdim)
255    }
256
257    /// Compute quantiles at specified levels
258    pub fn quantile(&self, q: &[f64], dim: Option<usize>, keepdim: bool) -> Result<Vec<Self>> {
259        let mut results = Vec::new();
260        for &quantile in q {
261            results.push(self.percentile(quantile * 100.0, dim, keepdim)?);
262        }
263        Ok(results)
264    }
265
266    /// Create histogram of tensor values
267    pub fn histogram(&self, config: &HistogramConfig) -> Result<Histogram> {
268        let data = self.to_vec()?;
269
270        if data.is_empty() {
271            return Ok(Histogram {
272                counts: vec![0; config.bins],
273                edges: (0..=config.bins).map(|i| i as f64).collect(),
274                total_count: 0,
275            });
276        }
277
278        // Compute min and max if not provided
279        let min_val = config.min_val.unwrap_or_else(|| {
280            data.iter()
281                .map(|&x| TensorElement::to_f64(&x).expect("f64 conversion should succeed"))
282                .fold(f64::INFINITY, f64::min)
283        });
284        let max_val = config.max_val.unwrap_or_else(|| {
285            data.iter()
286                .map(|&x| TensorElement::to_f64(&x).expect("f64 conversion should succeed"))
287                .fold(f64::NEG_INFINITY, f64::max)
288        });
289
290        if min_val >= max_val {
291            return Err(TorshError::InvalidArgument(
292                "Minimum value must be less than maximum value".to_string(),
293            ));
294        }
295
296        // Create bin edges
297        let bin_width = (max_val - min_val) / config.bins as f64;
298        let edges: Vec<f64> = (0..=config.bins)
299            .map(|i| min_val + i as f64 * bin_width)
300            .collect();
301
302        // Count values in each bin
303        let mut counts = vec![0; config.bins];
304        for &value in data.iter() {
305            let val = TensorElement::to_f64(&value).expect("f64 conversion should succeed");
306
307            let bin_idx = if val <= min_val {
308                if config.include_outliers {
309                    0
310                } else {
311                    continue;
312                }
313            } else if val >= max_val {
314                if config.include_outliers {
315                    config.bins - 1
316                } else {
317                    continue;
318                }
319            } else {
320                ((val - min_val) / bin_width).floor() as usize
321            };
322
323            let bin_idx = bin_idx.min(config.bins - 1);
324            counts[bin_idx] += 1;
325        }
326
327        Ok(Histogram {
328            counts,
329            edges,
330            total_count: data.len(),
331        })
332    }
333
334    /// Compute correlation coefficient with another tensor
335    pub fn correlation(&self, other: &Self, method: CorrelationMethod) -> Result<T> {
336        if self.shape() != other.shape() {
337            return Err(TorshError::ShapeMismatch {
338                expected: self.shape().dims().to_vec(),
339                got: other.shape().dims().to_vec(),
340            });
341        }
342
343        match method {
344            CorrelationMethod::Pearson => self.pearson_correlation(other),
345            CorrelationMethod::Spearman => self.spearman_correlation(other),
346            CorrelationMethod::Kendall => self.kendall_correlation(other),
347        }
348    }
349
350    /// Pearson correlation coefficient
351    fn pearson_correlation(&self, other: &Self) -> Result<T> {
352        let n = self.numel() as f64;
353        if n < 2.0 {
354            return Err(TorshError::InvalidArgument(
355                "Need at least 2 values for correlation".to_string(),
356            ));
357        }
358
359        // Compute means
360        let mean_x = self.mean(None, false)?;
361        let mean_y = other.mean(None, false)?;
362
363        let mean_x_data = mean_x.to_vec()?;
364        let mean_y_data = mean_y.to_vec()?;
365        let mean_x_val = mean_x_data[0];
366        let mean_y_val = mean_y_data[0];
367
368        // Compute deviations and products
369        let self_data = self.to_vec()?;
370        let other_data = other.to_vec()?;
371
372        let mut sum_xy = 0.0;
373        let mut sum_xx = 0.0;
374        let mut sum_yy = 0.0;
375
376        for (&x, &y) in self_data.iter().zip(other_data.iter()) {
377            let dx = TensorElement::to_f64(&x).expect("f64 conversion should succeed")
378                - TensorElement::to_f64(&mean_x_val).expect("f64 conversion should succeed");
379            let dy = TensorElement::to_f64(&y).expect("f64 conversion should succeed")
380                - TensorElement::to_f64(&mean_y_val).expect("f64 conversion should succeed");
381
382            sum_xy += dx * dy;
383            sum_xx += dx * dx;
384            sum_yy += dy * dy;
385        }
386
387        let denominator = (sum_xx * sum_yy).sqrt();
388        if denominator.abs() < f64::EPSILON {
389            return Err(TorshError::InvalidArgument(
390                "Cannot compute correlation: one variable has zero variance".to_string(),
391            ));
392        }
393
394        let correlation = sum_xy / denominator;
395        Ok(<T as TensorElement>::from_f64(correlation).expect("f64 conversion should succeed"))
396    }
397
398    /// Spearman rank correlation coefficient
399    fn spearman_correlation(&self, other: &Self) -> Result<T> {
400        // Convert to ranks and compute Pearson correlation of ranks
401        let self_ranks = self.rank()?;
402        let other_ranks = other.rank()?;
403        self_ranks.pearson_correlation(&other_ranks)
404    }
405
406    /// Kendall tau correlation coefficient
407    fn kendall_correlation(&self, other: &Self) -> Result<T> {
408        let n = self.numel();
409        if n < 2 {
410            return Err(TorshError::InvalidArgument(
411                "Need at least 2 values for Kendall correlation".to_string(),
412            ));
413        }
414
415        let self_data = self.to_vec()?;
416        let other_data = other.to_vec()?;
417
418        let mut concordant = 0;
419        let mut discordant = 0;
420        let mut tied_x = 0;
421        let mut tied_y = 0;
422        let mut tied_xy = 0;
423
424        for i in 0..n {
425            for j in i + 1..n {
426                let x1 =
427                    TensorElement::to_f64(&self_data[i]).expect("f64 conversion should succeed");
428                let x2 =
429                    TensorElement::to_f64(&self_data[j]).expect("f64 conversion should succeed");
430                let y1 =
431                    TensorElement::to_f64(&other_data[i]).expect("f64 conversion should succeed");
432                let y2 =
433                    TensorElement::to_f64(&other_data[j]).expect("f64 conversion should succeed");
434
435                let dx = x2 - x1;
436                let dy = y2 - y1;
437
438                if dx.abs() < f64::EPSILON && dy.abs() < f64::EPSILON {
439                    tied_xy += 1;
440                } else if dx.abs() < f64::EPSILON {
441                    tied_x += 1;
442                } else if dy.abs() < f64::EPSILON {
443                    tied_y += 1;
444                } else if dx * dy > 0.0 {
445                    concordant += 1;
446                } else {
447                    discordant += 1;
448                }
449            }
450        }
451
452        let total_pairs = n * (n - 1) / 2;
453        let effective_pairs = total_pairs - tied_x - tied_y - tied_xy;
454
455        if effective_pairs == 0 {
456            return Ok(<T as TensorElement>::from_f64(0.0).expect("f64 conversion should succeed"));
457        }
458
459        let tau = (concordant as f64 - discordant as f64) / effective_pairs as f64;
460        Ok(<T as TensorElement>::from_f64(tau).expect("f64 conversion should succeed"))
461    }
462
463    /// Compute ranks of tensor elements
464    fn rank(&self) -> Result<Self> {
465        let data = self.to_vec()?;
466        let n = data.len();
467
468        // Create index-value pairs and sort by value
469        let mut indexed_values: Vec<(usize, T)> =
470            data.iter().enumerate().map(|(i, &val)| (i, val)).collect();
471
472        indexed_values.sort_by(|a, b| {
473            TensorElement::to_f64(&a.1)
474                .expect("f64 conversion should succeed")
475                .partial_cmp(&TensorElement::to_f64(&b.1).expect("f64 conversion should succeed"))
476                .unwrap_or(std::cmp::Ordering::Equal)
477        });
478
479        // Assign ranks (1-based, handle ties with average rank)
480        let mut ranks = vec![T::default(); n];
481        let mut i = 0;
482
483        while i < n {
484            let mut j = i;
485            while j < n
486                && TensorElement::to_f64(&indexed_values[j].1)
487                    .expect("f64 conversion should succeed")
488                    == TensorElement::to_f64(&indexed_values[i].1)
489                        .expect("f64 conversion should succeed")
490            {
491                j += 1;
492            }
493
494            // Average rank for tied values
495            let avg_rank = (i + j + 1) as f64 / 2.0;
496            for k in i..j {
497                ranks[indexed_values[k].0] = <T as TensorElement>::from_f64(avg_rank)
498                    .expect("f64 conversion should succeed");
499            }
500            i = j;
501        }
502
503        Self::from_data(ranks, self.shape().dims().to_vec(), self.device)
504    }
505
506    /// Generate comprehensive statistical summary
507    pub fn describe(&self) -> Result<StatSummary> {
508        let data = self.to_vec()?;
509        if data.is_empty() {
510            return Err(TorshError::InvalidArgument(
511                "Cannot compute statistics on empty tensor".to_string(),
512            ));
513        }
514
515        let count = data.len();
516        let values: Vec<f64> = data
517            .iter()
518            .map(|&x| TensorElement::to_f64(&x).expect("f64 conversion should succeed"))
519            .collect();
520
521        // Basic statistics
522        let sum: f64 = values.iter().sum();
523        let mean = sum / count as f64;
524
525        let variance =
526            values.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / (count - 1).max(1) as f64;
527        let std = variance.sqrt();
528
529        let min = values.iter().fold(f64::INFINITY, |a, &b| a.min(b));
530        let max = values.iter().fold(f64::NEG_INFINITY, |a, &b| a.max(b));
531
532        // Percentiles
533        let mut sorted_values = values.clone();
534        sorted_values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
535
536        let q25 = percentile_sorted(&sorted_values, 25.0);
537        let q50 = percentile_sorted(&sorted_values, 50.0);
538        let q75 = percentile_sorted(&sorted_values, 75.0);
539
540        Ok(StatSummary {
541            count,
542            mean,
543            std,
544            min,
545            max,
546            q25,
547            q50,
548            q75,
549        })
550    }
551
552    /// Compute covariance matrix for 2D tensor (each column is a variable)
553    pub fn cov(&self, mode: StatMode) -> Result<Self> {
554        let shape = self.shape();
555        if shape.ndim() != 2 {
556            return Err(TorshError::dimension_error(
557                "Covariance matrix requires 2D tensor",
558                "covariance computation",
559            ));
560        }
561
562        let (n_samples, n_features) = (shape.dims()[0], shape.dims()[1]);
563        if n_samples < 2 {
564            return Err(TorshError::InvalidArgument(
565                "Need at least 2 samples for covariance".to_string(),
566            ));
567        }
568
569        // Center the data by subtracting per-column means
570        // Compute mean of each column (feature) independently
571        let data = self.to_vec()?;
572        let mut centered_data = data.clone();
573        for feat in 0..n_features {
574            let col_sum = (0..n_samples).fold(<T as num_traits::Zero>::zero(), |acc, row| {
575                acc + data[row * n_features + feat]
576            });
577            let col_mean = col_sum
578                / <T as num_traits::FromPrimitive>::from_usize(n_samples)
579                    .unwrap_or_else(|| <T as num_traits::One>::one());
580            for row in 0..n_samples {
581                centered_data[row * n_features + feat] = data[row * n_features + feat] - col_mean;
582            }
583        }
584        let centered = Self::from_data(centered_data, vec![n_samples, n_features], self.device())?;
585
586        // Compute covariance matrix: (X^T * X) / (n - 1)
587        let centered_t = centered.transpose(1, 0)?;
588        let cov_unnormalized = centered_t.matmul(&centered)?;
589
590        let divisor = match mode {
591            StatMode::Population => n_samples,
592            StatMode::Sample => n_samples - 1,
593        };
594
595        cov_unnormalized.div_scalar(
596            <T as num_traits::FromPrimitive>::from_usize(divisor)
597                .unwrap_or_else(|| <T as num_traits::One>::one()),
598        )
599    }
600
601    /// Compute correlation matrix for 2D tensor
602    pub fn corrcoef(&self) -> Result<Self> {
603        let cov_matrix = self.cov(StatMode::Sample)?;
604        let cov_data = cov_matrix.to_vec()?;
605        let n_features = cov_matrix.shape().dims()[0];
606
607        // Extract diagonal elements (variances)
608        let mut std_devs = Vec::with_capacity(n_features);
609        for i in 0..n_features {
610            let variance = TensorElement::to_f64(&cov_data[i * n_features + i])
611                .expect("f64 conversion should succeed");
612            std_devs.push(variance.sqrt());
613        }
614
615        // Normalize covariance matrix to get correlation matrix
616        let mut corr_data = Vec::with_capacity(cov_data.len());
617        for i in 0..n_features {
618            for j in 0..n_features {
619                let cov_val = TensorElement::to_f64(&cov_data[i * n_features + j])
620                    .expect("f64 conversion should succeed");
621                let corr_val = if std_devs[i] > f64::EPSILON && std_devs[j] > f64::EPSILON {
622                    cov_val / (std_devs[i] * std_devs[j])
623                } else {
624                    0.0
625                };
626                corr_data.push(
627                    <T as TensorElement>::from_f64(corr_val)
628                        .expect("f64 conversion should succeed"),
629                );
630            }
631        }
632
633        Self::from_data(corr_data, vec![n_features, n_features], self.device)
634    }
635}
636
637/// Helper function to compute percentile from sorted array
638fn percentile_sorted(sorted_values: &[f64], q: f64) -> f64 {
639    if sorted_values.is_empty() {
640        return 0.0;
641    }
642
643    let pos = q / 100.0 * (sorted_values.len() - 1) as f64;
644    let lower_idx = pos.floor() as usize;
645    let upper_idx = (pos.ceil() as usize).min(sorted_values.len() - 1);
646
647    if lower_idx == upper_idx {
648        sorted_values[lower_idx]
649    } else {
650        let weight = pos - pos.floor();
651        sorted_values[lower_idx] * (1.0 - weight) + sorted_values[upper_idx] * weight
652    }
653}
654
655#[cfg(test)]
656mod tests {
657    use super::*;
658    use torsh_core::device::DeviceType;
659
660    #[test]
661    fn test_basic_statistics() {
662        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
663        let tensor = Tensor::from_data(data, vec![5], DeviceType::Cpu)
664            .expect("tensor creation should succeed");
665
666        let mean = tensor.mean(None, false).expect("mean should succeed");
667        assert!(
668            (mean.to_vec().expect("to_vec conversion should succeed")[0] - 3.0_f32).abs()
669                < 1e-6_f32
670        );
671
672        let var_sample = tensor
673            .var(None, false, StatMode::Sample)
674            .expect("variance should succeed");
675        assert!(
676            (var_sample
677                .to_vec()
678                .expect("to_vec conversion should succeed")[0]
679                - 2.5_f32)
680                .abs()
681                < 1e-6_f32
682        );
683
684        let std_sample = tensor
685            .std(None, false, StatMode::Sample)
686            .expect("std should succeed");
687        assert!(
688            (std_sample
689                .to_vec()
690                .expect("to_vec conversion should succeed")[0]
691                - 2.5_f32.sqrt())
692            .abs()
693                < 1e-6
694        );
695    }
696
697    #[test]
698    fn test_percentiles() {
699        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
700        let tensor = Tensor::from_data(data, vec![10], DeviceType::Cpu)
701            .expect("tensor creation should succeed");
702
703        let median = tensor.median(None, false).expect("median should succeed");
704        assert!(
705            (median.to_vec().expect("to_vec conversion should succeed")[0] - 5.5_f32).abs()
706                < 1e-6_f32
707        );
708
709        let q25 = tensor
710            .percentile(25.0, None, false)
711            .expect("percentile should succeed");
712        assert!(
713            (q25.to_vec().expect("to_vec conversion should succeed")[0] - 3.25_f32).abs()
714                < 1e-6_f32
715        );
716
717        let q75 = tensor
718            .percentile(75.0, None, false)
719            .expect("percentile should succeed");
720        assert!(
721            (q75.to_vec().expect("to_vec conversion should succeed")[0] - 7.75_f32).abs()
722                < 1e-6_f32
723        );
724    }
725
726    #[test]
727    fn test_histogram() {
728        let data = vec![1.0, 2.0, 2.0, 3.0, 3.0, 3.0, 4.0, 4.0, 5.0];
729        let tensor = Tensor::from_data(data, vec![9], DeviceType::Cpu)
730            .expect("tensor creation should succeed");
731
732        let config = HistogramConfig {
733            bins: 5,
734            min_val: Some(1.0),
735            max_val: Some(5.0),
736            include_outliers: true,
737        };
738
739        let hist = tensor.histogram(&config).expect("histogram should succeed");
740        assert_eq!(hist.total_count, 9);
741        assert_eq!(hist.counts.len(), 5);
742        assert_eq!(hist.edges.len(), 6);
743    }
744
745    #[test]
746    fn test_correlation() {
747        let x_data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
748        let y_data = vec![2.0, 4.0, 6.0, 8.0, 10.0]; // Perfect positive correlation
749
750        let x = Tensor::from_data(x_data, vec![5], DeviceType::Cpu)
751            .expect("tensor creation should succeed");
752        let y = Tensor::from_data(y_data, vec![5], DeviceType::Cpu)
753            .expect("tensor creation should succeed");
754
755        let corr = x
756            .correlation(&y, CorrelationMethod::Pearson)
757            .expect("correlation should succeed");
758        assert!((corr - 1.0_f32).abs() < 1e-6_f32); // Should be 1.0 for perfect positive correlation
759    }
760
761    #[test]
762    fn test_statistical_summary() {
763        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
764        let tensor = Tensor::from_data(data, vec![10], DeviceType::Cpu)
765            .expect("tensor creation should succeed");
766
767        let summary = tensor.describe().expect("describe should succeed");
768        assert_eq!(summary.count, 10);
769        assert!((summary.mean - 5.5).abs() < 1e-6);
770        assert!((summary.q50 - 5.5).abs() < 1e-6); // Median
771        assert_eq!(summary.min, 1.0);
772        assert_eq!(summary.max, 10.0);
773    }
774
775    #[test]
776    fn test_covariance_matrix() {
777        // Create a 2D tensor (samples x features)
778        let data = vec![1.0, 2.0, 2.0, 4.0, 3.0, 6.0, 4.0, 8.0];
779        let tensor = Tensor::from_data(data, vec![4, 2], DeviceType::Cpu)
780            .expect("tensor creation should succeed");
781
782        let cov_matrix = tensor
783            .cov(StatMode::Sample)
784            .expect("covariance should succeed");
785        assert_eq!(cov_matrix.shape().dims(), &[2, 2]);
786
787        // Check that it's symmetric (off-diagonal elements should be equal)
788        let cov_data = cov_matrix
789            .to_vec()
790            .expect("to_vec conversion should succeed");
791        // For a 2x2 matrix [[a, b], [c, d]] stored as [a, b, c, d]:
792        // Symmetry means b == c (cov_data[1] == cov_data[2])
793        assert!((cov_data[1] as f64 - cov_data[2] as f64).abs() < 1e-6);
794    }
795
796    #[test]
797    fn test_covariance_column_wise() {
798        // Two features with different column means: col0 has mean 2.0, col1 has mean 20.0
799        // Data: [[1, 10], [2, 20], [3, 30]]  =>  col means: [2, 20]
800        // Centered: [[-1,-10],[0,0],[1,10]]
801        // Cov sample (n-1=2):
802        //   cov(0,0) = ((-1)^2 + 0^2 + 1^2) / 2 = 1.0
803        //   cov(0,1) = cov(1,0) = ((-1)(-10)+0+1*10) / 2 = 20/2 = 10.0
804        //   cov(1,1) = (100+0+100) / 2 = 100.0
805        let data: Vec<f32> = vec![1.0, 10.0, 2.0, 20.0, 3.0, 30.0];
806        let tensor = Tensor::from_data(data, vec![3, 2], DeviceType::Cpu)
807            .expect("tensor creation should succeed");
808
809        let cov = tensor.cov(StatMode::Sample).expect("cov should succeed");
810        let cov_data = cov.to_vec().expect("to_vec should succeed");
811
812        // cov_data = [cov00, cov01, cov10, cov11]
813        assert!(
814            (cov_data[0] - 1.0_f32).abs() < 1e-4,
815            "cov00 = {}",
816            cov_data[0]
817        );
818        assert!(
819            (cov_data[1] - 10.0_f32).abs() < 1e-4,
820            "cov01 = {}",
821            cov_data[1]
822        );
823        assert!(
824            (cov_data[2] - 10.0_f32).abs() < 1e-4,
825            "cov10 = {}",
826            cov_data[2]
827        );
828        assert!(
829            (cov_data[3] - 100.0_f32).abs() < 1e-4,
830            "cov11 = {}",
831            cov_data[3]
832        );
833    }
834
835    #[test]
836    fn test_ranks() {
837        let data = vec![3.0, 1.0, 4.0, 2.0, 2.0]; // Contains ties
838        let tensor = Tensor::from_data(data, vec![5], DeviceType::Cpu)
839            .expect("tensor creation should succeed");
840
841        let ranks = tensor.rank().expect("rank should be available");
842        let rank_data = ranks.to_vec().expect("to_vec conversion should succeed");
843
844        // Check that ranks are in expected range
845        for &rank in rank_data.iter() {
846            assert!((1.0_f32..=5.0_f32).contains(&rank));
847        }
848    }
849}