Skip to main content

datui_lib/analysis/
statistics.rs

1use color_eyre::Result;
2use color_eyre::eyre::Report;
3use polars::polars_compute::rolling::QuantileMethod;
4use polars::prelude::*;
5use std::collections::HashMap;
6use std::ops::Range;
7
8/// Collect a LazyFrame, with the streaming engine when the `streaming` feature is on and
9/// `use_streaming` asks. Returns `PolarsError` for callers that store or show it
10/// (e.g. `DataTableState::error`).
11pub fn collect_lazy(
12    lf: LazyFrame,
13    use_streaming: bool,
14) -> std::result::Result<DataFrame, PolarsError> {
15    #[cfg(feature = "streaming")]
16    {
17        if may_stream(&lf, use_streaming) && !sorts_by_one_wide_key(&lf) {
18            // A plain collect is always one frame; `Multiple` only comes from sink_multiple.
19            lf.collect_with_engine(Engine::Streaming)
20                .map(|result| result.unwrap_single())
21        } else {
22            lf.collect()
23        }
24    }
25    #[cfg(not(feature = "streaming"))]
26    {
27        let _ = use_streaming; // ignored when streaming feature is disabled
28        lf.collect()
29    }
30}
31
32/// Whether a query over `lf` may use the streaming engine: asked for, and possible.
33/// Polars 0.55's streaming engine cannot run an anonymous scan (a SQLite table): it
34/// stops at a `todo!`.
35pub fn may_stream(lf: &LazyFrame, wanted: bool) -> bool {
36    use polars::lazy::dsl::{DslPlan, FileScanDsl};
37    wanted
38        && !lf.logical_plan.into_iter().any(|node| match node {
39            DslPlan::Scan { scan_type, .. } => {
40                matches!(scan_type.as_ref(), FileScanDsl::Anonymous { .. })
41            }
42            _ => false,
43        })
44}
45
46/// Whether `lf` takes the first rows of an unstable sort by one Decimal or Int128 key:
47/// Polars 0.55 streams that as a top-k, which panics on those dtypes, so it runs in
48/// memory. A spec's `scale` reads as Decimal and a query's `by` group sort is unstable.
49/// Stable sorts (with a row index key), full sorts and deeper slices still stream.
50#[cfg(feature = "streaming")]
51fn sorts_by_one_wide_key(lf: &LazyFrame) -> bool {
52    use polars::lazy::dsl::DslPlan;
53    let wide_sort = |node: &DslPlan| match node {
54        DslPlan::Sort {
55            input,
56            by_column,
57            sort_options,
58            ..
59        } if by_column.len() == 1 && !sort_options.maintain_order => {
60            LazyFrame::from((**input).clone())
61                .select([by_column[0].clone()])
62                .collect_schema()
63                .ok()
64                .and_then(|schema| schema.get_at_index(0).map(|(_, dtype)| dtype.clone()))
65                .is_some_and(|dtype| dtype.is_decimal() || dtype == DataType::Int128)
66        }
67        _ => false,
68    };
69    lf.logical_plan.into_iter().any(|node| match node {
70        DslPlan::Slice {
71            input, offset: 0, ..
72        } => input.into_iter().any(wide_sort),
73        DslPlan::Sort { sort_options, .. } => sort_options.limit.is_some() && wide_sort(node),
74        _ => false,
75    })
76}
77
78#[derive(Clone)]
79pub struct ColumnStatistics {
80    pub name: String,
81    pub count: usize,
82    pub null_count: usize,
83    pub numeric_stats: Option<NumericStatistics>,
84    pub categorical_stats: Option<CategoricalStatistics>,
85    pub temporal_stats: Option<TemporalStatistics>,
86}
87
88#[derive(Clone)]
89pub struct NumericStatistics {
90    pub mean: f64,
91    pub std: f64,
92    pub min: f64,
93    pub max: f64,
94    pub median: f64,
95    pub q25: f64,
96    pub q75: f64,
97    pub percentiles: HashMap<u8, f64>, // 1, 5, 25, 50, 75, 95, 99
98    pub skewness: f64,
99    pub kurtosis: f64,
100}
101
102#[derive(Clone)]
103pub struct CategoricalStatistics {
104    pub min: Option<String>, // Lexicographically smallest string
105    pub max: Option<String>, // Lexicographically largest string
106}
107
108/// Describe for a Date, Datetime, Time or Duration column: each statistic a value of the
109/// column's own type, written as the table writes it; `None` for a null.
110#[derive(Clone, Default)]
111pub struct TemporalStatistics {
112    pub mean: Option<String>,
113    pub min: Option<String>,
114    pub q25: Option<String>,
115    pub median: Option<String>,
116    pub q75: Option<String>,
117    pub max: Option<String>,
118}
119
120#[derive(Clone)]
121pub struct DistributionAnalysis {
122    pub column_name: String,
123    pub distribution_type: DistributionType,
124    /// The chosen family's p-value; with no clear fit, the best any family managed.
125    pub confidence: f64,
126    pub characteristics: DistributionCharacteristics,
127    pub outliers: OutlierAnalysis,
128    pub percentiles: PercentileBreakdown,
129    /// At most five thousand of the column's finite values, spread across it, sorted.
130    pub sorted_sample_values: Vec<f64>,
131    /// Every family's fit and test, or why it does not apply.
132    pub fits: Vec<(
133        DistributionType,
134        crate::analysis::distribution_fit::FitOutcome,
135    )>,
136    /// Each fitted family's quantiles at the plotting positions of
137    /// `sorted_sample_values`, for its Q-Q plot: computed with the fit, not per frame.
138    pub qq: Vec<(DistributionType, Vec<f64>)>,
139    /// The histogram last drawn, kept for the next frame; see [`Self::histogram`].
140    pub histogram: HistogramCache,
141}
142
143/// What a histogram of [`DistributionAnalysis::sorted_sample_values`] is drawn for:
144/// one family's fit, a number of bins over a range of values, and the scale.
145#[derive(Debug, Clone, Copy, PartialEq)]
146pub struct HistogramKey {
147    pub family: DistributionType,
148    pub bins: usize,
149    /// Bins equal in log space, the curve at each bin's geometric middle.
150    pub log: bool,
151    /// The values the bins span, positive on a log scale.
152    pub range: (f64, f64),
153    /// Points along a continuous family's density on linear bins.
154    pub samples: usize,
155}
156
157/// A histogram and its family's expected counts, ready to draw.
158#[derive(Debug)]
159pub struct Histogram {
160    /// Values in each bin.
161    pub counts: Vec<usize>,
162    /// The count axis's top: the tallest bar or expected count, rounded up to even
163    /// so the middle label is a whole count.
164    pub top: f64,
165    /// The fit's expected counts as a curve, x where the axis puts each point and y
166    /// on the 0-100 scale the bars stand on. Empty when the family does not apply.
167    pub curve: Vec<(f64, f64)>,
168}
169
170/// The last [`Histogram`] built for an analysis. A copy starts empty.
171#[derive(Debug, Default)]
172pub struct HistogramCache(std::sync::Mutex<Option<(HistogramKey, std::sync::Arc<Histogram>)>>);
173
174impl Clone for HistogramCache {
175    fn clone(&self) -> Self {
176        Self::default()
177    }
178}
179
180#[cfg(test)]
181thread_local! {
182    /// Histograms this thread has built, for the tests that count them.
183    pub(crate) static HISTOGRAMS_BUILT: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
184}
185
186impl DistributionAnalysis {
187    pub fn fit(
188        &self,
189        family: DistributionType,
190    ) -> Option<&crate::analysis::distribution_fit::FitOutcome> {
191        self.fits
192            .iter()
193            .find(|(fitted, _)| *fitted == family)
194            .map(|(_, outcome)| outcome)
195    }
196
197    pub fn qq(&self, family: DistributionType) -> Option<&[f64]> {
198        self.qq
199            .iter()
200            .find(|(fitted, _)| *fitted == family)
201            .map(|(_, quantiles)| quantiles.as_slice())
202    }
203
204    /// The histogram for `key`: the one last built when the key is the same, so a
205    /// frame that changes nothing counts nothing and evaluates no CDF.
206    pub fn histogram(&self, key: HistogramKey) -> std::sync::Arc<Histogram> {
207        let mut cache = self
208            .histogram
209            .0
210            .lock()
211            .unwrap_or_else(std::sync::PoisonError::into_inner);
212        if let Some((cached, histogram)) = cache.as_ref()
213            && *cached == key
214        {
215            return std::sync::Arc::clone(histogram);
216        }
217        #[cfg(test)]
218        HISTOGRAMS_BUILT.with(|built| built.set(built.get() + 1));
219        let histogram = std::sync::Arc::new(self.build_histogram(key));
220        *cache = Some((key, std::sync::Arc::clone(&histogram)));
221        histogram
222    }
223
224    fn build_histogram(&self, key: HistogramKey) -> Histogram {
225        let HistogramKey {
226            family,
227            bins,
228            log,
229            range: (low, high),
230            samples,
231        } = key;
232        let sorted = &self.sorted_sample_values;
233        let n = sorted.len() as f64;
234        let edges: Vec<f64> = if log {
235            let (log_low, log_high) = (low.ln(), high.ln());
236            let width = (log_high - log_low) / bins as f64;
237            (0..=bins)
238                .map(|i| (log_low + i as f64 * width).exp())
239                .collect()
240        } else {
241            let width = (high - low) / bins as f64;
242            (0..=bins).map(|i| low + i as f64 * width).collect()
243        };
244        // Sorted, so a bin's count is the distance between where its edges fall:
245        // each bin holds its lower edge, the last its upper edge too.
246        let below = |edge: f64| sorted.partition_point(|value| *value < edge);
247        let counts: Vec<usize> = (0..bins)
248            .map(|i| {
249                let end = if i + 1 == bins {
250                    sorted.partition_point(|value| *value <= edges[i + 1])
251                } else {
252                    below(edges[i + 1])
253                };
254                end.saturating_sub(below(edges[i]))
255            })
256            .collect();
257
258        // Expected counts from the fit every view of this family uses, by the CDF
259        // across each bin: exact for log-scaled and whole-number bins, where a density
260        // at the center is not.
261        let fitted = self
262            .fit(family)
263            .and_then(|outcome| outcome.test())
264            .map(|test| &test.fitted);
265        let expected: Vec<f64> = match fitted {
266            Some(fitted) => edges
267                .windows(2)
268                .enumerate()
269                .map(|(i, edge)| {
270                    let upper = if i + 1 == bins {
271                        fitted.cdf(edge[1])
272                    } else {
273                        fitted.cdf_below(edge[1])
274                    };
275                    (upper - fitted.cdf_below(edge[0])).max(0.0) * n
276                })
277                .collect(),
278            None => vec![0.0; bins],
279        };
280        let tallest = counts.iter().copied().max().unwrap_or(0);
281        let expected_top = expected.iter().copied().fold(0.0, f64::max);
282        let top = (tallest.max(expected_top.ceil() as usize).max(1) as f64 / 2.0).ceil() * 2.0;
283        let height = |count: f64| count / top * 100.0;
284
285        let curve = match fitted {
286            // A continuous family on linear bins is drawn as its density, scaled to a
287            // bin's count: a smooth curve rather than a staircase.
288            Some(fitted) if !fitted.discrete() && !log && high > low => {
289                let bin_width = (high - low) / bins as f64;
290                (0..samples)
291                    .map(|i| {
292                        let x = low + i as f64 / (samples - 1) as f64 * (high - low);
293                        (x, height(fitted.density(x) * bin_width * n))
294                    })
295                    .filter(|(_, y)| y.is_finite())
296                    .collect()
297            }
298            // Counts, and log-scaled bins, by each bin's expected count at its center:
299            // on Log a position is the log of the value, and a center the geometric
300            // middle.
301            Some(_) => edges
302                .windows(2)
303                .zip(&expected)
304                .map(|(edge, count)| {
305                    let center = if log {
306                        (edge[0] * edge[1]).sqrt().ln()
307                    } else {
308                        (edge[0] + edge[1]) / 2.0
309                    };
310                    (center, height(*count))
311                })
312                .collect(),
313            None => Vec::new(),
314        };
315        Histogram { counts, top, curve }
316    }
317}
318
319#[derive(Clone)]
320pub struct DistributionCharacteristics {
321    pub shapiro_wilk_stat: Option<f64>,
322    pub shapiro_wilk_pvalue: Option<f64>,
323    pub skewness: f64,
324    pub kurtosis: f64,
325    pub mean: f64,
326    pub median: f64,
327    pub std_dev: f64,
328    pub coefficient_of_variation: f64,
329}
330
331#[derive(Clone)]
332pub struct OutlierAnalysis {
333    pub total_count: usize,
334    pub percentage: f64,
335    pub iqr_count: usize,
336    pub zscore_count: usize,
337}
338
339#[derive(Clone)]
340pub struct PercentileBreakdown {
341    pub p25: f64,
342    pub p50: f64,
343    pub p75: f64,
344    pub p99: f64,
345}
346
347/// Which coefficient the correlation matrix shows.
348#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
349pub enum CorrelationMethod {
350    /// Pearson's r: how close the pairs fall to a line.
351    #[default]
352    Pearson,
353    /// Spearman's ρ: Pearson's r of the pairs' ranks, for any monotone relation.
354    Spearman,
355}
356
357impl CorrelationMethod {
358    pub fn toggled(self) -> Self {
359        match self {
360            Self::Pearson => Self::Spearman,
361            Self::Spearman => Self::Pearson,
362        }
363    }
364}
365
366// Correlation matrix structures
367#[derive(Clone)]
368pub struct CorrelationMatrix {
369    pub columns: Vec<String>,            // Numeric column names
370    pub correlations: Vec<Vec<f64>>,     // Square matrix of Pearson correlations
371    pub p_values: Option<Vec<Vec<f64>>>, // Statistical significance (optional)
372    pub sample_sizes: Vec<Vec<usize>>,   // Sample size for each pair
373    /// Spearman's ρ for each pair, over the same pairs as Pearson's r; `None` when
374    /// the rows read hold more values than [`RANK_VALUES`].
375    pub rank_correlations: Option<Vec<Vec<f64>>>,
376    pub rank_p_values: Option<Vec<Vec<f64>>>,
377}
378
379/// The most values Spearman's ρ ranks: eight bytes each, so 512 MiB beside the rows
380/// read. A sample's 100,000 rows rank up to 671 columns; a read of every row of a
381/// large table can be past it, and the matrix then has Pearson's r only.
382pub const RANK_VALUES: usize = 64 * 1024 * 1024;
383
384impl CorrelationMatrix {
385    /// The pair's coefficient by `method`; NaN where there is none.
386    pub fn coefficient(&self, method: CorrelationMethod, row: usize, col: usize) -> f64 {
387        let matrix = match method {
388            CorrelationMethod::Pearson => Some(&self.correlations),
389            CorrelationMethod::Spearman => self.rank_correlations.as_ref(),
390        };
391        matrix
392            .and_then(|m| m.get(row))
393            .and_then(|r| r.get(col))
394            .copied()
395            .unwrap_or(f64::NAN)
396    }
397
398    /// The pair's p-value by `method`, when the matrix has them.
399    pub fn p_value(&self, method: CorrelationMethod, row: usize, col: usize) -> Option<f64> {
400        let matrix = match method {
401            CorrelationMethod::Pearson => self.p_values.as_ref(),
402            CorrelationMethod::Spearman => self.rank_p_values.as_ref(),
403        };
404        matrix.and_then(|m| m.get(row)?.get(col).copied())
405    }
406}
407
408#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash)]
409pub enum DistributionType {
410    #[default]
411    Normal,
412    LogNormal,
413    Uniform,
414    PowerLaw,
415    Exponential,
416    Beta,
417    Gamma,
418    ChiSquared,
419    StudentsT,
420    Poisson,
421    Bernoulli,
422    Binomial,
423    Geometric,
424    Weibull,
425    /// One value throughout: nothing to fit.
426    Constant,
427    /// Every candidate was rejected: shown as "No clear fit".
428    Unknown,
429}
430
431impl std::fmt::Display for DistributionType {
432    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
433        match self {
434            DistributionType::Normal => write!(f, "Normal"),
435            DistributionType::LogNormal => write!(f, "Log-Normal"),
436            DistributionType::Uniform => write!(f, "Uniform"),
437            DistributionType::PowerLaw => write!(f, "Power Law"),
438            DistributionType::Exponential => write!(f, "Exponential"),
439            DistributionType::Beta => write!(f, "Beta"),
440            DistributionType::Gamma => write!(f, "Gamma"),
441            DistributionType::ChiSquared => write!(f, "Chi-Squared"),
442            DistributionType::StudentsT => write!(f, "Student's t"),
443            DistributionType::Poisson => write!(f, "Poisson"),
444            DistributionType::Bernoulli => write!(f, "Bernoulli"),
445            DistributionType::Binomial => write!(f, "Binomial"),
446            DistributionType::Geometric => write!(f, "Geometric"),
447            DistributionType::Weibull => write!(f, "Weibull"),
448            DistributionType::Constant => write!(f, "Constant"),
449            DistributionType::Unknown => write!(f, "No clear fit"),
450        }
451    }
452}
453
454#[derive(Clone)]
455pub struct AnalysisResults {
456    pub column_statistics: Vec<ColumnStatistics>,
457    pub total_rows: usize,
458    pub sample_size: Option<usize>,
459    /// Rows an equal-per-value sample kept of each value. See [`crate::analysis::sampling::PerValue`].
460    pub per_value: Option<usize>,
461    pub correlation_matrix: Option<CorrelationMatrix>,
462    pub distribution_analyses: Vec<DistributionAnalysis>,
463}
464
465#[derive(Debug, Clone, Copy)]
466pub struct ComputeOptions {
467    /// Fit distributions to the numeric columns. Analyses need this and
468    /// `include_distribution_analyses` both.
469    pub include_distribution_info: bool,
470    pub include_distribution_analyses: bool,
471    pub include_correlation_matrix: bool,
472    pub include_skewness_kurtosis_outliers: bool,
473    /// When true, use Polars streaming engine for LazyFrame collect when the streaming feature is enabled.
474    pub polars_streaming: bool,
475}
476
477impl Default for ComputeOptions {
478    fn default() -> Self {
479        Self {
480            include_distribution_info: false,
481            include_distribution_analyses: false,
482            include_correlation_matrix: false,
483            include_skewness_kurtosis_outliers: false,
484            polars_streaming: true,
485        }
486    }
487}
488
489/// Describe's temporal statistics for one collected column, through the same
490/// aggregation the lazy path runs.
491fn temporal_stats_of(series: &Series) -> Result<Option<TemporalStatistics>> {
492    if !is_temporal_type(series.dtype()) {
493        return Ok(None);
494    }
495    let frame = DataFrame::new_infer_height(vec![series.clone().into()])?;
496    let schema = frame.schema().clone();
497    let agg_df = frame
498        .lazy()
499        .select(build_describe_aggregation_exprs(&schema))
500        .collect()?;
501    Ok(parse_describe_agg_row(&agg_df, &schema)
502        .pop()
503        .and_then(|stats| stats.temporal_stats))
504}
505
506/// A value as the table writes it; `None` for a null.
507fn get_value_str(df: &DataFrame, col_name: &str, row: usize) -> Option<String> {
508    match df.column(col_name).ok()?.get(row).ok()? {
509        AnyValue::Null => None,
510        v => Some(crate::exact::str_value(&v).to_string()),
511    }
512}
513
514/// Statistics for a LazyFrame, over the rows a [`crate::analysis::sampling::Sample`] picks from
515/// `lf`, which is already cut to the sample's scope. Describe for every column;
516/// numeric percentiles, skewness, kurtosis, distribution fits and the correlation
517/// matrix as `options` asks.
518pub fn compute_statistics_for_sample(
519    lf: &LazyFrame,
520    sample: &crate::analysis::sampling::Sample,
521    known_total: Option<usize>,
522    options: ComputeOptions,
523) -> Result<AnalysisResults> {
524    let schema = lf.clone().collect_schema()?;
525    let use_streaming = options.polars_streaming;
526    let rows = crate::analysis::sampling::read(lf, sample, known_total, use_streaming)?;
527    let total_rows = rows.total_rows;
528    let actual_sample_size = rows.sample_size;
529    let per_value = rows.per_value.as_ref().map(|per_value| per_value.kept);
530    let df = rows.df;
531
532    let mut column_statistics = Vec::new();
533    let mut distribution_analyses = Vec::new();
534
535    for (name, dtype) in schema.iter() {
536        let col = df.column(name)?;
537        let series = col.as_materialized_series();
538        let count = series.len();
539        let null_count = series.null_count();
540
541        let numeric = if is_numeric_type(dtype) {
542            Some(NumericColumn::of(series)?)
543        } else {
544            None
545        };
546        let numeric_stats = numeric
547            .as_ref()
548            .map(|column| compute_numeric_stats(column, options.include_skewness_kurtosis_outliers))
549            .transpose()?;
550
551        let categorical_stats = if is_categorical_type(dtype) {
552            Some(compute_categorical_stats(series)?)
553        } else {
554            None
555        };
556
557        if options.include_distribution_info
558            && options.include_distribution_analyses
559            && null_count < count
560            && let (Some(column), Some(stats)) = (&numeric, &numeric_stats)
561        {
562            distribution_analyses.push(distribution_analysis(
563                name,
564                column,
565                stats,
566                actual_sample_size.unwrap_or(count),
567            ));
568        }
569
570        column_statistics.push(ColumnStatistics {
571            name: name.to_string(),
572            count,
573            null_count,
574            numeric_stats,
575            categorical_stats,
576            temporal_stats: temporal_stats_of(series)?,
577        });
578    }
579
580    let correlation_matrix = if options.include_correlation_matrix {
581        compute_correlation_matrix(&df).ok()
582    } else {
583        None
584    };
585
586    Ok(AnalysisResults {
587        column_statistics,
588        total_rows,
589        sample_size: actual_sample_size,
590        per_value,
591        correlation_matrix,
592        distribution_analyses,
593    })
594}
595
596/// Builds describe-only AnalysisResults from a list of column statistics.
597///
598/// Used when completing chunked describe; correlation and distribution analyses stay empty/None.
599pub fn analysis_results_from_describe(
600    column_statistics: Vec<ColumnStatistics>,
601    total_rows: usize,
602    sample_size: Option<usize>,
603) -> AnalysisResults {
604    AnalysisResults {
605        column_statistics,
606        total_rows,
607        sample_size,
608        per_value: None,
609        correlation_matrix: None,
610        distribution_analyses: Vec::new(),
611    }
612}
613
614/// Builds aggregation expressions for describe (count, null_count, mean, std, percentiles, min, max).
615/// Used so we can run a single collect on a LazyFrame without materializing all rows.
616fn build_describe_aggregation_exprs(schema: &Schema) -> Vec<Expr> {
617    let mut exprs = Vec::new();
618    for (name, dtype) in schema.iter() {
619        let name = name.as_str();
620        let prefix = format!("{}::", name);
621        exprs.push(col(name).count().alias(format!("{}count", prefix)));
622        exprs.push(
623            col(name)
624                .null_count()
625                .alias(format!("{}null_count", prefix)),
626        );
627        if is_numeric_type(dtype) {
628            let c = col(name).cast(DataType::Float64);
629            exprs.push(c.clone().mean().alias(format!("{}mean", prefix)));
630            exprs.push(c.clone().std(1).alias(format!("{}std", prefix)));
631            exprs.push(c.clone().min().alias(format!("{}min", prefix)));
632            // Without NaN, which sorts above every number and would be the upper
633            // quantiles of a column holding enough of it.
634            let numbers = c.clone().drop_nans();
635            exprs.push(
636                numbers
637                    .clone()
638                    .quantile(lit(0.25), QuantileMethod::Nearest)
639                    .alias(format!("{}q25", prefix)),
640            );
641            exprs.push(
642                numbers
643                    .clone()
644                    .quantile(lit(0.5), QuantileMethod::Nearest)
645                    .alias(format!("{}median", prefix)),
646            );
647            exprs.push(
648                numbers
649                    .clone()
650                    .quantile(lit(0.75), QuantileMethod::Nearest)
651                    .alias(format!("{}q75", prefix)),
652            );
653            exprs.push(c.max().alias(format!("{}max", prefix)));
654        } else if is_categorical_type(dtype) {
655            exprs.push(col(name).min().alias(format!("{}min", prefix)));
656            exprs.push(col(name).max().alias(format!("{}max", prefix)));
657        } else if is_temporal_type(dtype) {
658            // Mean and quantiles run on the physical integers and are cast back, so
659            // each is a value of the column's own type, as in Polars' describe.
660            let physical = dtype.to_physical();
661            let back = |e: Expr| e.cast(physical.clone()).cast(dtype.clone());
662            let p = col(name).cast(physical.clone());
663            exprs.push(
664                back(p.clone().cast(DataType::Float64).mean()).alias(format!("{}mean", prefix)),
665            );
666            exprs.push(col(name).min().alias(format!("{}min", prefix)));
667            for (q, stat) in [(0.25, "q25"), (0.5, "median"), (0.75, "q75")] {
668                exprs.push(
669                    back(p.clone().quantile(lit(q), QuantileMethod::Nearest))
670                        .alias(format!("{}{}", prefix, stat)),
671                );
672            }
673            exprs.push(col(name).max().alias(format!("{}max", prefix)));
674        }
675    }
676    exprs
677}
678
679/// Parses the single-row aggregation result from describe into column statistics.
680fn parse_describe_agg_row(agg_df: &DataFrame, schema: &Schema) -> Vec<ColumnStatistics> {
681    let row = 0usize;
682    let mut column_statistics = Vec::with_capacity(schema.len());
683    for (name, dtype) in schema.iter() {
684        let name_str = name.as_str();
685        let prefix = format!("{}::", name_str);
686        let count: usize = agg_df
687            .column(&format!("{}count", prefix))
688            .ok()
689            .map(|s| match s.get(row) {
690                Ok(AnyValue::UInt32(x)) => x as usize,
691                _ => 0,
692            })
693            .unwrap_or(0);
694        let null_count: usize = agg_df
695            .column(&format!("{}null_count", prefix))
696            .ok()
697            .map(|s| match s.get(row) {
698                Ok(AnyValue::UInt32(x)) => x as usize,
699                _ => 0,
700            })
701            .unwrap_or(0);
702        let numeric_stats = if is_numeric_type(dtype) {
703            let mean = get_f64(agg_df, &format!("{}mean", prefix), row);
704            let std = get_f64(agg_df, &format!("{}std", prefix), row);
705            let min = get_f64(agg_df, &format!("{}min", prefix), row);
706            let q25 = get_f64(agg_df, &format!("{}q25", prefix), row);
707            let median = get_f64(agg_df, &format!("{}median", prefix), row);
708            let q75 = get_f64(agg_df, &format!("{}q75", prefix), row);
709            let max = get_f64(agg_df, &format!("{}max", prefix), row);
710            let mut percentiles = HashMap::new();
711            percentiles.insert(25u8, q25);
712            percentiles.insert(50u8, median);
713            percentiles.insert(75u8, q75);
714            Some(NumericStatistics {
715                mean,
716                std,
717                min,
718                max,
719                median,
720                q25,
721                q75,
722                percentiles,
723                skewness: 0.0,
724                kurtosis: 3.0,
725            })
726        } else {
727            None
728        };
729        let categorical_stats = if is_categorical_type(dtype) {
730            let min = get_str(agg_df, &format!("{}min", prefix), row);
731            let max = get_str(agg_df, &format!("{}max", prefix), row);
732            Some(CategoricalStatistics { min, max })
733        } else {
734            None
735        };
736        let temporal_stats = is_temporal_type(dtype).then(|| {
737            let value = |stat: &str| get_value_str(agg_df, &format!("{}{}", prefix, stat), row);
738            TemporalStatistics {
739                mean: value("mean"),
740                min: value("min"),
741                q25: value("q25"),
742                median: value("median"),
743                q75: value("q75"),
744                max: value("max"),
745            }
746        });
747        column_statistics.push(ColumnStatistics {
748            name: name_str.to_string(),
749            count,
750            null_count,
751            numeric_stats,
752            categorical_stats,
753            temporal_stats,
754        });
755    }
756    column_statistics
757}
758
759/// Describe statistics for a frame. With `sample_size`, a table with more rows than
760/// that is described from a sample (see [`crate::analysis::sampling::analysis_rows`]); without it, every row is
761/// aggregated in one streaming pass, never held. `known_total` saves a count.
762pub fn compute_describe_from_lazy(
763    lf: &LazyFrame,
764    known_total: Option<usize>,
765    sample: &crate::analysis::sampling::Sample,
766    polars_streaming: bool,
767) -> Result<AnalysisResults> {
768    let schema = lf.clone().collect_schema()?;
769    if sample.method != crate::analysis::sampling::SampleMethod::EveryRow {
770        let rows = crate::analysis::sampling::read(lf, sample, known_total, polars_streaming)?;
771        let mut results = compute_describe_single_aggregation(
772            &rows.df,
773            &schema,
774            rows.total_rows,
775            rows.sample_size,
776            polars_streaming,
777        )?;
778        results.per_value = rows.per_value.map(|per_value| per_value.kept);
779        return Ok(results);
780    }
781    let total_rows = match known_total {
782        Some(total) => total,
783        None => crate::analysis::sampling::count_rows(lf, polars_streaming)?,
784    };
785    let exprs = build_describe_aggregation_exprs(&schema);
786    let agg_df = collect_lazy(lf.clone().select(exprs), polars_streaming).map_err(Report::from)?;
787    let column_statistics = parse_describe_agg_row(&agg_df, &schema);
788    Ok(analysis_results_from_describe(
789        column_statistics,
790        total_rows,
791        None,
792    ))
793}
794
795/// Computes describe statistics in a single aggregation pass over the DataFrame.
796/// Uses one collect() with aggregated expressions for all columns (count, null_count, mean, std, min, percentiles, max).
797pub fn compute_describe_single_aggregation(
798    df: &DataFrame,
799    schema: &Schema,
800    total_rows: usize,
801    sample_size: Option<usize>,
802    polars_streaming: bool,
803) -> Result<AnalysisResults> {
804    let exprs = build_describe_aggregation_exprs(schema);
805    let agg_df =
806        collect_lazy(df.clone().lazy().select(exprs), polars_streaming).map_err(Report::from)?;
807    let column_statistics = parse_describe_agg_row(&agg_df, schema);
808    Ok(analysis_results_from_describe(
809        column_statistics,
810        total_rows,
811        sample_size,
812    ))
813}
814
815fn get_f64(df: &DataFrame, col_name: &str, row: usize) -> f64 {
816    df.column(col_name)
817        .ok()
818        .and_then(|s| {
819            let v = s.get(row).ok()?;
820            match v {
821                AnyValue::Float64(x) => Some(x),
822                AnyValue::Float32(x) => Some(x as f64),
823                AnyValue::Int32(x) => Some(x as f64),
824                AnyValue::Int64(x) => Some(x as f64),
825                AnyValue::UInt32(x) => Some(x as f64),
826                AnyValue::Null => Some(f64::NAN),
827                _ => None,
828            }
829        })
830        .unwrap_or(f64::NAN)
831}
832
833fn get_str(df: &DataFrame, col_name: &str, row: usize) -> Option<String> {
834    df.column(col_name).ok().and_then(|s| {
835        s.get(row)
836            .ok()
837            .map(|v| crate::exact::str_value(&v).to_string())
838    })
839}
840
841/// Uses Polars' definition so Int128, UInt128, Decimal, and future numeric types are included.
842fn is_numeric_type(dtype: &DataType) -> bool {
843    dtype.is_numeric()
844}
845
846fn is_categorical_type(dtype: &DataType) -> bool {
847    matches!(dtype, DataType::String | DataType::Categorical(..))
848}
849
850fn is_temporal_type(dtype: &DataType) -> bool {
851    matches!(
852        dtype,
853        DataType::Date | DataType::Datetime(..) | DataType::Time | DataType::Duration(_)
854    )
855}
856
857/// A numeric column cast to `f64` once, for everything the analysis computes from it.
858struct NumericColumn {
859    floats: Float64Chunked,
860    /// Every finite value: nulls, NaN and infinities left out. No distribution has
861    /// NaN, and one is enough to leave a sort by `partial_cmp` out of order.
862    finite: Vec<f64>,
863}
864
865impl NumericColumn {
866    fn of(series: &Series) -> Result<Self> {
867        let floats = series.cast(&DataType::Float64)?.f64()?.clone();
868        let finite = floats.iter().flatten().filter(|v| v.is_finite()).collect();
869        Ok(Self { floats, finite })
870    }
871
872    /// Up to ten thousand finite values, every k-th row's rather than the first ten
873    /// thousand: a sample is spread across the table, and its head is one stretch of it.
874    fn spread(&self) -> Vec<f64> {
875        const MAX_VALUES: usize = 10_000;
876        let step = self.floats.len().div_ceil(MAX_VALUES).max(1);
877        self.floats
878            .iter()
879            .step_by(step)
880            .flatten()
881            .filter(|v| v.is_finite())
882            .collect()
883    }
884}
885
886fn compute_numeric_stats(
887    column: &NumericColumn,
888    include_advanced: bool,
889) -> Result<NumericStatistics> {
890    // Cast and aggregate as Describe does (`build_describe_aggregation_exprs`), so a
891    // sample's median is one number wherever it is shown.
892    let floats = column.floats.clone().into_series();
893    let mean = floats.mean().unwrap_or(f64::NAN);
894    let std = floats.std(1).unwrap_or(f64::NAN);
895    let min = floats.min::<f64>()?.unwrap_or(f64::NAN);
896    let max = floats.max::<f64>()?.unwrap_or(f64::NAN);
897
898    // NaN sorts above every number, so a column with some would have them as its
899    // upper percentiles and fences; Describe leaves them out too. One sort for all.
900    let floats = floats.f64()?;
901    let numbers = floats.filter(&floats.is_not_nan())?;
902    const PERCENTILES: [u8; 7] = [1, 5, 25, 50, 75, 95, 99];
903    let quantiles = PERCENTILES.map(|p| f64::from(p) / 100.0);
904    let values = numbers.quantiles(&quantiles, QuantileMethod::Nearest)?;
905    let percentiles: HashMap<u8, f64> = PERCENTILES
906        .into_iter()
907        .zip(values)
908        .map(|(p, value)| (p, value.unwrap_or(f64::NAN)))
909        .collect();
910
911    let median = percentiles[&50];
912    let q25 = percentiles[&25];
913    let q75 = percentiles[&75];
914
915    let (skewness, kurtosis) = if include_advanced {
916        skewness_and_kurtosis(&column.finite)
917    } else {
918        (0.0, 3.0)
919    };
920
921    Ok(NumericStatistics {
922        mean,
923        std,
924        min,
925        max,
926        median,
927        q25,
928        q75,
929        percentiles,
930        skewness,
931        kurtosis,
932    })
933}
934
935/// Mean and sample standard deviation (ddof 1) of a set of values.
936fn mean_and_std(values: &[f64]) -> (f64, f64) {
937    let n = values.len() as f64;
938    if values.len() < 2 {
939        return (values.first().copied().unwrap_or(f64::NAN), f64::NAN);
940    }
941    let mean = values.iter().sum::<f64>() / n;
942    let sum_squares: f64 = values.iter().map(|v| (v - mean).powi(2)).sum();
943    (mean, (sum_squares / (n - 1.0)).sqrt())
944}
945
946/// Skewness and kurtosis of one set of values, `n` being their count: the
947/// bias-corrected forms of Polars' `skew(bias=False)` and `kurtosis(bias=False)`,
948/// kurtosis on the scale where a normal is 3. One value throughout, or too few to
949/// say, is 0 and 3.
950fn skewness_and_kurtosis(values: &[f64]) -> (f64, f64) {
951    let count = values.len();
952    if count < 3 || values.iter().all(|v| *v == values[0]) {
953        return (0.0, 3.0);
954    }
955    let n = count as f64;
956    let (mean, std) = mean_and_std(values);
957    let (mut cubes, mut fourths) = (0.0, 0.0);
958    for v in values {
959        let z = (v - mean) / std;
960        let z2 = z * z;
961        cubes += z2 * z;
962        fourths += z2 * z2;
963    }
964    let skewness = n / ((n - 1.0) * (n - 2.0)) * cubes;
965    if count < 4 {
966        return (skewness, 3.0);
967    }
968    let excess = n * (n + 1.0) / ((n - 1.0) * (n - 2.0) * (n - 3.0)) * fourths
969        - 3.0 * (n - 1.0) * (n - 1.0) / ((n - 2.0) * (n - 3.0));
970    (skewness, excess + 3.0)
971}
972
973/// Where a value falls against the IQR fences and three standard deviations.
974struct OutlierTest {
975    lower_fence: f64,
976    upper_fence: f64,
977    mean: f64,
978    std: f64,
979}
980
981impl OutlierTest {
982    /// Fences from the quartiles; the z-score from the values' own mean and std, so
983    /// a NaN elsewhere in the column cannot void it.
984    fn new(values: &[f64], q25: f64, q75: f64) -> Option<Self> {
985        let (mean, std) = mean_and_std(values);
986        if q25.is_nan() || q75.is_nan() || std.is_nan() || std == 0.0 {
987            return None;
988        }
989        let iqr = q75 - q25;
990        Some(Self {
991            lower_fence: q25 - 1.5 * iqr,
992            upper_fence: q75 + 1.5 * iqr,
993            mean,
994            std,
995        })
996    }
997
998    fn beyond_fences(&self, value: f64) -> bool {
999        value < self.lower_fence || value > self.upper_fence
1000    }
1001
1002    fn z_score(&self, value: f64) -> f64 {
1003        (value - self.mean).abs() / self.std
1004    }
1005}
1006
1007const Z_THRESHOLD: f64 = 3.0;
1008
1009fn compute_categorical_stats(series: &Series) -> Result<CategoricalStatistics> {
1010    let min = if let Ok(str_series) = series.str() {
1011        let mut min_val: Option<String> = None;
1012        for s in str_series.iter().flatten() {
1013            let s_str = s.to_string();
1014            min_val = match min_val {
1015                None => Some(s_str.clone()),
1016                Some(ref current) if s_str < *current => Some(s_str),
1017                Some(current) => Some(current),
1018            };
1019        }
1020        min_val
1021    } else {
1022        None
1023    };
1024
1025    let max = if let Ok(str_series) = series.str() {
1026        let mut max_val: Option<String> = None;
1027        for s in str_series.iter().flatten() {
1028            let s_str = s.to_string();
1029            max_val = match max_val {
1030                None => Some(s_str.clone()),
1031                Some(ref current) if s_str > *current => Some(s_str),
1032                Some(current) => Some(current),
1033            };
1034        }
1035        max_val
1036    } else {
1037        None
1038    };
1039
1040    Ok(CategoricalStatistics { min, max })
1041}
1042
1043/// Seeds the fit tests' simulations, so the same values get the same p-values.
1044const FIT_SEED: u64 = 0x5eed_d157;
1045
1046/// The family a column's values follow, its p-value and every family's fit. `rows`
1047/// is how many rows the column was read from.
1048struct ColumnFit {
1049    distribution_type: DistributionType,
1050    confidence: f64,
1051    fits: Vec<(
1052        DistributionType,
1053        crate::analysis::distribution_fit::FitOutcome,
1054    )>,
1055}
1056
1057fn infer_distribution(values: &[f64], rows: usize) -> ColumnFit {
1058    let unknown = ColumnFit {
1059        distribution_type: DistributionType::Unknown,
1060        confidence: 0.0,
1061        fits: Vec::new(),
1062    };
1063    if rows < 3 || values.is_empty() {
1064        return unknown;
1065    }
1066
1067    let mean: f64 = values.iter().sum::<f64>() / values.len() as f64;
1068    let variance: f64 =
1069        values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (values.len() - 1) as f64;
1070    let std = variance.sqrt();
1071
1072    // One value throughout fits every distribution's degenerate case and none of
1073    // them usefully; a year column in a partitioned table is the usual one.
1074    if std == 0.0 {
1075        return ColumnFit {
1076            distribution_type: DistributionType::Constant,
1077            confidence: 1.0,
1078            fits: Vec::new(),
1079        };
1080    }
1081
1082    // Counts are described by a count distribution when one holds.
1083    let counts = values.iter().all(|v| *v >= 0.0 && *v == v.floor());
1084    let fits = crate::analysis::distribution_fit::test_all(values, FIT_SEED);
1085    let distribution_type = crate::analysis::distribution_fit::select(&fits, counts);
1086    // The figure beside the name is that family's p-value; with no clear fit, the best
1087    // any family managed, so the table can say how far from fitting it was.
1088    let confidence = fits
1089        .iter()
1090        .find(|(family, _)| *family == distribution_type)
1091        .and_then(|(_, outcome)| outcome.p_value())
1092        .or_else(|| {
1093            fits.iter()
1094                .filter_map(|(_, outcome)| outcome.p_value())
1095                .max_by(f64::total_cmp)
1096        })
1097        .unwrap_or(0.0);
1098    ColumnFit {
1099        distribution_type,
1100        confidence,
1101        fits,
1102    }
1103}
1104
1105/// The Shapiro-Francia statistic of `sorted` against normal scores at Blom's plotting
1106/// positions `(i + 1 - 3/8) / (n + 1/4)`, and its p-value.
1107fn approximate_shapiro_wilk(sorted: &[f64]) -> (Option<f64>, Option<f64>) {
1108    let n = sorted.len();
1109    if n < 3 {
1110        return (None, None);
1111    }
1112
1113    let mean: f64 = sorted.iter().sum::<f64>() / n as f64;
1114    let variance: f64 = sorted.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (n - 1) as f64;
1115    let std = variance.sqrt();
1116
1117    if std == 0.0 {
1118        return (None, None);
1119    }
1120
1121    let mut sum_expected_sq = 0.0;
1122    let mut sum_data_sq = 0.0;
1123    let mut sum_product = 0.0;
1124
1125    for (i, &value) in sorted.iter().enumerate() {
1126        let p = (i as f64 + 1.0 - 0.375) / (n as f64 + 0.25);
1127        let expected_quantile = crate::analysis::distribution_fit::normal_quantile(p);
1128        let standardized_value = (value - mean) / std;
1129
1130        sum_expected_sq += expected_quantile * expected_quantile;
1131        sum_data_sq += standardized_value * standardized_value;
1132        sum_product += expected_quantile * standardized_value;
1133    }
1134
1135    let sw_stat = if sum_expected_sq > 0.0 && sum_data_sq > 0.0 {
1136        (sum_product * sum_product) / (sum_expected_sq * sum_data_sq)
1137    } else {
1138        0.0
1139    };
1140
1141    let sw_stat = sw_stat.clamp(0.0, 1.0);
1142    (Some(sw_stat), shapiro_francia_pvalue(sw_stat, n))
1143}
1144
1145/// The p-value of the Shapiro-Francia W' computed above (squared correlation of sorted
1146/// values with normal scores), by Royston's (1993) approximation for 5 to 5,000 values;
1147/// `None` outside.
1148fn shapiro_francia_pvalue(w: f64, n: usize) -> Option<f64> {
1149    if !(5..=5_000).contains(&n) {
1150        return None;
1151    }
1152    if w >= 1.0 {
1153        return Some(1.0);
1154    }
1155    let u = (n as f64).ln();
1156    let v = u.ln();
1157    let mu = -1.2725 + 1.0521 * (v - u);
1158    let sigma = 1.0308 - 0.26758 * (v + 2.0 / u);
1159    let z = ((1.0 - w).ln() - mu) / sigma;
1160    Some((1.0 - crate::analysis::distribution_fit::normal_cdf(z)).clamp(0.0, 1.0))
1161}
1162
1163/// One numeric column's distribution: its fits, normality, outliers and the sorted
1164/// values its Q-Q plot draws. `rows` is how many rows the column was read from.
1165fn distribution_analysis(
1166    column_name: &str,
1167    column: &NumericColumn,
1168    numeric_stats: &NumericStatistics,
1169    rows: usize,
1170) -> DistributionAnalysis {
1171    let spread = column.spread();
1172    let fit = infer_distribution(&spread, rows);
1173    // At most five thousand, spread across the rows: the head of a table sorted by
1174    // date is its first few years.
1175    const MAX_VALUES: usize = 5_000;
1176    let step = spread.len().div_ceil(MAX_VALUES).max(1);
1177    let mut sorted_sample_values: Vec<f64> = spread.into_iter().step_by(step).collect();
1178    sorted_sample_values.sort_by(f64::total_cmp);
1179
1180    let (sw_stat, sw_pvalue) = approximate_shapiro_wilk(&sorted_sample_values);
1181    let coefficient_of_variation = if numeric_stats.mean != 0.0 {
1182        numeric_stats.std / numeric_stats.mean.abs()
1183    } else {
1184        0.0
1185    };
1186
1187    let characteristics = DistributionCharacteristics {
1188        shapiro_wilk_stat: sw_stat,
1189        shapiro_wilk_pvalue: sw_pvalue,
1190        skewness: numeric_stats.skewness,
1191        kurtosis: numeric_stats.kurtosis,
1192        mean: numeric_stats.mean,
1193        median: numeric_stats.median,
1194        std_dev: numeric_stats.std,
1195        coefficient_of_variation,
1196    };
1197
1198    let qq = fit
1199        .fits
1200        .iter()
1201        .filter_map(|(family, outcome)| {
1202            let test = outcome.test()?;
1203            Some((
1204                *family,
1205                crate::analysis::distribution_fit::qq_quantiles(
1206                    &test.fitted,
1207                    sorted_sample_values.len(),
1208                ),
1209            ))
1210        })
1211        .collect();
1212
1213    let outliers = compute_outlier_analysis(&column.finite, numeric_stats);
1214
1215    let percentiles = PercentileBreakdown {
1216        p25: numeric_stats.q25,
1217        p50: numeric_stats.median,
1218        p75: numeric_stats.q75,
1219        p99: numeric_stats
1220            .percentiles
1221            .get(&99)
1222            .copied()
1223            .unwrap_or(f64::NAN),
1224    };
1225
1226    DistributionAnalysis {
1227        column_name: column_name.to_string(),
1228        distribution_type: fit.distribution_type,
1229        confidence: fit.confidence,
1230        characteristics,
1231        outliers,
1232        percentiles,
1233        sorted_sample_values,
1234        fits: fit.fits,
1235        qq,
1236        histogram: HistogramCache::default(),
1237    }
1238}
1239
1240/// Outliers among every finite value of the column, counted by each test and as a
1241/// share of the values.
1242fn compute_outlier_analysis(values: &[f64], numeric_stats: &NumericStatistics) -> OutlierAnalysis {
1243    let mut analysis = OutlierAnalysis {
1244        total_count: 0,
1245        percentage: 0.0,
1246        iqr_count: 0,
1247        zscore_count: 0,
1248    };
1249    let Some(test) = OutlierTest::new(values, numeric_stats.q25, numeric_stats.q75) else {
1250        return analysis;
1251    };
1252
1253    for &value in values {
1254        let beyond_fences = test.beyond_fences(value);
1255        let beyond_z = test.z_score(value) > Z_THRESHOLD;
1256        if !beyond_fences && !beyond_z {
1257            continue;
1258        }
1259        analysis.total_count += 1;
1260        analysis.iqr_count += usize::from(beyond_fences);
1261        analysis.zscore_count += usize::from(beyond_z);
1262    }
1263
1264    analysis.percentage = analysis.total_count as f64 / values.len() as f64 * 100.0;
1265    analysis
1266}
1267
1268/// Pairwise Pearson correlations of all numeric columns, with p-values and sample
1269/// sizes; needs at least two numeric columns.
1270pub fn compute_correlation_matrix(df: &DataFrame) -> Result<CorrelationMatrix> {
1271    let columns = df
1272        .schema()
1273        .iter()
1274        .filter(|(_, dtype)| is_numeric_type(dtype))
1275        .count()
1276        .max(1);
1277    let band = CORRELATION_SCRATCH_BYTES / (columns * std::mem::size_of::<f64>());
1278    correlation_matrix_in_bands(df, band.max(MIN_BAND_ROWS))
1279}
1280
1281/// Bytes of converted values a correlation matrix holds at once, beyond the sample
1282/// itself and the matrix: one band of the 100,000-row sample is the whole of it up
1283/// to 83 columns.
1284const CORRELATION_SCRATCH_BYTES: usize = 64 * 1024 * 1024;
1285
1286/// The fewest rows in a band, so a schema of thousands of columns is not read a few
1287/// rows at a time.
1288const MIN_BAND_ROWS: usize = 1024;
1289
1290/// Rows cast to floats at a time, so no cast is the size of a column.
1291const CAST_ROWS: usize = 16 * 1024;
1292
1293/// [`compute_correlation_matrix`] converting `band` rows of every column at a time.
1294fn correlation_matrix_in_bands(df: &DataFrame, band: usize) -> Result<CorrelationMatrix> {
1295    let schema = df.schema();
1296    let numeric_cols: Vec<String> = schema
1297        .iter()
1298        .filter(|(_, dtype)| is_numeric_type(dtype))
1299        .map(|(name, _)| name.to_string())
1300        .collect();
1301
1302    if numeric_cols.len() < 2 {
1303        return Err(color_eyre::eyre::eyre!(
1304            "Need at least 2 numeric columns for correlation matrix"
1305        ));
1306    }
1307
1308    let series = numeric_cols
1309        .iter()
1310        .map(|name| Ok(df.column(name)?.as_materialized_series()))
1311        .collect::<Result<Vec<_>>>()?;
1312
1313    let n = numeric_cols.len();
1314    let rows = df.height();
1315    let band = band.clamp(1, rows.max(1));
1316
1317    // Means first, then bands of rows of every column, less those means; each pair's sums
1318    // run across bands in row order, converting each column once and holding one band.
1319    let mut shifts = vec![Shift::default(); n];
1320    across_threads(
1321        series.iter().zip(shifts.iter_mut()).collect(),
1322        |(series, shift)| *shift = Shift::new(series),
1323    );
1324    let mut sums: Vec<Vec<PairSums>> = (0..n)
1325        .map(|i| vec![PairSums::default(); n - i - 1])
1326        .collect();
1327    let mut bands: Vec<Vec<f64>> = (0..n).map(|_| Vec::with_capacity(band)).collect();
1328    for start in (0..rows).step_by(band) {
1329        let within = start..(start + band).min(rows);
1330        across_threads(
1331            series.iter().zip(&shifts).zip(bands.iter_mut()).collect(),
1332            |((series, shift), values)| shift.fill(series, within.clone(), values),
1333        );
1334        // Every pair is one pass over two columns; fifty columns are 1,225 pairs, so
1335        // the rows of the matrix are shared out across threads, interleaved to even
1336        // the load.
1337        let (bands, shifts) = (&bands, &shifts);
1338        across_threads(sums.iter_mut().enumerate().collect(), |(i, row)| {
1339            for (k, sums) in row.iter_mut().enumerate() {
1340                let j = i + 1 + k;
1341                let both = shifts[i].complete && shifts[j].complete;
1342                sums.add_pairs(&bands[i], &bands[j], both);
1343            }
1344        });
1345    }
1346
1347    let mut correlations = vec![vec![1.0; n]; n];
1348    let mut p_values = vec![vec![0.0; n]; n];
1349    let mut sample_sizes = vec![vec![0; n]; n];
1350    for (i, row) in sums.iter().enumerate() {
1351        for (k, sums) in row.iter().enumerate() {
1352            let j = i + 1 + k;
1353            let sample_size = sums.count;
1354            sample_sizes[i][j] = sample_size;
1355            sample_sizes[j][i] = sample_size;
1356            // Fewer than three pairs say nothing.
1357            let correlation = if sample_size < 3 {
1358                f64::NAN
1359            } else {
1360                sums.correlation()
1361            };
1362            correlations[i][j] = correlation;
1363            correlations[j][i] = correlation;
1364            if !correlation.is_nan() {
1365                let p_value = compute_correlation_p_value(correlation, sample_size);
1366                p_values[i][j] = p_value;
1367                p_values[j][i] = p_value;
1368            }
1369        }
1370    }
1371
1372    let ranked = (rows.saturating_mul(n) <= RANK_VALUES)
1373        .then(|| rank_correlation_matrix(&series, &sample_sizes));
1374    let (rank_correlations, rank_p_values) = ranked.unzip();
1375    Ok(CorrelationMatrix {
1376        columns: numeric_cols,
1377        correlations,
1378        p_values: Some(p_values),
1379        sample_sizes,
1380        rank_correlations,
1381        rank_p_values,
1382    })
1383}
1384
1385/// Marks a row with no finite value in [`Ranked::ranks`].
1386const NO_RANK: u32 = u32::MAX;
1387
1388/// One column's ranks, for Spearman's ρ.
1389struct Ranked {
1390    /// Twice each finite value's average rank (from 1) among the column's finite
1391    /// values, so a tie's half rank stays whole; [`NO_RANK`] where there is none.
1392    /// Equal values share a rank, so the ranks also tell ties apart.
1393    ranks: Vec<u32>,
1394    /// The rows with a finite value, in order of value.
1395    order: Vec<u32>,
1396    /// Every row has a finite value.
1397    complete: bool,
1398}
1399
1400impl Ranked {
1401    fn new(series: &Series) -> Option<Self> {
1402        let rows = series.len();
1403        // Doubled ranks reach twice the rows.
1404        if rows >= (NO_RANK / 2) as usize {
1405            return None;
1406        }
1407        let mut values = Vec::with_capacity(rows);
1408        let mut row = 0u32;
1409        for_each_float(series, 0..rows, |v| {
1410            if let Some(v) = v.filter(|v| v.is_finite()) {
1411                values.push((v, row));
1412            }
1413            row += 1;
1414        });
1415        values.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
1416        let mut ranks = vec![NO_RANK; rows];
1417        let mut start = 0;
1418        while start < values.len() {
1419            // `==` so that -0 and 0 tie, which total_cmp sorts side by side.
1420            let end = start
1421                + values[start..]
1422                    .iter()
1423                    .take_while(|(v, _)| *v == values[start].0)
1424                    .count();
1425            let doubled = (start + 1 + end) as u32;
1426            for &(_, row) in &values[start..end] {
1427                ranks[row as usize] = doubled;
1428            }
1429            start = end;
1430        }
1431        Some(Self {
1432            complete: values.len() == rows,
1433            order: values.into_iter().map(|(_, row)| row).collect(),
1434            ranks,
1435        })
1436    }
1437
1438    /// `out` becomes this column's doubled ranks among the rows where `other` also
1439    /// has a value, [`NO_RANK`] elsewhere; the number of those rows is returned.
1440    /// Walking the rows in order of value, a run of equal global ranks is a run of
1441    /// equal values. `kept` is scratch, reused from pair to pair.
1442    fn ranks_beside(&self, other: &Ranked, out: &mut Vec<u32>, kept: &mut Vec<u32>) -> usize {
1443        out.clear();
1444        out.resize(self.ranks.len(), NO_RANK);
1445        kept.clear();
1446        kept.extend(
1447            self.order
1448                .iter()
1449                .copied()
1450                .filter(|&row| other.ranks[row as usize] != NO_RANK),
1451        );
1452        let mut start = 0;
1453        while start < kept.len() {
1454            let tie = self.ranks[kept[start] as usize];
1455            let end = start
1456                + kept[start..]
1457                    .iter()
1458                    .take_while(|&&row| self.ranks[row as usize] == tie)
1459                    .count();
1460            let doubled = (start + 1 + end) as u32;
1461            for &row in &kept[start..end] {
1462                out[row as usize] = doubled;
1463            }
1464            start = end;
1465        }
1466        kept.len()
1467    }
1468}
1469
1470/// Spearman's ρ for every pair of `series` with p-values: Pearson's r of ranks, each
1471/// pair ranked over rows where both are finite. Columns with no missing values reuse
1472/// their own ranks. Holds about the sample's size in ranks and order.
1473fn rank_correlation_matrix(
1474    series: &[&Series],
1475    sample_sizes: &[Vec<usize>],
1476) -> (Vec<Vec<f64>>, Vec<Vec<f64>>) {
1477    let n = series.len();
1478    let mut ranked: Vec<Option<Ranked>> = (0..n).map(|_| None).collect();
1479    across_threads(
1480        series.iter().zip(ranked.iter_mut()).collect(),
1481        |(series, ranked)| *ranked = Ranked::new(series),
1482    );
1483    let mut rows: Vec<Vec<f64>> = (0..n).map(|i| vec![f64::NAN; n - i - 1]).collect();
1484    let ranked = &ranked;
1485    across_threads(rows.iter_mut().enumerate().collect(), |(i, row)| {
1486        // Ranks as doubled u32s, half the size of floats, held once per row of the
1487        // matrix rather than once per pair.
1488        let (mut a, mut b, mut kept) = (Vec::new(), Vec::new(), Vec::new());
1489        for (k, rho) in row.iter_mut().enumerate() {
1490            let (Some(x), Some(y)) = (&ranked[i], &ranked[i + 1 + k]) else {
1491                continue;
1492            };
1493            let mut sums = PairSums::default();
1494            if x.complete && y.complete {
1495                // Ranks centered on their mean, which is the row count plus one.
1496                let mean = (x.ranks.len() + 1) as f64;
1497                for (&rx, &ry) in x.ranks.iter().zip(&y.ranks) {
1498                    sums.add(rx as f64 - mean, ry as f64 - mean);
1499                }
1500            } else {
1501                let pairs = x.ranks_beside(y, &mut a, &mut kept);
1502                y.ranks_beside(x, &mut b, &mut kept);
1503                let mean = (pairs + 1) as f64;
1504                for (&rx, &ry) in a.iter().zip(&b) {
1505                    if rx != NO_RANK {
1506                        sums.add(rx as f64 - mean, ry as f64 - mean);
1507                    }
1508                }
1509            }
1510            if sums.count >= 3 {
1511                *rho = sums.correlation();
1512            }
1513        }
1514    });
1515    let mut rho = vec![vec![1.0; n]; n];
1516    let mut p_values = vec![vec![0.0; n]; n];
1517    for (i, row) in rows.iter().enumerate() {
1518        for (k, &r) in row.iter().enumerate() {
1519            let j = i + 1 + k;
1520            rho[i][j] = r;
1521            rho[j][i] = r;
1522            if !r.is_nan() {
1523                // The same t approximation as Pearson's, over the same pairs.
1524                let p = compute_correlation_p_value(r, sample_sizes[i][j]);
1525                p_values[i][j] = p;
1526                p_values[j][i] = p;
1527            }
1528        }
1529    }
1530    (rho, p_values)
1531}
1532
1533/// Runs `work` on every item, the items dealt out across threads in turn. A worker's
1534/// panic is raised again here rather than leaving its items undone.
1535fn across_threads<T: Send>(items: Vec<T>, work: impl Fn(T) + Sync) {
1536    let threads = std::thread::available_parallelism()
1537        .map_or(1, usize::from)
1538        .min(items.len())
1539        .max(1);
1540    let mut shares: Vec<Vec<T>> = (0..threads).map(|_| Vec::new()).collect();
1541    for (k, item) in items.into_iter().enumerate() {
1542        shares[k % threads].push(item);
1543    }
1544    let work = &work;
1545    std::thread::scope(|scope| {
1546        let handles: Vec<_> = shares
1547            .into_iter()
1548            .map(|share| scope.spawn(move || share.into_iter().for_each(work)))
1549            .collect();
1550        for handle in handles {
1551            if let Err(panic) = handle.join() {
1552                std::panic::resume_unwind(panic);
1553            }
1554        }
1555    });
1556}
1557
1558/// `len` values of `series` from `start`, as floats.
1559fn float_piece(series: &Series, start: usize, len: usize) -> Option<Float64Chunked> {
1560    let piece = series
1561        .slice(start as i64, len)
1562        .cast(&DataType::Float64)
1563        .ok()?;
1564    piece.f64().ok().cloned()
1565}
1566
1567/// The values of `series` in `rows` as floats, None where null, cast [`CAST_ROWS`]
1568/// at a time.
1569fn for_each_float(series: &Series, rows: Range<usize>, mut f: impl FnMut(Option<f64>)) {
1570    for start in rows.clone().step_by(CAST_ROWS) {
1571        let len = CAST_ROWS.min(rows.end - start);
1572        match float_piece(series, start, len) {
1573            Some(floats) => floats.iter().for_each(&mut f),
1574            None => (0..len).for_each(|_| f(None)),
1575        }
1576    }
1577}
1578
1579/// A numeric column's mean over its finite values, taken off each value before it
1580/// is correlated: centered, the one-pass sums below stay exact enough.
1581#[derive(Clone, Copy, Default)]
1582struct Shift {
1583    mean: f64,
1584    /// No value is missing, so every row pairs.
1585    complete: bool,
1586}
1587
1588impl Shift {
1589    fn new(series: &Series) -> Self {
1590        let (mut sum, mut count) = (0.0, 0usize);
1591        for_each_float(series, 0..series.len(), |v| {
1592            if let Some(v) = v.filter(|v| v.is_finite()) {
1593                sum += v;
1594                count += 1;
1595            }
1596        });
1597        Self {
1598            mean: if count > 0 { sum / count as f64 } else { 0.0 },
1599            complete: count == series.len(),
1600        }
1601    }
1602
1603    /// `values` becomes the column's `rows` less the mean, NaN where a value is null
1604    /// or not finite.
1605    fn fill(&self, series: &Series, rows: Range<usize>, values: &mut Vec<f64>) {
1606        values.clear();
1607        for_each_float(series, rows, |v| {
1608            values.push(
1609                v.filter(|v| v.is_finite())
1610                    .map_or(f64::NAN, |v| v - self.mean),
1611            );
1612        });
1613    }
1614}
1615
1616/// One pass of sums over paired values, each less a shift near its mean. The
1617/// spreads about the pairs' own means follow exactly whatever the shift; a shift
1618/// near the mean keeps them clear of rounding.
1619#[derive(Clone, Copy, Default)]
1620struct PairSums {
1621    count: usize,
1622    x: f64,
1623    y: f64,
1624    xx: f64,
1625    yy: f64,
1626    xy: f64,
1627}
1628
1629impl PairSums {
1630    fn add(&mut self, v1: f64, v2: f64) {
1631        self.count += 1;
1632        self.x += v1;
1633        self.y += v2;
1634        self.xx += v1 * v1;
1635        self.yy += v2 * v2;
1636        self.xy += v1 * v2;
1637    }
1638
1639    /// Adds the rows where both centered columns have a value: with `both` complete,
1640    /// every row.
1641    fn add_pairs(&mut self, a: &[f64], b: &[f64], both: bool) {
1642        // Summed in a copy, which the loop can keep in registers.
1643        let mut sums = *self;
1644        for (&v1, &v2) in a.iter().zip(b) {
1645            if !both && (v1.is_nan() || v2.is_nan()) {
1646                continue;
1647            }
1648            sums.add(v1, v2);
1649        }
1650        *self = sums;
1651    }
1652
1653    /// The sums of squares and of products about the pairs' means.
1654    fn spreads(&self) -> (f64, f64, f64) {
1655        let n = self.count as f64;
1656        (
1657            self.xx - self.x * self.x / n,
1658            self.yy - self.y * self.y / n,
1659            self.xy - self.x * self.y / n,
1660        )
1661    }
1662
1663    /// NaN for fewer than two pairs or a column of one value.
1664    fn correlation(&self) -> f64 {
1665        let (sxx, syy, sxy) = self.spreads();
1666        // Measured against the sums of squares: what a column with one value leaves
1667        // behind is rounding, not spread.
1668        if self.count < 2 || sxx <= self.xx * 1e-12 || syy <= self.yy * 1e-12 {
1669            return f64::NAN;
1670        }
1671        (sxy / (sxx * syy).sqrt()).clamp(-1.0, 1.0)
1672    }
1673}
1674
1675/// The two-sided p-value of Pearson's r over `n` pairs: Student's t with `n - 2`
1676/// degrees of freedom. With `t² = r²·df / (1 - r²)`, both tails together are
1677/// `I_x(df/2, 1/2)` at `x = df / (df + t²) = 1 - r²`, taken directly so that a small p
1678/// is not lost to `1 - cdf`.
1679fn compute_correlation_p_value(correlation: f64, n: usize) -> f64 {
1680    if n < 3 || correlation.is_nan() {
1681        return 1.0;
1682    }
1683    if correlation.abs() >= 1.0 {
1684        return 0.0;
1685    }
1686    let df = (n - 2) as f64;
1687    crate::analysis::distribution_fit::beta_inc(df / 2.0, 0.5, 1.0 - correlation * correlation)
1688        .clamp(0.0, 1.0)
1689}
1690
1691#[cfg(test)]
1692mod tests;
1693
1694#[cfg(test)]
1695mod normality_tests;
1696
1697#[cfg(test)]
1698pub(crate) mod describe_tests;
1699
1700#[cfg(all(test, feature = "streaming"))]
1701mod streaming_guard_tests;