datafusion-arrowmetal 0.4.1

A DataFusion 55.1 physical optimizer rule that runs full sorts, and count(*), DISTINCT and integer MIN/MAX group-bys over tables of 50M rows or more in the measured group-count ranges, on Apple silicon GPUs through ArrowMetal; hash joins, top-k and filters are available and off by default
Documentation
//! The run-time choice for a replaced aggregate: ArrowMetal or DataFusion's own operators, from the
//! group-count estimate and the measured table (`agg_table.rs`, generated by
//! `scripts/groupby_table.py` from the sweep CSV it names).
//!
//! A replaced aggregate is described by its *shape*: the aggregate families it computes, whether it
//! groups by one key or several, and the class of its key types. The table holds, per shape and
//! group-count bucket, the row counts at which the sweep measured ArrowMetal at least
//! `MIN_RATIO x HEADROOM` ahead of DataFusion alone. A node is run on ArrowMetal only when every
//! family it computes is taken at the estimated bucket and the input's row count.

use arrow::datatypes::DataType;
use datafusion::datasource::memory::MemorySourceConfig;
use datafusion::datasource::source::DataSourceExec;
use datafusion::physical_plan::ExecutionPlan;

use crate::agg_table as table;
use crate::exec::{AggKind, MetalOp};

/// The shape of a replaced aggregate, as the table keys it.
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Shape {
    /// Aggregate families, sorted, deduplicated (`distinct` for a GROUP BY with no aggregates).
    pub families: Vec<&'static str>,
    /// "1" or "2+".
    pub keys: &'static str,
    /// "i32" (integer keys of at most 32 bits, UInt32 excluded), "i64" (any 64-bit or UInt32
    /// key among them), "float", "string" (the widest class among the keys: string > float > i64
    /// > i32).
    pub key_class: &'static str,
    /// "memory" when the aggregate reads an in-memory scan directly (a `MemTable`, a
    /// `MemorySourceConfig`), else "other" (a file scan, a filter, a join, ...).
    pub source: &'static str,
}

impl std::fmt::Display for Shape {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(
            f,
            "{} over {} {} key{} ({} input)",
            self.families.join(" + "),
            self.keys,
            self.key_class,
            if self.keys == "1" { "" } else { "s" },
            self.source
        )
    }
}

fn key_class(t: &DataType) -> &'static str {
    match t {
        DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::UInt8 | DataType::UInt16 => "i32",
        DataType::Int64 | DataType::UInt32 | DataType::UInt64 => "i64",
        DataType::Float32 | DataType::Float64 => "float",
        _ => "string",
    }
}

fn is_int(t: &DataType) -> bool {
    t.is_integer()
}

/// "memory" or "other" (see [`Shape::source`]).
fn source_class(input: &dyn ExecutionPlan) -> &'static str {
    match input.downcast_ref::<DataSourceExec>() {
        Some(d) if d.data_source().downcast_ref::<MemorySourceConfig>().is_some() => "memory",
        _ => "other",
    }
}

/// The shape of an aggregate op reading `input`.
pub(crate) fn shape(op: &MetalOp, input: &dyn ExecutionPlan) -> Option<Shape> {
    let MetalOp::Aggregate { keys, aggs } = op else { return None };
    let schema = input.schema();
    let rank = |c: &str| ["i32", "i64", "float", "string"].iter().position(|x| *x == c).unwrap_or(3);
    let mut kc = "i32";
    for &k in keys {
        let c = key_class(schema.field(k).data_type());
        if rank(c) > rank(kc) {
            kc = c;
        }
    }
    let mut families: Vec<&'static str> = aggs
        .iter()
        .map(|a| {
            let t = a.column.map(|c| schema.field(c).data_type().clone());
            match (a.kind, t) {
                (AggKind::CountAll | AggKind::Count, _) => "count",
                (AggKind::Sum | AggKind::Mean, Some(t)) if is_int(&t) => "sum_avg_int",
                (AggKind::Sum | AggKind::Mean, Some(DataType::Float64)) => "sum_avg_f64",
                (AggKind::Sum | AggKind::Mean, _) => "sum_avg_f32",
                (AggKind::Min | AggKind::Max, Some(t)) if is_int(&t) => "minmax_int",
                (AggKind::Min | AggKind::Max, _) => "minmax_f64",
            }
        })
        .collect();
    if families.is_empty() {
        families.push("distinct");
    }
    families.sort_unstable();
    families.dedup();
    Some(Shape { families, keys: if keys.len() == 1 { "1" } else { "2+" }, key_class: kc, source: source_class(input) })
}

