1use std::fmt;
4use std::sync::{Arc, Mutex};
5
6use arrow::array::{Array, AsArray};
7use arrow::datatypes::{DataType, SchemaRef};
8use arrow::record_batch::RecordBatch;
9use datafusion::common::runtime::SpawnedTask;
10use datafusion::common::tree_node::{Transformed, TreeNode, TreeNodeRecursion};
11use datafusion::common::{internal_err, DataFusionError, Result};
12use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation};
13use datafusion::execution::TaskContext;
14use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr};
15use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
16use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
17use datafusion::physical_plan::metrics::{Count, ExecutionPlanMetricsSet, Gauge, MetricBuilder, MetricsSet, Time};
18use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
19use datafusion::physical_plan::{
20 execute_stream, DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, Partitioning,
21 PlanProperties, SendableRecordBatchStream,
22};
23use futures::stream::BoxStream;
24use futures::{FutureExt, StreamExt, TryStreamExt};
25
26use crate::rule::{lock, AggregateChoice, Decision, SharedLog};
27
28#[derive(Debug, Clone, PartialEq)]
30#[non_exhaustive]
31pub struct SortKey {
32 pub column: usize,
34 pub descending: bool,
36 pub nulls_first: bool,
38}
39
40#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42#[non_exhaustive]
43pub enum AggKind {
44 Sum,
46 Min,
48 Max,
50 Count,
52 CountAll,
54 Mean,
56}
57
58#[derive(Debug, Clone, PartialEq)]
60#[non_exhaustive]
61pub struct AggSpec {
62 pub kind: AggKind,
64 pub column: Option<usize>,
66 pub float: bool,
69 pub out_type: DataType,
71}
72
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76#[non_exhaustive]
77pub enum JoinHow {
78 Inner,
80 Left,
82 Right,
84}
85
86#[derive(Debug, Clone, PartialEq)]
88#[non_exhaustive]
89pub enum MetalOp {
90 Sort {
92 keys: Vec<SortKey>,
94 fetch: Option<usize>,
96 },
97 Aggregate {
99 keys: Vec<usize>,
101 aggs: Vec<AggSpec>,
103 },
104 Filter {
106 predicate: String,
108 projection: Option<Vec<usize>>,
110 },
111 Join {
115 how: JoinHow,
117 left_keys: Vec<usize>,
119 right_keys: Vec<usize>,
121 projection: Vec<usize>,
123 left_columns: usize,
125 },
126}
127
128pub struct MetalExec {
148 op: MetalOp,
149 inputs: Vec<Arc<dyn ExecutionPlan>>,
151 original: Arc<dyn ExecutionPlan>,
153 props: Arc<PlanProperties>,
154 log: SharedLog,
155 plan: u64,
157 metrics: ExecutionPlanMetricsSet,
158 choice: AggregateChoice,
160 table_rows: Option<usize>,
162 rows_hint: Option<usize>,
164 join_slot: JoinSlot,
166}
167
168#[derive(Debug, Clone, Copy)]
170pub(crate) struct AggSettings {
171 pub choice: AggregateChoice,
172 pub table_rows: Option<usize>,
174 pub rows_hint: Option<usize>,
176}
177
178struct PhaseMetrics {
180 input_time: Time,
181 import_time: Time,
182 kernel_time: Time,
183 export_time: Time,
184 probe_time: Time,
185 input_batches: Count,
186 handed_back: Count,
187 groups_estimate: Gauge,
188}
189
190impl PhaseMetrics {
191 fn add(&self, t: &crate::gpu::GpuTimes) {
192 self.import_time.add_duration(t.import);
193 self.kernel_time.add_duration(t.kernel);
194 self.export_time.add_duration(t.export);
195 }
196}
197
198fn slices(b: RecordBatch, batch_size: usize) -> Vec<RecordBatch> {
200 let mut v = Vec::new();
201 let mut off = 0;
202 while off < b.num_rows() {
203 let n = batch_size.min(b.num_rows() - off);
204 v.push(b.slice(off, n));
205 off += n;
206 }
207 v
208}
209
210pub(crate) fn batch_bytes(b: &RecordBatch) -> usize {
213 b.columns().iter().map(|a| array_bytes(a.as_ref())).sum()
214}
215
216fn array_bytes(a: &dyn Array) -> usize {
217 let n = a.len();
218 let nulls = a.nulls().map_or(0, |_| n.div_ceil(8));
219 let values = match a.data_type() {
220 DataType::Boolean => n.div_ceil(8),
221 DataType::Utf8 => {
222 let o = a.as_string::<i32>().value_offsets();
223 (n + 1) * 4 + (o[n] - o[0]) as usize
224 }
225 DataType::LargeUtf8 => {
226 let o = a.as_string::<i64>().value_offsets();
227 (n + 1) * 8 + (o[n] - o[0]) as usize
228 }
229 DataType::Utf8View => {
230 let v = a.as_string_view();
231 n * 16 + v.views().iter().map(|w| *w as u32 as usize).filter(|&l| l > 12).sum::<usize>()
232 }
233 t => match t.primitive_width() {
234 Some(w) => n * w,
235 None => return a.get_array_memory_size(),
236 },
237 };
238 values + nulls
239}
240
241struct Collected {
243 parts: Vec<Vec<RecordBatch>>,
245 rest: Vec<Option<SendableRecordBatchStream>>,
247 held: Vec<MemoryReservation>,
249 refused: Option<String>,
251 per_input: Vec<usize>,
253}
254
255impl Collected {
256 fn rows(&self) -> usize {
257 self.parts.iter().flatten().map(|b| b.num_rows()).sum()
258 }
259}
260
261const PREFIX_ROWS: usize = 262_144;
265
266const DECIDE_SAMPLE: usize = 8_192;
270
271const CONFIRM_SAMPLE: usize = 2_048;
273
274async fn collect_streams(
279 streams: Vec<SendableRecordBatchStream>,
280 ctx: &TaskContext,
281 quota: Option<usize>,
282) -> Result<Collected> {
283 let drain = |mut s: SendableRecordBatchStream, r: MemoryReservation| async move {
284 let mut v = Vec::new();
285 let mut rows = 0usize;
286 while let Some(b) = s.next().await {
287 let b = b?;
288 rows += b.num_rows();
289 let refused = r.try_grow(batch_bytes(&b)).err().map(|e| e.to_string());
290 v.push(b);
291 if refused.is_some() || quota.is_some_and(|q| rows >= q) {
292 return Ok::<_, DataFusionError>((v, Some(s), r, refused));
293 }
294 }
295 Ok((v, None, r, None))
296 };
297 let reservation = || MemoryConsumer::new("MetalExec").register(ctx.memory_pool());
298 let results: Vec<_> = if quota.is_some() || streams.len() == 1 {
299 futures::future::join_all(streams.into_iter().map(|s| drain(s, reservation())))
300 .await
301 .into_iter()
302 .collect::<Result<_>>()?
303 } else {
304 let tasks: Vec<_> = streams.into_iter().map(|s| SpawnedTask::spawn(drain(s, reservation()))).collect();
305 let mut v = Vec::with_capacity(tasks.len());
306 for t in tasks {
307 v.push(t.join().await.map_err(|e| DataFusionError::External(Box::new(e)))??);
308 }
309 v
310 };
311 let mut out = Collected { parts: Vec::new(), rest: Vec::new(), held: Vec::new(), refused: None, per_input: Vec::new() };
312 for (v, rest, r, refused) in results {
313 out.parts.push(v);
314 out.rest.push(rest);
315 out.held.push(r);
316 if out.refused.is_none() {
317 out.refused = refused;
318 }
319 }
320 if out.parts.is_empty() {
321 out.parts.push(Vec::new());
322 out.rest.push(None);
323 }
324 out.per_input = vec![out.parts.len()];
325 Ok(out)
326}
327
328async fn finish(mut c: Collected, ctx: &TaskContext) -> Result<Collected> {
331 let open: Vec<usize> = (0..c.rest.len()).filter(|&i| c.rest[i].is_some()).collect();
332 if open.is_empty() {
333 return Ok(c);
334 }
335 let streams: Vec<SendableRecordBatchStream> = open.iter().filter_map(|&i| c.rest[i].take()).collect();
336 let more = collect_streams(streams, ctx, None).await?;
337 for ((i, v), rest) in open.into_iter().zip(more.parts).zip(more.rest) {
338 c.parts[i].extend(v);
339 c.rest[i] = rest;
340 }
341 c.held.extend(more.held);
342 if c.refused.is_none() {
343 c.refused = more.refused;
344 }
345 Ok(c)
346}
347
348type ReplaySlot = Option<(Vec<RecordBatch>, Option<SendableRecordBatchStream>)>;
353
354struct ReplayExec {
355 schema: SchemaRef,
356 slots: Mutex<Vec<ReplaySlot>>,
357 props: Arc<PlanProperties>,
358}
359
360impl fmt::Debug for ReplayExec {
361 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
362 f.debug_struct("ReplayExec").finish()
363 }
364}
365
366impl ReplayExec {
367 fn new(schema: SchemaRef, parts: Vec<Vec<RecordBatch>>, rest: Vec<Option<SendableRecordBatchStream>>) -> Self {
368 let n = parts.len();
369 let props = PlanProperties::new(
370 EquivalenceProperties::new(Arc::clone(&schema)),
371 Partitioning::UnknownPartitioning(n),
372 EmissionType::Incremental,
373 Boundedness::Bounded,
374 );
375 let slots = parts.into_iter().zip(rest).map(Some).collect();
376 Self { schema, slots: Mutex::new(slots), props: Arc::new(props) }
377 }
378}
379
380impl DisplayAs for ReplayExec {
381 fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
382 write!(f, "ReplayExec")
383 }
384}
385
386impl ExecutionPlan for ReplayExec {
387 fn name(&self) -> &str {
388 "ReplayExec"
389 }
390 fn properties(&self) -> &Arc<PlanProperties> {
391 &self.props
392 }
393 fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
394 vec![]
395 }
396 fn apply_expressions(
397 &self,
398 _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
399 ) -> Result<TreeNodeRecursion> {
400 Ok(TreeNodeRecursion::Continue)
401 }
402 fn with_new_children(self: Arc<Self>, children: Vec<Arc<dyn ExecutionPlan>>) -> Result<Arc<dyn ExecutionPlan>> {
403 if children.is_empty() {
404 Ok(self)
405 } else {
406 internal_err!("ReplayExec has no children")
407 }
408 }
409 fn execute(&self, partition: usize, _ctx: Arc<TaskContext>) -> Result<SendableRecordBatchStream> {
410 let slot = self.slots.lock().unwrap_or_else(|p| p.into_inner()).get_mut(partition).and_then(Option::take);
411 let Some((batches, rest)) = slot else {
412 return internal_err!("ReplayExec partition {partition} executed twice or out of range");
413 };
414 let head = futures::stream::iter(batches.into_iter().map(Ok));
415 let s: BoxStream<'static, Result<RecordBatch>> = match rest {
416 Some(r) => Box::pin(head.chain(r)),
417 None => Box::pin(head),
418 };
419 Ok(Box::pin(RecordBatchStreamAdapter::new(Arc::clone(&self.schema), s)))
420 }
421}
422
423impl fmt::Debug for MetalExec {
424 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
425 f.debug_struct("MetalExec").field("op", &self.op).finish()
426 }
427}
428
429struct Job {
431 op: MetalOp,
432 inputs: Vec<Arc<dyn ExecutionPlan>>,
433 original: Arc<dyn ExecutionPlan>,
434 out_schema: SchemaRef,
435 log: SharedLog,
436 plan: u64,
437 choice: AggregateChoice,
438 table_rows: Option<usize>,
439 rows_hint: Option<usize>,
440 mm: PhaseMetrics,
441 ctx: Arc<TaskContext>,
442}
443
444impl Job {
445 fn record(&self, d: Decision) {
446 lock(&self.log).push(self.plan, d);
447 }
448
449 fn hand_back(&self, c: Collected) -> Result<BoxStream<'static, Result<RecordBatch>>> {
456 let plan = self.hand_back_plan(c)?;
457 Ok(Box::pin(execute_stream(plan, Arc::clone(&self.ctx))?))
458 }
459
460 fn hand_back_plan(&self, c: Collected) -> Result<Arc<dyn ExecutionPlan>> {
462 drop(c.held);
463 let per_input = if c.per_input.len() == self.inputs.len() { c.per_input } else { vec![c.parts.len()] };
464 let (mut parts, mut rest) = (c.parts.into_iter(), c.rest.into_iter());
465 let mut plan = Arc::clone(&self.original);
466 for (input, n) in self.inputs.iter().zip(per_input) {
467 let p: Vec<Vec<RecordBatch>> = parts.by_ref().take(n).collect();
468 let r: Vec<Option<SendableRecordBatchStream>> = rest.by_ref().take(n).collect();
469 let leaf: Arc<dyn ExecutionPlan> = Arc::new(ReplayExec::new(input.schema(), p, r));
470 plan = replace_leaf(&plan, input, &leaf)?;
471 }
472 Ok(plan)
473 }
474
475 async fn run_gpu(self, c: Collected) -> Result<BoxStream<'static, Result<RecordBatch>>> {
478 let ctx = Arc::clone(&self.ctx);
479 match self.run_gpu_outcome(c).await? {
480 Outcome::Gpu(batches, held) => Ok(Box::pin(futures::stream::iter(batches.into_iter().map(Ok)).map(move |b| {
481 let _held = &held;
482 b
483 }))),
484 Outcome::Back(plan) => Ok(Box::pin(execute_stream(plan, ctx)?)),
485 }
486 }
487
488 async fn run_gpu_outcome(self, c: Collected) -> Result<Outcome> {
491 let op = self.op.clone();
492 let s2 = Arc::clone(&self.out_schema);
493 let parts = c.parts;
494 let per_input = c.per_input.clone();
495 let (res, parts, times) = tokio::task::spawn_blocking(move || {
496 let mut groups: Vec<Vec<&RecordBatch>> = Vec::new();
498 let mut it = parts.iter();
499 for &n in &per_input {
500 groups.push(it.by_ref().take(n).flatten().collect());
501 }
502 let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| crate::gpu::run(&op, &groups, &s2)))
503 .unwrap_or_else(|p| {
504 let msg = p
505 .downcast_ref::<String>()
506 .cloned()
507 .or_else(|| p.downcast_ref::<&str>().map(|s| s.to_string()))
508 .unwrap_or_default();
509 Err(format!("panic in the GPU path: {msg}"))
510 });
511 (r, parts, crate::gpu::take_times())
512 })
513 .await
514 .map_err(|e| DataFusionError::External(Box::new(e)))?;
515 self.mm.add(×);
516 let c = Collected { parts, ..c };
517 match res {
518 Ok(b) => {
519 let out = MemoryConsumer::new("MetalExec output").register(self.ctx.memory_pool());
520 if let Err(e) = out.try_grow(batch_bytes(&b)) {
521 self.record(Decision::memory_hand_back(&self.op, "the result", &e.to_string()));
522 self.mm.handed_back.add(1);
523 return Ok(Outcome::Back(self.hand_back_plan(c)?));
524 }
525 drop(c);
527 let batch_size = self.ctx.session_config().batch_size();
528 Ok(Outcome::Gpu(slices(b, batch_size), out))
529 }
530 Err(msg) => {
531 self.record(Decision::runtime_fallback(&self.op, &msg));
532 Ok(Outcome::Back(self.hand_back_plan(c)?))
533 }
534 }
535 }
536
537 async fn run_sort_or_filter(self, stream: SendableRecordBatchStream) -> Result<BoxStream<'static, Result<RecordBatch>>> {
539 let t = std::time::Instant::now();
540 let c = collect_streams(vec![stream], &self.ctx, None).await?;
541 self.mm.input_time.add_duration(t.elapsed());
542 self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
543 if let Some(e) = &c.refused {
544 self.record(Decision::memory_hand_back(&self.op, "the input", e));
545 self.mm.handed_back.add(1);
546 return self.hand_back(c);
547 }
548 self.run_gpu(c).await
549 }
550
551 async fn run_aggregate(self, streams: Vec<SendableRecordBatchStream>) -> Result<BoxStream<'static, Result<RecordBatch>>> {
563 let MetalOp::Aggregate { keys, .. } = &self.op else {
564 return internal_err!("run_aggregate on {:?}", self.op);
565 };
566 let keys = keys.clone();
567 let t = std::time::Instant::now();
568 if self.choice == AggregateChoice::DataFusion {
569 let n = streams.len();
570 let c = Collected {
571 parts: vec![Vec::new(); n],
572 rest: streams.into_iter().map(Some).collect(),
573 held: Vec::new(),
574 refused: None,
575 per_input: vec![n],
576 };
577 self.record(Decision::runtime_choice(&self.op, false, "aggregate_choice is DataFusion".into(), self.rows_hint.unwrap_or(0), None));
578 self.mm.handed_back.add(1);
579 return self.hand_back(c);
580 }
581 let prefix = match (self.choice, self.rows_hint) {
582 (AggregateChoice::Measured, Some(_)) => Some(PREFIX_ROWS.div_ceil(streams.len().max(1))),
583 _ => None,
584 };
585 let mut c = collect_streams(streams, &self.ctx, prefix).await?;
586 if let Some(e) = &c.refused {
587 self.mm.input_time.add_duration(t.elapsed());
588 self.record(Decision::memory_hand_back(&self.op, "the input", e));
589 self.mm.handed_back.add(1);
590 return self.hand_back(c);
591 }
592 let mut from_prefix: Option<(String, crate::probe::GroupEstimate)> = None;
593 if let (Some(_), Some(total)) = (prefix, self.rows_hint) {
594 if c.rest.iter().any(Option::is_some) {
595 self.mm.input_time.add_duration(t.elapsed());
597 let t = std::time::Instant::now();
598 let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
599 let (take, reason, estimate) = decide(&self.op, self.inputs[0].as_ref(), &refs, &keys, total, self.table_rows);
600 self.mm.probe_time.add_duration(t.elapsed());
601 if !take {
602 if let Some(e) = &estimate {
603 self.mm.groups_estimate.set(e.estimate as usize);
604 }
605 let reason = format!("{reason} (from the first {} rows of each partition)", prefix.unwrap_or(0));
606 self.record(Decision::runtime_choice(&self.op, false, reason, total, estimate));
607 self.mm.handed_back.add(1);
608 return self.hand_back(c);
609 }
610 if let Some(e) = estimate {
611 from_prefix = Some((reason, e));
612 }
613 let t = std::time::Instant::now();
614 c = finish(c, &self.ctx).await?;
615 self.mm.input_time.add_duration(t.elapsed());
616 if let Some(e) = &c.refused {
617 self.record(Decision::memory_hand_back(&self.op, "the input", e));
618 self.mm.handed_back.add(1);
619 return self.hand_back(c);
620 }
621 } else {
622 self.mm.input_time.add_duration(t.elapsed());
623 }
624 } else {
625 self.mm.input_time.add_duration(t.elapsed());
626 }
627 self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
628 let rows = c.rows();
629 let t = std::time::Instant::now();
630 let (on_gpu, reason, estimate) = match (self.choice, from_prefix) {
631 (AggregateChoice::Measured, Some((why, e))) => {
632 let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
636 let (agree, f) = confirm(&refs, &keys, rows, &e);
637 let check = format!(
638 "a {}-row sample of the whole input puts it at {} to {} groups",
639 f.sample_rows, f.low, f.high
640 );
641 if agree {
642 (true, format!("{why} (from the first rows of each partition; {check})"), Some(e))
643 } else {
644 (false, format!("{why} from the first rows of each partition, but {check}; left to DataFusion"), Some(f))
645 }
646 }
647 (AggregateChoice::Measured, None) => {
648 let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
649 decide(&self.op, self.inputs[0].as_ref(), &refs, &keys, rows, self.table_rows)
650 }
651 _ => (true, "aggregate_choice is ArrowMetal".to_string(), None),
652 };
653 self.mm.probe_time.add_duration(t.elapsed());
654 if let Some(e) = &estimate {
655 self.mm.groups_estimate.set(e.estimate as usize);
656 }
657 self.record(Decision::runtime_choice(&self.op, on_gpu, reason, rows, estimate));
658 if !on_gpu {
659 self.mm.handed_back.add(1);
660 return self.hand_back(c);
661 }
662 self.run_gpu(c).await
663 }
664
665 async fn run_join(self, streams: Vec<SendableRecordBatchStream>, left: usize) -> Result<Outcome> {
668 let t = std::time::Instant::now();
669 let total = streams.len();
670 let mut c = collect_streams(streams, &self.ctx, None).await?;
671 c.per_input = vec![left, total - left];
672 self.mm.input_time.add_duration(t.elapsed());
673 self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
674 if let Some(e) = &c.refused {
675 self.record(Decision::memory_hand_back(&self.op, "the input", e));
676 self.mm.handed_back.add(1);
677 return Ok(Outcome::Back(self.hand_back_plan(c)?));
678 }
679 self.run_gpu_outcome(c).await
680 }
681}
682
683enum Outcome {
685 Gpu(Vec<RecordBatch>, MemoryReservation),
687 Back(Arc<dyn ExecutionPlan>),
689}
690
691type SharedOutcome = futures::future::Shared<futures::future::BoxFuture<'static, std::result::Result<Arc<Outcome>, Arc<DataFusionError>>>>;
694
695type JoinSlot = Arc<Mutex<Option<(SharedOutcome, usize)>>>;
697
698fn partition_of(shared: SharedOutcome, p: usize, n: usize, ctx: Arc<TaskContext>) -> BoxStream<'static, Result<RecordBatch>> {
701 let s = futures::stream::once(shared).map(move |r| -> Result<BoxStream<'static, Result<RecordBatch>>> {
702 let out = r.map_err(DataFusionError::Shared)?;
703 match &*out {
704 Outcome::Gpu(batches, _) => {
705 let mine: Vec<RecordBatch> = batches.iter().skip(p).step_by(n).cloned().collect();
706 let held = Arc::clone(&out);
707 Ok(Box::pin(futures::stream::iter(mine.into_iter().map(Ok)).map(move |b| {
708 let _held = &held;
709 b
710 })))
711 }
712 Outcome::Back(plan) => Ok(Box::pin(plan.execute(p, Arc::clone(&ctx))?)),
714 }
715 });
716 Box::pin(s.try_flatten())
717}
718
719fn fit_partitions(plan: Arc<dyn ExecutionPlan>, n: usize) -> Result<Arc<dyn ExecutionPlan>> {
721 let have = plan.output_partitioning().partition_count();
722 Ok(if have == n {
723 plan
724 } else if n == 1 {
725 Arc::new(CoalescePartitionsExec::new(plan))
726 } else {
727 Arc::new(datafusion::physical_plan::repartition::RepartitionExec::try_new(plan, Partitioning::RoundRobinBatch(n))?)
728 })
729}
730
731impl MetalExec {
732 pub(crate) fn new(
735 op: MetalOp,
736 inputs: Vec<Arc<dyn ExecutionPlan>>,
737 original: Arc<dyn ExecutionPlan>,
738 (log, plan): (SharedLog, u64),
739 agg: AggSettings,
740 ) -> Self {
741 let AggSettings { choice, table_rows, rows_hint } = agg;
742 let mut eq = original.equivalence_properties().clone();
743 eq.clear_per_partition_constants();
744 let props = PlanProperties::new(
745 eq,
746 Partitioning::UnknownPartitioning(1),
747 EmissionType::Final,
748 Boundedness::Bounded,
749 );
750 Self {
751 op,
752 inputs,
753 original,
754 props: Arc::new(props),
755 log,
756 plan,
757 metrics: ExecutionPlanMetricsSet::new(),
758 choice,
759 table_rows,
760 rows_hint,
761 join_slot: Arc::new(Mutex::new(None)),
762 }
763 }
764
765 pub(crate) fn with_output_partitions(mut self, n: usize) -> Self {
769 if matches!(self.op, MetalOp::Join { .. }) && n > 1 {
770 let props = PlanProperties::new(
771 self.props.eq_properties.clone(),
772 Partitioning::UnknownPartitioning(n),
773 EmissionType::Final,
774 Boundedness::Bounded,
775 );
776 self.props = Arc::new(props);
777 }
778 self
779 }
780
781 fn job(&self, ctx: &Arc<TaskContext>) -> Job {
783 Job {
784 op: self.op.clone(),
785 inputs: self.inputs.clone(),
786 original: Arc::clone(&self.original),
787 out_schema: self.schema(),
788 log: Arc::clone(&self.log),
789 plan: self.plan,
790 choice: self.choice,
791 table_rows: self.table_rows,
792 rows_hint: self.rows_hint,
793 mm: self.phase_metrics(),
794 ctx: Arc::clone(ctx),
795 }
796 }
797
798 pub fn op(&self) -> &MetalOp {
800 &self.op
801 }
802
803 pub fn input(&self) -> &Arc<dyn ExecutionPlan> {
805 &self.inputs[0]
806 }
807
808 pub fn inputs(&self) -> &[Arc<dyn ExecutionPlan>] {
810 &self.inputs
811 }
812
813 fn phase_metrics(&self) -> PhaseMetrics {
817 let m = |name: &'static str| MetricBuilder::new(&self.metrics).subset_time(name, 0);
818 PhaseMetrics {
819 input_time: m("input_time"),
820 import_time: m("import_time"),
821 kernel_time: m("kernel_time"),
822 export_time: m("export_time"),
823 probe_time: m("probe_time"),
824 input_batches: MetricBuilder::new(&self.metrics).counter("input_batches", 0),
825 handed_back: MetricBuilder::new(&self.metrics).counter("handed_back", 0),
826 groups_estimate: MetricBuilder::new(&self.metrics).gauge("groups_estimate", 0),
827 }
828 }
829}
830
831pub(crate) fn confirm(
834 refs: &[&RecordBatch],
835 keys: &[usize],
836 rows: usize,
837 e: &crate::probe::GroupEstimate,
838) -> (bool, crate::probe::GroupEstimate) {
839 let f = crate::probe::estimate_up_to(refs, keys, None, Some(rows), CONFIRM_SAMPLE);
840 (f.low <= e.high && f.high >= e.low, f)
841}
842
843fn decide(
845 op: &MetalOp,
846 input: &dyn ExecutionPlan,
847 refs: &[&RecordBatch],
848 keys: &[usize],
849 rows: usize,
850 table_rows: Option<usize>,
851) -> (bool, String, Option<crate::probe::GroupEstimate>) {
852 let Some(shape) = crate::choice::shape(op, input) else {
853 return (false, "not an aggregate".into(), None);
854 };
855 let at = table_rows.unwrap_or(rows) as u64;
856 let settled = crate::choice::settled_for(&shape, at);
857 let e = crate::probe::estimate_up_to(refs, keys, Some(&settled), Some(rows), DECIDE_SAMPLE);
858 let how = if e.exact {
859 format!("{} groups, counted over all {} rows", e.estimate, e.rows)
860 } else {
861 format!("an estimated {} groups ({} to {}, Chao1 over a {}-row sample)", e.estimate, e.low, e.high, e.sample_rows)
862 };
863 let at_text = match table_rows {
864 Some(n) => format!("{} rows (table_rows; input {rows})", crate::choice::rows_text(n as u64)),
865 None => format!("{rows} rows"),
866 };
867 if !settled(e.low, e.high) {
868 return (
869 false,
870 format!("{how} at {at_text}: the range reaches group counts the table decides differently; left to DataFusion"),
871 Some(e),
872 );
873 }
874 let b = crate::choice::bucket(e.estimate, at);
875 match crate::choice::takes(&shape, b, at) {
876 Ok(why) => (true, format!("{how} at {at_text}: {why}"), Some(e)),
877 Err(why) => (false, format!("{how} at {at_text}: {why}; left to DataFusion"), Some(e)),
878 }
879}
880
881impl DisplayAs for MetalExec {
882 fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
883 let schema = self.inputs[0].schema();
884 let name = |i: usize| schema.field(i).name().clone();
885 match &self.op {
886 MetalOp::Join { how, left_keys, right_keys, projection, .. } => {
887 let right = self.inputs.get(1).map(|r| r.schema()).unwrap_or_else(|| Arc::clone(&schema));
888 let on: Vec<String> = left_keys
889 .iter()
890 .zip(right_keys)
891 .map(|(&l, &r)| format!("({}, {})", name(l), right.field(r).name()))
892 .collect();
893 write!(f, "MetalExec: join={how:?}, on=[{}], projection={projection:?}", on.join(", "))
894 }
895 MetalOp::Sort { keys, fetch } => {
896 let ks: Vec<String> = keys
897 .iter()
898 .map(|k| {
899 format!(
900 "{} {} NULLS {}",
901 name(k.column),
902 if k.descending { "DESC" } else { "ASC" },
903 if k.nulls_first { "FIRST" } else { "LAST" }
904 )
905 })
906 .collect();
907 write!(f, "MetalExec: sort=[{}]", ks.join(", "))?;
908 if let Some(n) = fetch {
909 write!(f, ", fetch={n}")?;
910 }
911 Ok(())
912 }
913 MetalOp::Aggregate { keys, aggs } => {
914 let ks: Vec<String> = keys.iter().map(|&k| name(k)).collect();
915 let asx: Vec<String> = aggs
916 .iter()
917 .map(|a| format!("{:?}({})", a.kind, a.column.map(name).unwrap_or_else(|| "*".into())))
918 .collect();
919 write!(f, "MetalExec: group_by=[{}], aggr=[{}]", ks.join(", "), asx.join(", "))
920 }
921 MetalOp::Filter { predicate, projection, .. } => {
922 write!(f, "MetalExec: filter={predicate}")?;
923 if let Some(p) = projection {
924 write!(f, ", projection={p:?}")?;
925 }
926 Ok(())
927 }
928 }
929 }
930}
931
932fn replace_leaf(
934 plan: &Arc<dyn ExecutionPlan>,
935 old: &Arc<dyn ExecutionPlan>,
936 new: &Arc<dyn ExecutionPlan>,
937) -> Result<Arc<dyn ExecutionPlan>> {
938 if Arc::ptr_eq(plan, old) {
939 return Ok(Arc::clone(new));
940 }
941 Ok(Arc::clone(plan)
942 .transform_down(|n| {
943 if Arc::ptr_eq(&n, old) {
944 Ok(Transformed::new(Arc::clone(new), true, TreeNodeRecursion::Jump))
945 } else {
946 Ok(Transformed::no(n))
947 }
948 })?
949 .data)
950}
951
952impl ExecutionPlan for MetalExec {
953 fn name(&self) -> &str {
954 "MetalExec"
955 }
956
957 fn properties(&self) -> &Arc<PlanProperties> {
958 &self.props
959 }
960
961 fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
962 self.inputs.iter().collect()
963 }
964
965 fn benefits_from_input_partitioning(&self) -> Vec<bool> {
966 vec![false; self.inputs.len()]
967 }
968
969 fn apply_expressions(
970 &self,
971 _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
972 ) -> Result<TreeNodeRecursion> {
973 Ok(TreeNodeRecursion::Continue)
974 }
975
976 fn with_new_children(
977 self: Arc<Self>,
978 children: Vec<Arc<dyn ExecutionPlan>>,
979 ) -> Result<Arc<dyn ExecutionPlan>> {
980 if children.len() != self.inputs.len() {
981 return internal_err!("MetalExec takes {} children, got {}", self.inputs.len(), children.len());
982 }
983 let mut original = Arc::clone(&self.original);
984 for (old, new) in self.inputs.iter().zip(&children) {
985 original = replace_leaf(&original, old, new)?;
986 }
987 Ok(Arc::new(MetalExec {
988 op: self.op.clone(),
989 inputs: children,
990 original,
991 props: Arc::clone(&self.props),
992 log: Arc::clone(&self.log),
993 plan: self.plan,
994 metrics: ExecutionPlanMetricsSet::new(),
995 choice: self.choice,
996 table_rows: self.table_rows,
997 rows_hint: self.rows_hint,
998 join_slot: Arc::new(Mutex::new(None)),
999 }))
1000 }
1001
1002 fn execute(&self, partition: usize, ctx: Arc<TaskContext>) -> Result<SendableRecordBatchStream> {
1003 let n = self.props.partitioning.partition_count();
1004 if partition >= n {
1005 return internal_err!("MetalExec has {n} partition(s), asked for {partition}");
1006 }
1007 let schema = self.schema();
1008 if matches!(self.op, MetalOp::Join { .. }) {
1009 let mut slot = self.join_slot.lock().unwrap_or_else(|p| p.into_inner());
1012 let fresh = slot.as_ref().is_none_or(|(_, taken)| *taken >= n);
1013 if fresh {
1014 let job = self.job(&ctx);
1015 let mut streams = Vec::new();
1016 let mut left = 0;
1017 for (i, input) in self.inputs.iter().enumerate() {
1018 for p in 0..input.output_partitioning().partition_count() {
1019 streams.push(input.execute(p, Arc::clone(&ctx))?);
1020 }
1021 if i == 0 {
1022 left = streams.len();
1023 }
1024 }
1025 let fut: futures::future::BoxFuture<'static, std::result::Result<Arc<Outcome>, Arc<DataFusionError>>> =
1026 Box::pin(async move {
1027 let out = match job.run_join(streams, left).await {
1028 Ok(Outcome::Back(plan)) => fit_partitions(plan, n).map(Outcome::Back),
1029 other => other,
1030 };
1031 out.map(Arc::new).map_err(Arc::new)
1032 });
1033 *slot = Some((fut.shared(), 0));
1034 }
1035 let Some((shared, taken)) = slot.as_mut() else {
1036 return internal_err!("MetalExec join slot empty");
1037 };
1038 *taken += 1;
1039 let shared = shared.clone();
1040 if *taken >= n {
1041 *slot = None;
1043 }
1044 let s = partition_of(shared, partition, n, ctx);
1045 return Ok(Box::pin(RecordBatchStreamAdapter::new(schema, s)));
1046 }
1047 let job = self.job(&ctx);
1048 let fut: futures::future::BoxFuture<'static, Result<BoxStream<'static, Result<RecordBatch>>>> =
1049 if matches!(self.op, MetalOp::Aggregate { .. }) {
1050 let input = &self.inputs[0];
1052 let mut streams = Vec::new();
1053 for p in 0..input.output_partitioning().partition_count() {
1054 streams.push(input.execute(p, Arc::clone(&ctx))?);
1055 }
1056 Box::pin(job.run_aggregate(streams))
1057 } else {
1058 let input = &self.inputs[0];
1060 let source: Arc<dyn ExecutionPlan> = if input.output_partitioning().partition_count() > 1 {
1061 Arc::new(CoalescePartitionsExec::new(Arc::clone(input)))
1062 } else {
1063 Arc::clone(input)
1064 };
1065 let stream = source.execute(0, Arc::clone(&ctx))?;
1066 Box::pin(job.run_sort_or_filter(stream))
1067 };
1068 let s = futures::stream::once(fut).try_flatten();
1069 Ok(Box::pin(RecordBatchStreamAdapter::new(schema, s)))
1070 }
1071
1072 fn metrics(&self) -> Option<MetricsSet> {
1073 Some(self.metrics.clone_inner())
1074 }
1075}
1076
1077#[cfg(test)]
1078mod tests {
1079 use super::*;
1080 use arrow::array::Int64Array;
1081 use arrow::datatypes::{Field, Schema};
1082
1083 fn batches(keys: Vec<i64>) -> Vec<RecordBatch> {
1084 let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int64, false)]));
1085 keys.chunks(8192)
1086 .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(Int64Array::from(c.to_vec()))]).unwrap())
1087 .collect()
1088 }
1089
1090 #[test]
1093 fn a_prefix_that_under_counts_is_not_confirmed() {
1094 let n = 1_200_000usize;
1095 let head = 262_144usize;
1096 let keys: Vec<i64> = (0..n).map(|i| if i < head { (i % 20_000) as i64 } else { i as i64 }).collect();
1097 let all = batches(keys);
1098 let refs: Vec<&RecordBatch> = all.iter().collect();
1099 let prefix: Vec<&RecordBatch> = refs[..head / 8192].to_vec();
1100 let e = crate::probe::estimate_up_to(&prefix, &[0], None, Some(n), 8_192);
1101 assert!(e.high < 40_000, "{e:?}");
1102 let (agree, f) = confirm(&refs, &[0], n, &e);
1103 assert!(!agree, "prefix {e:?} whole {f:?}");
1104 let mut shuffled: Vec<i64> = (0..n).map(|i| if i < head { (i % 20_000) as i64 } else { i as i64 }).collect();
1106 let mut state = 0x1234_5678u64;
1107 for i in (1..n).rev() {
1108 state ^= state << 13;
1109 state ^= state >> 7;
1110 state ^= state << 17;
1111 shuffled.swap(i, (state % (i as u64 + 1)) as usize);
1112 }
1113 let all = batches(shuffled);
1114 let refs: Vec<&RecordBatch> = all.iter().collect();
1115 let prefix: Vec<&RecordBatch> = refs[..head / 8192].to_vec();
1116 let e = crate::probe::estimate_up_to(&prefix, &[0], None, Some(n), 8_192);
1117 let (agree, f) = confirm(&refs, &[0], n, &e);
1118 assert!(agree, "prefix {e:?} whole {f:?}");
1119 }
1120}