Skip to main content

datafusion_arrowmetal/
lib.rs

1//! ArrowMetal inside Apache DataFusion: a physical optimizer rule and a custom `ExecutionPlan`.
2//!
3//! [`ArrowMetalRule`] walks DataFusion's optimized physical plan and replaces the nodes ArrowMetal
4//! can run with the same answer by a [`MetalExec`]: `SortExec` (with or without `fetch`), a hash
5//! `AggregateExec` over column keys with `sum`/`min`/`max`/`count`/`avg` or none (DISTINCT),
6//! a `FilterExec` whose predicate translates, and a `HashJoinExec` (inner, left, right) on equal
7//! int32, int64 or Utf8 keys. `MetalExec` collects its input's partitions, runs the operation
8//! through ArrowMetal's plan runner on the GPU, and emits `RecordBatch`es with the schema
9//! DataFusion expects. Anything it does not support is left unchanged, and the reason is recorded
10//! in the rule's [`Report`].
11//!
12//! The default config ([`ArrowMetalConfig::default`]) takes full sorts from 250,000 input rows,
13//! aggregates where the measured table (`src/agg_table.rs`) takes their shape (a replaced
14//! aggregate estimates its group count from a sample of its keys when it runs and either runs on
15//! ArrowMetal or hands the node back to DataFusion's own operators, [`AggregateChoice`]), and
16//! joins where the measured join table (`src/join_table.rs`) takes them ([`JoinChoice`]). Top-k
17//! and filters are left.
18//!
19//! ```no_run
20//! # async fn f() -> datafusion::error::Result<()> {
21//! use datafusion::prelude::*;
22//! use datafusion_arrowmetal::{session_context, ArrowMetalConfig, ArrowMetalRule};
23//!
24//! // The measured take-list; `ArrowMetalConfig::all()` takes every shape the rule can translate.
25//! let rule = ArrowMetalRule::new(ArrowMetalConfig::default());
26//! let ctx = session_context(SessionConfig::new(), rule.clone());
27//! // register tables, run SQL ...
28//! for d in rule.report().decisions() { println!("{d}"); }
29//! # Ok(()) }
30//! ```
31
32#![warn(missing_docs)]
33
34mod agg_table;
35mod choice;
36mod exec;
37mod gpu;
38mod join_table;
39mod probe;
40mod rule;
41mod translate;
42
43pub use exec::{AggKind, AggSpec, JoinHow, MetalExec, MetalOp, SortKey};
44pub use probe::GroupEstimate;
45pub use rule::{AggregateChoice, ArrowMetalConfig, ArrowMetalRule, Decision, GroupChoice, JoinChoice, Report};
46
47use std::sync::Arc;
48
49use datafusion::execution::session_state::SessionStateBuilder;
50use datafusion::physical_optimizer::optimizer::PhysicalOptimizer;
51use datafusion::physical_optimizer::PhysicalOptimizerRule;
52use datafusion::prelude::{SessionConfig, SessionContext};
53
54/// DataFusion's default physical optimizer rules with `rule` inserted before the last two:
55/// the post-optimization `FilterPushdown` (which wires dynamic filters to the operators that own
56/// them, so a node replaced after it would leave a dynamic filter nobody updates) and
57/// `SanityCheckPlan` (so the rewritten plan is still checked for distribution and ordering).
58pub fn physical_optimizer_rules(
59    rule: ArrowMetalRule,
60) -> Vec<Arc<dyn PhysicalOptimizerRule + Send + Sync>> {
61    let mut rules = PhysicalOptimizer::new().rules;
62    let at = rules
63        .iter()
64        .rposition(|r| r.name() == "SanityCheckPlan")
65        .unwrap_or(rules.len());
66    // The post-optimization FilterPushdown sits just before SanityCheckPlan in 55.1.0.
67    let at = if at > 0 && rules[at - 1].name().starts_with("FilterPushdown") { at - 1 } else { at };
68    rules.insert(at, Arc::new(rule));
69    rules
70}
71
72/// Registers `rule` on a `SessionStateBuilder` (see [`physical_optimizer_rules`] for where).
73pub fn with_arrowmetal(builder: SessionStateBuilder, rule: ArrowMetalRule) -> SessionStateBuilder {
74    builder.with_physical_optimizer_rules(physical_optimizer_rules(rule))
75}
76
77/// A `SessionContext` with DataFusion's defaults plus `rule`.
78pub fn session_context(config: SessionConfig, rule: ArrowMetalRule) -> SessionContext {
79    let state = with_arrowmetal(
80        SessionStateBuilder::new().with_config(config).with_default_features(),
81        rule,
82    )
83    .build();
84    SessionContext::new_with_state(state)
85}