diff --git a/src/distributed_planner/task_estimator.rs b/src/distributed_planner/task_estimator.rs index fbfc7b09..4232b5c5 100644 --- a/src/distributed_planner/task_estimator.rs +++ b/src/distributed_planner/task_estimator.rs @@ -275,10 +275,16 @@ impl TaskEstimator for FileScanConfigTaskEstimator { let file_scan: &FileScanConfig = dse.data_source().as_any().downcast_ref()?; let mut file_scan_template = file_scan.clone(); - let mut file_groups = VecDeque::with_capacity(file_scan.file_groups.len() * task_count); - for file_group in file_scan_template.file_groups.drain(..) { - file_groups.extend(file_group.split_files(task_count)); - } + let input_group_count = file_scan_template.file_groups.len().max(1); + let all_partitioned_files = file_scan_template + .file_groups + .iter() + .flat_map(|file_group| file_group.iter().cloned()) + .collect::>(); + file_scan_template.file_groups.clear(); + let rebalanced = + rebalance_round_robin(all_partitioned_files, input_group_count * task_count); + let mut file_groups: VecDeque = rebalanced.into_iter().map(Into::into).collect(); let expected_partitions = plan.output_partitioning().partition_count(); let dle = DistributedLeafExec::new( @@ -299,6 +305,17 @@ impl TaskEstimator for FileScanConfigTaskEstimator { } } +fn rebalance_round_robin(items: Vec, target_groups: usize) -> Vec> { + let target_groups = target_groups.min(items.len()); + let mut groups = (0..target_groups) + .map(|_| Vec::new()) + .collect::>>(); + for (idx, item) in items.into_iter().enumerate() { + groups[idx % target_groups].push(item); + } + groups +} + /// Tries multiple user-provided [TaskEstimator]s until one returns an estimation. If none /// returns an estimation, a set of default [TaskEstimation] implementations is tried. Right /// now the only default [TaskEstimation] is [FileScanConfigTaskEstimator]. @@ -402,6 +419,22 @@ mod tests { Ok(()) } + #[test] + fn test_rebalance_round_robin_fixes_group_boundary_skew() { + let items = (0..8).collect::>(); + let groups = rebalance_round_robin(items, 5); + let sizes = groups.iter().map(Vec::len).collect::>(); + assert_eq!(sizes, vec![2, 2, 2, 1, 1]); + } + + #[test] + fn test_rebalance_round_robin_caps_partitions_to_file_count() { + let items = vec![10, 20, 30]; + let groups = rebalance_round_robin(items, 5); + let sizes = groups.iter().map(Vec::len).collect::>(); + assert_eq!(sizes, vec![1, 1, 1]); + } + impl CombinedTaskEstimator { fn push(&mut self, value: impl TaskEstimator + Send + Sync + 'static) { self.user_provided.push(Arc::new(value));