/// The bucket of `groups` groups over `rows` input rows (`None`: between the last bucket and
/// rows / NEAR_ROWS, or no groups).
pub(crate) fn bucket(groups: u64, rows: u64) -> Option<&'static str> {
    if groups == 0 {
        return None;
    }
    if groups.saturating_mul(table::NEAR_ROWS) >= rows {
        return Some(table::ROWS_BUCKET);
    }
    table::BUCKETS.iter().find(|(_, lo, hi)| *lo <= groups && groups <= *hi).map(|b| b.0)
}

fn row(family: &str, s: &Shape, bucket: &str) -> Option<&'static table::Row> {
    table::TABLE
        .iter()
        .find(|r| {
            r.family == family && r.keys == s.keys && r.key_class == s.key_class && r.source == s.source && r.bucket == bucket
        })
}

/// Whether the table takes family `family` of shape `s` at `bucket` and `rows`: Ok(the rows it
/// takes from) or Err(why not).
fn family_takes(family: &str, s: &Shape, bucket: &str, rows: u64) -> Result<u64, String> {
    let Some(r) = row(family, s, bucket) else {
        return Err(format!(
            "{family} over {} {} key(s), {} input, at the {bucket} bucket was not measured",
            s.keys, s.key_class, s.source
        ));
    };
    let ratios = r.ratios.iter().map(|(n, x)| format!("{}: {x:.2}x", rows_text(*n))).collect::<Vec<_>>().join(", ");
    match r.min_rows {
        None => Err(format!(
            "{family} at the {bucket} bucket is not taken at any measured size (needs {:.2}x at two or more sizes; {ratios})",
            table::MIN_RATIO * table::HEADROOM
        )),
        Some(m) if rows < m => Err(format!("{family} at the {bucket} bucket is taken from {} rows ({ratios})", rows_text(m))),
        Some(_) if r.max_rows.is_some_and(|x| rows > x) => Err(format!(
            "{family} at the {bucket} bucket is taken up to {} rows only, the largest size measured, its ratio falling there ({ratios})",
            rows_text(r.max_rows.unwrap_or_default())
        )),
        Some(m) => Ok(m),
    }
}

pub(crate) fn rows_text(n: u64) -> String {
    if n >= 1_000_000 && n.is_multiple_of(1_000_000) {
        format!("{}M", n / 1_000_000)
    } else if n >= 1_000 && n.is_multiple_of(1_000) {
        format!("{}k", n / 1_000)
    } else {
        n.to_string()
    }
}

/// Whether every family of `s` is taken at `bucket` and `rows`: Ok(reason) or Err(reason).
pub(crate) fn takes(s: &Shape, bucket: Option<&str>, rows: u64) -> Result<String, String> {
    let Some(b) = bucket else {
        return Err("no bucket holds that group count at this row count".into());
    };
    let mut from = 0;
    for f in &s.families {
        from = from.max(family_takes(f, s, b, rows)?);
    }
    Ok(format!("{} at the {b} bucket is taken from {} rows", s.families.join(" + "), rows_text(from)))
}

/// The group-count regions at `rows` input rows: [(fewest, most, bucket or None)], in order.
fn regions(rows: u64) -> Vec<(u64, u64, Option<&'static str>)> {
    let edge = rows.div_ceil(table::NEAR_ROWS).max(1);
    let mut out = Vec::new();
    let mut top = 0;
    for &(name, lo, hi) in table::BUCKETS {
        if lo >= edge {
            break;
        }
        top = hi.min(edge - 1);
        out.push((lo, top, Some(name)));
    }
    if top + 1 < edge {
        out.push((top + 1, edge - 1, None));
    }
    out.push((edge, u64::MAX, Some(table::ROWS_BUCKET)));
    out
}

/// `settled(lo, hi)`: whether every group count from `lo` to `hi` gets the same decision for `s` at
/// `rows` input rows (the probe samples until its range is settled).
pub(crate) fn settled_for(s: &Shape, rows: u64) -> impl Fn(u64, u64) -> bool {
    let regions: Vec<(u64, u64, bool)> =
        regions(rows).into_iter().map(|(lo, hi, b)| (lo, hi, takes(s, b, rows).is_ok())).collect();
    move |lo, hi| {
        let mut seen: Option<bool> = None;
        for &(rlo, rhi, t) in &regions {
            if rlo <= hi && rhi >= lo {
                if seen.is_some_and(|x| x != t) {
                    return false;
                }
                seen = Some(t);
            }
        }
        true
    }
}

