Skip to main content

datafusion_arrowmetal/
probe.rs

1//! The group-count probe: how many groups a GROUP BY's keys hold, estimated on the CPU from a
2//! sample of the collected input, before anything is imported to the GPU.
3//!
4//! A port of the Polars engine's probe (`python/arrowmetal/polars_engine.py`, "The group-count
5//! probe"). The distinct key tuples of a fixed sample of n of the input's N rows are scaled to the
6//! input by the bias-corrected Chao1 estimator (Chao, Biometrics 2005),
7//!
8//! ```text
9//!     D = d + f1 (f1 - 1) / (2 (f2 + 1)),
10//! ```
11//!
12//! where d is the number of distinct tuples in the sample and f1, f2 the number seen exactly once
13//! and exactly twice, clipped to [d, N]. The sample is stratified: row `i * (N / n) + h(i)` for
14//! `i < n`, `h` a fixed-seed splitmix64 hash of `i`, so the same input always gets the same sample
15//! and the same estimate.
16//!
17//! The samples are 512, 2,048, 8,192, ... rows, at most 65,536 and a quarter of the input, and the
18//! probe stops at the first whose range settles the decision: the range is D with f2 moved by two
19//! of its Poisson standard deviations (and one or two more) either way, and the caller's
20//! `settled(lo, hi)` says whether every group count in it gets the same answer. An input of at most
21//! 4,096 rows is counted exactly.
22
23use std::time::{Duration, Instant};
24
25use arrow::array::{Array, AsArray};
26use arrow::datatypes::{
27    DataType, Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, UInt16Type, UInt32Type,
28    UInt64Type, UInt8Type,
29};
30use arrow::record_batch::RecordBatch;
31
32const FIRST: usize = 512;
33const GROWTH: usize = 4;
34const SAMPLE_MAX: usize = 65_536;
35const EXACT: usize = 4_096;
36const SEED: u64 = 0xA6_5EED;
37const SETTLED_F2: u64 = 8;
38/// What a null key hashes to (a null is a group of its own, as in SQL GROUP BY).
39const NULL_KEY: u64 = 0x6E75_6C6C_6B65_7931;
40
41/// One probe's answer.
42#[derive(Debug, Clone, Copy, PartialEq)]
43#[non_exhaustive]
44pub struct GroupEstimate {
45    /// The Chao1 estimate (or the exact count), in groups.
46    pub estimate: u64,
47    /// The low end of its range (equal to `estimate` when counted exactly).
48    pub low: u64,
49    /// The high end of its range.
50    pub high: u64,
51    /// Rows in the last sample (the whole input when counted exactly).
52    pub sample_rows: usize,
53    /// Rows of the input (the whole input, also when the sample was drawn from a part of it).
54    pub rows: usize,
55    /// True when every row was counted (inputs of at most 4,096 rows).
56    pub exact: bool,
57    /// Wall time of the probe.
58    pub time: Duration,
59}
60
61fn splitmix(mut z: u64) -> u64 {
62    z = z.wrapping_add(0x9E37_79B9_7F4A_7C15);
63    z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
64    z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
65    z ^ (z >> 31)
66}
67
68fn fnv1a(b: &[u8]) -> u64 {
69    let mut h: u64 = 0xCBF2_9CE4_8422_2325;
70    for &x in b {
71        h ^= x as u64;
72        h = h.wrapping_mul(0x0100_0000_01B3);
73    }
74    h
75}
76
77/// The sample's row positions: stratified, increasing, fixed for (rows, n).
78fn positions(rows: usize, n: usize) -> Vec<usize> {
79    let stride = (rows / n).max(1);
80    (0..n)
81        .map(|i| i * stride + (splitmix(i as u64 + SEED) % stride as u64) as usize)
82        .filter(|&p| p < rows)
83        .collect()
84}
85
86/// Folds the key values of rows `idx` of column `a` into `out` (one slot per row): each value
87/// becomes a u64 that is equal for equal keys (DataFusion's grouping: -0.0 and +0.0 one group, each
88/// NaN bit pattern its own, null its own), mixed, and combined with what `out` holds unless
89/// `first`. One type dispatch per column and batch, not per value.
90fn fold_column(a: &dyn Array, idx: &[usize], out: &mut [u64], first: bool) {
91    fn put(out: &mut [u64], j: usize, v: u64, first: bool) {
92        let v = splitmix(v);
93        out[j] = if first { v } else { out[j].rotate_left(23).wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ v };
94    }
95    let nulls = a.nulls();
96    let valid = |i: usize| nulls.is_none_or(|n| n.is_valid(i));
97    macro_rules! prim {
98        ($t:ty, $f:expr) => {{
99            let v = a.as_primitive::<$t>().values();
100            for (j, &i) in idx.iter().enumerate() {
101                put(out, j, if valid(i) { $f(v[i]) } else { NULL_KEY }, first);
102            }
103        }};
104    }
105    macro_rules! strs {
106        ($arr:expr) => {{
107            let s = $arr;
108            for (j, &i) in idx.iter().enumerate() {
109                put(out, j, if valid(i) { fnv1a(s.value(i).as_bytes()) } else { NULL_KEY }, first);
110            }
111        }};
112    }
113    match a.data_type() {
114        DataType::Int8 => prim!(Int8Type, |x: i8| x as i64 as u64),
115        DataType::Int16 => prim!(Int16Type, |x: i16| x as i64 as u64),
116        DataType::Int32 => prim!(Int32Type, |x: i32| x as i64 as u64),
117        DataType::Int64 => prim!(Int64Type, |x: i64| x as u64),
118        DataType::UInt8 => prim!(UInt8Type, |x: u8| x as u64),
119        DataType::UInt16 => prim!(UInt16Type, |x: u16| x as u64),
120        DataType::UInt32 => prim!(UInt32Type, |x: u32| x as u64),
121        DataType::UInt64 => prim!(UInt64Type, |x: u64| x),
122        DataType::Float64 => prim!(Float64Type, |x: f64| (if x == 0.0 { 0.0 } else { x }).to_bits()),
123        DataType::Float32 => prim!(Float32Type, |x: f32| (if x == 0.0 { 0.0f32 } else { x }).to_bits() as u64),
124        DataType::Utf8 => strs!(a.as_string::<i32>()),
125        DataType::LargeUtf8 => strs!(a.as_string::<i64>()),
126        DataType::Utf8View => strs!(a.as_string_view()),
127        // Other key types are not taken by the rule (the caller never gets here): only nulls
128        // are told apart.
129        _ => {
130            for (j, &i) in idx.iter().enumerate() {
131                put(out, j, if valid(i) { 0 } else { NULL_KEY }, first);
132            }
133        }
134    }
135}
136
137/// One u64 per sampled row (`pos`: increasing positions over the batches in order), equal for
138/// equal key tuples.
139fn sample(batches: &[&RecordBatch], keys: &[usize], pos: &[usize]) -> Vec<u64> {
140    let mut out = vec![0u64; pos.len()];
141    let mut start = 0usize;
142    let mut at = 0usize;
143    let mut idx = Vec::new();
144    for b in batches {
145        let end = start + b.num_rows();
146        idx.clear();
147        let from = at;
148        while at < pos.len() && pos[at] < end {
149            idx.push(pos[at] - start);
150            at += 1;
151        }
152        if !idx.is_empty() {
153            for (j, &k) in keys.iter().enumerate() {
154                fold_column(b.column(k).as_ref(), &idx, &mut out[from..at], j == 0);
155            }
156        }
157        start = end;
158        if at == pos.len() {
159            break;
160        }
161    }
162    out
163}
164
165/// (d, f1, f2) of a sample of key tuples: an open-addressing count (the values are already mixed
166/// hashes), about a quarter of the time of sorting them.
167fn stats(v: Vec<u64>) -> (u64, u64, u64) {
168    let size = (v.len() * 2).next_power_of_two().max(16);
169    let mask = size - 1;
170    let mut keys = vec![0u64; size];
171    let mut counts = vec![0u32; size];
172    let mut d = 0u64;
173    for x in v {
174        let mut i = (x as usize) & mask;
175        loop {
176            if counts[i] == 0 {
177                keys[i] = x;
178                counts[i] = 1;
179                d += 1;
180                break;
181            }
182            if keys[i] == x {
183                counts[i] += 1;
184                break;
185            }
186            i = (i + 1) & mask;
187        }
188    }
189    let f1 = counts.iter().filter(|&&c| c == 1).count() as u64;
190    let f2 = counts.iter().filter(|&&c| c == 2).count() as u64;
191    (d, f1, f2)
192}
193
194/// (estimate, low end, high end), each clipped to [d, rows]. The high end is the whole input while
195/// f2 could still be 0 (below about 6 pairs seen), because Chao1 cannot see a group count much past
196/// n^2 / 2 without pairs.
197fn chao1(d: u64, f1: u64, f2: u64, rows: usize) -> (f64, f64, f64) {
198    let (d, f1, f2, rows) = (d as f64, f1 as f64, f2 as f64, rows as f64);
199    let s = f2.sqrt();
200    let at = |f: f64| rows.min(d.max(d + f1 * (f1 - 1.0) / (2.0 * (f + 1.0))));
201    let low_f2 = f2 - 2.0 * s - 1.0;
202    let hi = if low_f2 > 0.0 {
203        at(low_f2)
204    } else if f1 == 0.0 {
205        d
206    } else {
207        rows
208    };
209    (at(f2), at(f2 + 2.0 * s + 2.0), hi)
210}
211
212/// Estimates the groups of `keys` (column indices) over an input of `total` rows (`None`: the rows
213/// of `batches`) from `batches`: the whole input, or a part of it (the first batches of each
214/// partition), which the sample is drawn from. `settled(lo, hi)`: whether every group count from
215/// `lo` to `hi` gets the same decision; `None` stops once the estimate is within twice d, f2
216/// reaches 8, or the low end is a quarter of the input. (The run-time choice uses
217/// [`estimate_up_to`] with a smaller largest sample.)
218#[cfg(test)]
219pub(crate) fn estimate(
220    batches: &[&RecordBatch],
221    keys: &[usize],
222    settled: Option<&dyn Fn(u64, u64) -> bool>,
223    total: Option<usize>,
224) -> GroupEstimate {
225    estimate_up_to(batches, keys, settled, total, SAMPLE_MAX)
226}
227
228/// [`estimate`] with samples of at most `max_sample` rows.
229pub(crate) fn estimate_up_to(
230    batches: &[&RecordBatch],
231    keys: &[usize],
232    settled: Option<&dyn Fn(u64, u64) -> bool>,
233    total: Option<usize>,
234    max_sample: usize,
235) -> GroupEstimate {
236    let t = Instant::now();
237    let batches: Vec<&RecordBatch> = batches.iter().copied().filter(|b| b.num_rows() > 0).collect();
238    let avail: usize = batches.iter().map(|b| b.num_rows()).sum();
239    let rows = total.unwrap_or(avail).max(avail);
240    if avail == rows && rows <= EXACT {
241        let all: Vec<usize> = (0..rows).collect();
242        let (d, _, _) = stats(sample(&batches, keys, &all));
243        return GroupEstimate {
244            estimate: d,
245            low: d,
246            high: d,
247            sample_rows: rows,
248            rows,
249            exact: true,
250            time: t.elapsed(),
251        };
252    }
253    let cap = max_sample.min(SAMPLE_MAX).min(avail / 4).max(FIRST.min(avail));
254    let mut n = FIRST.min(avail);
255    loop {
256        let (d, f1, f2) = stats(sample(&batches, keys, &positions(avail, n)));
257        let (est, lo, hi) = chao1(d, f1, f2, rows);
258        let done = n * GROWTH > cap
259            || match settled {
260                Some(s) => s(lo.round() as u64, hi.round() as u64),
261                None => est <= 2.0 * d as f64 || f2 >= SETTLED_F2 || lo * 4.0 >= rows as f64,
262            };
263        if done {
264            return GroupEstimate {
265                estimate: est.round() as u64,
266                low: lo.round() as u64,
267                high: hi.round() as u64,
268                sample_rows: n,
269                rows,
270                exact: false,
271                time: t.elapsed(),
272            };
273        }
274        n *= GROWTH;
275    }
276}
277
278#[cfg(test)]
279mod tests {
280    use super::*;
281    use arrow::array::{Int32Array, StringArray};
282    use arrow::datatypes::{Field, Schema};
283    use std::sync::Arc;
284
285    fn batches(keys: Vec<i32>, chunk: usize) -> Vec<RecordBatch> {
286        let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int32, true)]));
287        keys.chunks(chunk)
288            .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(Int32Array::from(c.to_vec()))]).unwrap())
289            .collect()
290    }
291
292    fn uniform(rows: usize, groups: u64) -> Vec<i32> {
293        (0..rows as u64).map(|i| (splitmix(i ^ 0x55) % groups) as i32).collect()
294    }
295
296    #[test]
297    fn small_inputs_are_counted_exactly() {
298        let b = batches((0..4000).map(|i| i % 37).collect(), 1000);
299        let r: Vec<&RecordBatch> = b.iter().collect();
300        let e = estimate(&r, &[0], None, None);
301        assert!(e.exact);
302        assert_eq!((e.estimate, e.low, e.high), (37, 37, 37));
303    }
304
305    #[test]
306    fn estimates_land_near_the_truth() {
307        // (groups in the key domain, rows); the realised distinct count is what is compared.
308        for (groups, rows) in [(200u64, 1_000_000usize), (10_000, 1_000_000), (100_000, 2_000_000), (1_000_000, 4_000_000)] {
309            let keys = uniform(rows, groups);
310            let mut seen = keys.clone();
311            seen.sort_unstable();
312            seen.dedup();
313            let truth = seen.len() as f64;
314            let b = batches(keys, 8192);
315            let r: Vec<&RecordBatch> = b.iter().collect();
316            let e = estimate(&r, &[0], None, None);
317            let ratio = e.estimate as f64 / truth;
318            assert!((0.5..2.0).contains(&ratio), "groups {groups}, rows {rows}: estimate {e:?}, truth {truth}");
319            assert!(e.low as f64 <= truth * 1.2 && e.high as f64 >= truth * 0.8, "{e:?} vs {truth}");
320        }
321    }
322
323    #[test]
324    fn sorted_keys_are_not_underestimated() {
325        // Keys in order (each value twice in a row): a stratified sample sees them spread out.
326        let rows = 1_000_000;
327        let keys: Vec<i32> = (0..rows as i32).map(|i| i / 2).collect();
328        let b = batches(keys, 8192);
329        let r: Vec<&RecordBatch> = b.iter().collect();
330        let e = estimate(&r, &[0], None, None);
331        assert!(e.high as usize >= rows / 4, "{e:?}");
332    }
333
334    #[test]
335    fn two_keys_and_strings_hash_as_tuples() {
336        let n = 20_000;
337        let schema = Arc::new(Schema::new(vec![
338            Field::new("a", DataType::Int32, false),
339            Field::new("s", DataType::Utf8, true),
340        ]));
341        let a = Int32Array::from((0..n).map(|i| i % 10).collect::<Vec<i32>>());
342        let s = StringArray::from((0..n).map(|i| if i % 7 == 0 { None } else { Some(format!("s{}", i % 3)) }).collect::<Vec<_>>());
343        let b = RecordBatch::try_new(schema, vec![Arc::new(a), Arc::new(s)]).unwrap();
344        let e = estimate(&[&b], &[0, 1], None, None);
345        // 10 x (3 strings + null) = 40 tuples.
346        assert!((30..=48).contains(&e.estimate), "{e:?}");
347    }
348
349    #[test]
350    fn a_prefix_estimates_the_whole_input() {
351        // The first 262,144 of 2,000,000 rows, keys uniform over 100,000 values.
352        let keys = uniform(2_000_000, 100_000);
353        let b = batches(keys[..262_144].to_vec(), 8192);
354        let r: Vec<&RecordBatch> = b.iter().collect();
355        let e = estimate(&r, &[0], None, Some(2_000_000));
356        assert_eq!(e.rows, 2_000_000);
357        assert!((50_000..200_000).contains(&e.estimate), "{e:?}");
358    }
359
360    /// Probe timing (release build): `cargo test --release --lib probe_timing -- --ignored --nocapture`.
361    #[test]
362    #[ignore]
363    fn probe_timing() {
364        let keys: Vec<i64> = (0..262_144u64).map(|i| (splitmix(i) % 1_000_000) as i64).collect();
365        let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int64, false)]));
366        let b: Vec<RecordBatch> = keys
367            .chunks(8192)
368            .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(arrow::array::Int64Array::from(c.to_vec()))]).unwrap())
369            .collect();
370        let r: Vec<&RecordBatch> = b.iter().collect();
371        for n in [512usize, 2048, 8192, 32768] {
372            let pos = positions(262_144, n);
373            let t = Instant::now();
374            let mut v = Vec::new();
375            for _ in 0..100 {
376                v = sample(&r, &[0], &pos);
377            }
378            let ts = t.elapsed() / 100;
379            let t = Instant::now();
380            for _ in 0..100 {
381                std::hint::black_box(stats(v.clone()));
382            }
383            let tt = t.elapsed() / 100;
384            println!("n {n}: sample {ts:?}, stats {tt:?}");
385        }
386    }
387
388    #[test]
389    fn same_input_same_estimate() {
390        let b = batches(uniform(500_000, 50_000), 8192);
391        let r: Vec<&RecordBatch> = b.iter().collect();
392        let (x, y) = (estimate(&r, &[0], None, None), estimate(&r, &[0], None, None));
393        assert_eq!((x.estimate, x.low, x.high, x.sample_rows), (y.estimate, y.low, y.high, y.sample_rows));
394    }
395}