/// At plan time: whether any group count is taken for `s` at `rows` input rows. Ok(the buckets
/// taken, as text) or Err(why the node is left).
pub(crate) fn any_bucket(s: &Shape, rows: u64) -> Result<String, String> {
    let taken: Vec<&str> = regions(rows).into_iter().filter_map(|(_, _, b)| b).filter(|b| takes(s, Some(b), rows).is_ok()).collect();
    if taken.is_empty() {
        let smallest = smallest_rows(s);
        return Err(match smallest {
            Some(m) => format!(
                "{s}: the measured table takes no group count at {} rows (the first take is at {} rows; {})",
                rows_text(rows),
                rows_text(m),
                table::SOURCE
            ),
            None => format!("{s}: the measured table takes it at no group count and size ({})", table::SOURCE),
        });
    }
    Ok(format!("{s}: the measured table takes the {} bucket(s) at {} rows; MetalExec decides from a group-count estimate", taken.join(", "), rows_text(rows)))
}

// -------------------------------------------------------------------------------------------------
// Joins (`join_table.rs`, generated by `scripts/join_table.py`)
// -------------------------------------------------------------------------------------------------

/// The join table's key class of a join's keys: "i32" (Int32), "i64" (Int64) or "string"
/// (Utf8), the widest among them.
fn join_key_class(left: &dyn ExecutionPlan, keys: &[usize]) -> &'static str {
    let schema = left.schema();
    let rank = |c: &str| ["i32", "i64", "string"].iter().position(|x| *x == c).unwrap_or(2);
    let mut kc = "i32";
    for &k in keys {
        let c = match schema.field(k).data_type() {
            DataType::Int32 => "i32",
            DataType::Int64 => "i64",
            _ => "string",
        };
        if rank(c) > rank(kc) {
            kc = c;
        }
    }
    kc
}

/// At plan time: whether the measured join table takes a join `op` over `left` (build, `build`
/// rows) and `right` (probe, `probe` rows) inputs. Ok(reason) or Err(why it is left).
pub(crate) fn join_takes(
    op: &MetalOp,
    left: &dyn ExecutionPlan,
    right: &dyn ExecutionPlan,
    build: u64,
    probe: u64,
) -> Result<String, String> {
    use crate::join_table as jt;
    let MetalOp::Join { how, left_keys, .. } = op else {
        return Err("not a join".into());
    };
    let how = match how {
        crate::exec::JoinHow::Inner => "inner",
        crate::exec::JoinHow::Left => "left",
        crate::exec::JoinHow::Right => "right",
    };
    let key_class = join_key_class(left, left_keys);
    let source = if source_class(left) == "memory" && source_class(right) == "memory" { "memory" } else { "other" };
    let shape = format!("{how} join on {key_class} keys ({source} inputs)");
    let Some(&(bucket, _, _)) = jt::BUILD_BUCKETS.iter().find(|(_, lo, hi)| *lo <= build && build <= *hi) else {
        return Err(format!("{shape}: no build-side bucket of the measured join table holds {} rows ({})", rows_text(build), jt::SOURCE));
    };
    let Some(r) = jt::TABLE
        .iter()
        .find(|r| r.how == how && r.key_class == key_class && r.source == source && r.build == bucket)
    else {
        return Err(format!("{shape}, {bucket}-row build bucket: not measured ({})", jt::SOURCE));
    };
    match r.min_probe_rows {
        Some(m) if probe >= m => Ok(format!(
            "{shape}, {bucket}-row build bucket: the measured join table takes it from {} probe rows ({})",
            rows_text(m),
            r.reason
        )),
        Some(m) => Err(format!(
            "{shape}, {bucket}-row build bucket: the measured join table takes it from {} probe rows ({})",
            rows_text(m),
            r.reason
        )),
        None => Err(format!("{shape}, {bucket}-row build bucket: not taken ({}; {})", r.reason, jt::SOURCE)),
    }
}

/// The fewest rows at which the table takes `s` at some bucket.
fn smallest_rows(s: &Shape) -> Option<u64> {
    let mut best: Option<u64> = None;
    let buckets = table::BUCKETS.iter().map(|b| b.0).chain(std::iter::once(table::ROWS_BUCKET));
    for b in buckets {
        let mut need = Some(0u64);
        for f in &s.families {
            need = match (need, row(f, s, b).and_then(|r| r.min_rows)) {
                (Some(a), Some(m)) => Some(a.max(m)),
                _ => None,
            };
        }
        if let Some(n) = need {
            best = Some(best.map_or(n, |x: u64| x.min(n)));
        }
    }
    best
}