From d2d664c9ff24ea5443c46a348e3f1e891bb17a7a Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Wed, 16 Sep 2026 21:35:56 -0400 Subject: [PATCH] Bound runtime collection launch specialization --- src/LocalMath.jl | 1 + src/execution/candidate_grouping.jl | 5 +-- src/execution/collect_physical_support.jl | 12 +++--- src/execution/collect_stage.jl | 36 +++++++++-------- src/execution/keyed_reduce_stage.jl | 46 +++++++++++----------- src/execution/launch_support.jl | 7 ++++ src/execution/ordered_fold_stage.jl | 48 +++++++---------------- test/runtests.jl | 1 + test/test_launch_contract.jl | 48 +++++++++++++++++++++++ 9 files changed, 123 insertions(+), 81 deletions(-) create mode 100644 src/execution/launch_support.jl create mode 100644 test/test_launch_contract.jl diff --git a/src/LocalMath.jl b/src/LocalMath.jl index a73745c..da9c9a6 100644 --- a/src/LocalMath.jl +++ b/src/LocalMath.jl @@ -71,6 +71,7 @@ include("execution.jl") include("inspection.jl") include("execution/mechanism_support.jl") +include("execution/launch_support.jl") include("execution/validation_support.jl") include("execution/relation_views.jl") include("structural_binding.jl") diff --git a/src/execution/candidate_grouping.jl b/src/execution/candidate_grouping.jl index 6b8524f..1ea29d7 100644 --- a/src/execution/candidate_grouping.jl +++ b/src/execution/candidate_grouping.jl @@ -318,9 +318,8 @@ end function _group_destinations!(backend, grouping::_DestinationGrouping) local_extent = max(cld(Int(grouping.sort_capacity), _DESTINATION_GROUP_BLOCK), 1) * _DESTINATION_GROUP_BLOCK - _destination_grouping_local_sort_kernel!( - backend, _DESTINATION_GROUP_BLOCK, local_extent)(grouping; - ndrange = local_extent) + _launch_1d!(_destination_grouping_local_sort_kernel!, + backend, local_extent, Val(_DESTINATION_GROUP_BLOCK), grouping) source, destination = grouping.order_a, grouping.order_b width = _DESTINATION_GROUP_BLOCK while width < grouping.sort_capacity diff --git a/src/execution/collect_physical_support.jl b/src/execution/collect_physical_support.jl index b5afea8..fc00c7e 100644 --- a/src/execution/collect_physical_support.jl +++ b/src/execution/collect_physical_support.jl @@ -197,8 +197,8 @@ function _compacted_launch_prefix_scan!(backend, item_counts, prefix_storage, output = _compacted_scan_level(prefix_storage, prefix_offset, current) sums = _compacted_scan_level(sums_storage, sums_offset, blocks) extent = blocks * _COMPACTED_BLOCK - _compacted_scan_block_kernel!(backend, _COMPACTED_BLOCK, extent)( - item_counts, output, sums, Int32(current); ndrange = extent) + _launch_1d!(_compacted_scan_block_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), item_counts, output, sums, Int32(current)) prefix_offset += current sums_offset += blocks current = blocks @@ -208,8 +208,8 @@ function _compacted_launch_prefix_scan!(backend, item_counts, prefix_storage, output = _compacted_scan_level(prefix_storage, prefix_offset, current) sums = _compacted_scan_level(sums_storage, sums_offset, blocks) extent = blocks * _COMPACTED_BLOCK - _compacted_scan_block_kernel!(backend, _COMPACTED_BLOCK, extent)( - input, output, sums, Int32(length(input)); ndrange = extent) + _launch_1d!(_compacted_scan_block_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), input, output, sums, Int32(length(input))) prefix_offset += current sums_offset += blocks current <= _COMPACTED_BLOCK && break @@ -229,8 +229,8 @@ function _compacted_launch_prefix_scan!(backend, item_counts, prefix_storage, prefix = _compacted_scan_level(prefix_storage, child_offset, size) parent = _compacted_scan_level(prefix_storage, parent_offset, parent_size) extent = max(length(prefix), 1) - _compacted_scan_add_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - prefix, parent, Int32(length(prefix)); ndrange = extent) + _launch_1d!(_compacted_scan_add_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), prefix, parent, Int32(length(prefix))) end return nothing end diff --git a/src/execution/collect_stage.jl b/src/execution/collect_stage.jl index 3c5b4d0..77030ad 100644 --- a/src/execution/collect_stage.jl +++ b/src/execution/collect_stage.jl @@ -534,34 +534,37 @@ function _collect_launch_order!(backend, plan, workspace) items = div(candidates, _collect_width(plan)) prefix = _compacted_scan_level(workspace.prefix, 0, items) extent = max(items, 1) - _compacted_scatter_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + _launch_1d!(_compacted_scatter_kernel!, backend, extent, Val(_COMPACTED_BLOCK), workspace.valid, workspace.item_counts, prefix, workspace.order_a, workspace.positions, workspace.count, Val(_collect_width(plan)), - Int32(items); ndrange = extent) + Int32(items)) plan.sort_required || return nothing local_extent = max(cld(candidates, _COMPACTED_BLOCK), 1) * _COMPACTED_BLOCK - _compacted_local_bitonic_kernel!(backend, _COMPACTED_BLOCK, local_extent)( - plan, workspace, workspace.count, Int32(candidates); ndrange = local_extent) + _launch_1d!(_compacted_local_bitonic_kernel!, backend, local_extent, + Val(_COMPACTED_BLOCK), + plan, workspace, workspace.count, Int32(candidates)) width, to_b = _COMPACTED_BLOCK, true while width < candidates source, destination = to_b ? (workspace.order_a, workspace.order_b) : (workspace.order_b, workspace.order_a) - _compacted_merge_kernel!(backend, min(candidates, _COMPACTED_BLOCK), candidates)( + _launch_1d!(_compacted_merge_kernel!, backend, candidates, + Val(_COMPACTED_BLOCK), plan, workspace, source, destination, workspace.count, Int32(width), - Int32(candidates); ndrange = candidates) + Int32(candidates)) width *= 2 to_b = !to_b end if _is_grouped(plan.groups) groups = Int(plan.groups.count) + 1 - _compacted_directory_kernel!(backend, min(groups, _COMPACTED_BLOCK), groups)( - workspace, _collect_final_order(plan, workspace), Int32(plan.groups.count); - ndrange = groups) + _launch_1d!(_compacted_directory_kernel!, backend, groups, + Val(_COMPACTED_BLOCK), + workspace, _collect_final_order(plan, workspace), Int32(plan.groups.count)) end if _is_canonical_order(plan.order) extent = max(candidates, 1) - _compacted_validate_order_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - plan, workspace, _collect_final_order(plan, workspace); ndrange = extent) + _launch_1d!(_compacted_validate_order_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + plan, workspace, _collect_final_order(plan, workspace)) end return nothing end @@ -574,9 +577,9 @@ function _collect_publish_chunk!(backend, plans::Tuple, workspaces::Tuple, end groupeds = map(plan -> _collect_grouped(plan.groups), plans) extent = maximum(Int, extents) - _compacted_publish_ports_kernel!(backend, - min(extent, _COMPACTED_BLOCK), extent)(storages, workspaces, gate, - groupeds, extents; ndrange = extent) + _launch_1d!(_compacted_publish_ports_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + storages, workspaces, gate, groupeds, extents) return nothing end @@ -625,10 +628,11 @@ function _execute_collect_stage!(prepared::_CollectStagePreparation, predecessor_statuses = (relation_guard, predecessors...) extent = max(Int(execution.stage.source_count), maximum((Int(plan.candidate_count) for plan in execution.plans); init = 0), 1) - _collect_stage_reset_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + _launch_1d!(_collect_stage_reset_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), execution.workspaces, execution.status, execution.gate, prepared.validation, lease_index, - Int32(extent); ndrange = extent) + Int32(extent)) _launch_stage_relation_receipt!(backend, relation_guard, prepared.validation, program_validation, lease_index) _collect_stage_evaluate_kernel!(backend)(qualified, execution.plans, diff --git a/src/execution/keyed_reduce_stage.jl b/src/execution/keyed_reduce_stage.jl index 2d882a4..ac12cf6 100644 --- a/src/execution/keyed_reduce_stage.jl +++ b/src/execution/keyed_reduce_stage.jl @@ -368,22 +368,22 @@ function _keyed_reduce_launch_order!(backend, bounds, key_order, state) _compacted_launch_prefix_scan!(backend, state.item_counts, state.prefix, state.sums) prefix = _compacted_scan_level(state.prefix, 0, candidates) - _compacted_scatter_kernel!(backend, - min(max(candidates, 1), _COMPACTED_BLOCK), max(candidates, 1))( + _launch_1d!(_compacted_scatter_kernel!, backend, max(candidates, 1), + Val(_COMPACTED_BLOCK), state.valid, state.item_counts, prefix, state.order_a, - state.positions, state.count, Val(1), Int32(candidates); - ndrange = max(candidates, 1)) + state.positions, state.count, Val(1), Int32(candidates)) local_extent = max(cld(candidates, _COMPACTED_BLOCK), 1) * _COMPACTED_BLOCK - _compacted_local_bitonic_kernel!(backend, _COMPACTED_BLOCK, local_extent)( - key_order, state, state.count, Int32(candidates); - ndrange = local_extent) + _launch_1d!(_compacted_local_bitonic_kernel!, backend, local_extent, + Val(_COMPACTED_BLOCK), + key_order, state, state.count, Int32(candidates)) width, to_b = _COMPACTED_BLOCK, true while width < candidates source, destination = to_b ? (state.order_a, state.order_b) : (state.order_b, state.order_a) - _compacted_merge_kernel!(backend, min(candidates, _COMPACTED_BLOCK), candidates)( + _launch_1d!(_compacted_merge_kernel!, backend, candidates, + Val(_COMPACTED_BLOCK), key_order, state, source, destination, state.count, - Int32(width), Int32(candidates); ndrange = candidates) + Int32(width), Int32(candidates)) width *= 2 to_b = !to_b end @@ -571,9 +571,9 @@ function _execute_keyed_reduce_stage!(prepared::_KeyedReduceStagePreparation, _stage_runtime_parameters(parameters, execution.stage)) statuses = (relation_guard, predecessors...) extent = max(Int(bounds.candidate_count), 1) - _keyed_reduce_reset_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - bounds, states.reset, execution.storage, prepared.validation, lease_index; - ndrange = extent) + _launch_1d!(_keyed_reduce_reset_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + bounds, states.reset, execution.storage, prepared.validation, lease_index) _launch_stage_relation_receipt!(backend, relation_guard, prepared.validation, program_validation, lease_index) _keyed_reduce_evaluate_kernel!(backend)(qualified, plan.emission, @@ -583,19 +583,20 @@ function _execute_keyed_reduce_stage!(prepared::_KeyedReduceStagePreparation, _keyed_reduce_launch_order!(backend, bounds, plan.key_order, states.ordering) order = _keyed_reduce_final_order(bounds, states.ordering) - _keyed_reduce_segments_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + _launch_1d!(_keyed_reduce_segments_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), plan.key_order, states.segment, order, plan.fold.prior_capacity, - bounds.candidate_count; ndrange = extent) + bounds.candidate_count) _compacted_launch_prefix_scan!(backend, states.segment.item_counts, states.segment.prefix, states.ordering.sums) _keyed_reduce_unique_count_kernel!(backend, 1, 1)(states.segment; ndrange = 1) - _keyed_reduce_clear_counts_kernel!(backend, - min(extent, _COMPACTED_BLOCK), extent)(states.fold, - bounds.candidate_count; - ndrange = extent) - _keyed_reduce_fold_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - plan.fold, states.fold, order, bounds.candidate_count; ndrange = extent) + _launch_1d!(_keyed_reduce_clear_counts_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + states.fold, bounds.candidate_count) + _launch_1d!(_keyed_reduce_fold_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + plan.fold, states.fold, order, bounds.candidate_count) _compacted_launch_prefix_scan!(backend, states.fold.item_counts, states.fold.prefix, states.ordering.sums) _keyed_reduce_final_count_kernel!(backend, 1, 1)( @@ -604,7 +605,8 @@ function _execute_keyed_reduce_stage!(prepared::_KeyedReduceStagePreparation, execution.status, prepared.validation, program_validation, statuses, lease_index; ndrange = 1) - _keyed_reduce_publish_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - plan.publication, states.publication, execution.storage; ndrange = extent) + _launch_1d!(_keyed_reduce_publish_kernel!, backend, extent, + Val(_COMPACTED_BLOCK), + plan.publication, states.publication, execution.storage) return prepared end diff --git a/src/execution/launch_support.jl b/src/execution/launch_support.jl new file mode 100644 index 0000000..45e7fd7 --- /dev/null +++ b/src/execution/launch_support.jl @@ -0,0 +1,7 @@ +# KernelAbstractions specializes a constructor-provided ndrange on its value. +# The owning operation supplies its semantic workgroup; extent stays runtime data. +@inline function _launch_1d!( + kernel, backend, extent::Int, ::Val{Workgroup}, arguments... + ) where {Workgroup} + return kernel(backend, Workgroup)(arguments...; ndrange = max(extent, 1)) +end diff --git a/src/execution/ordered_fold_stage.jl b/src/execution/ordered_fold_stage.jl index 6605551..0e7bcbd 100644 --- a/src/execution/ordered_fold_stage.jl +++ b/src/execution/ordered_fold_stage.jl @@ -353,15 +353,11 @@ function _ordered_fold_stage_launch_order!(backend, order_law, workspace) while width <= extent distance = width >>> 1 while distance >= 1 - _ordered_fold_stage_bitonic_kernel!( - backend, - min(extent, _ORDERED_FOLD_BLOCK), extent - )( + _launch_1d!(_ordered_fold_stage_bitonic_kernel!, backend, extent, + Val(_ORDERED_FOLD_BLOCK), order_law, workspace.values, workspace.order, - Int32(distance), Int32(width), Int32(extent); - ndrange = extent - ) + Int32(distance), Int32(width), Int32(extent)) distance >>>= 1 end width <<= 1 @@ -592,49 +588,33 @@ function _execute_ordered_fold_stage!( predecessors = prefix, program_validation, ) order_extent = length(prepared.workspace.order) - _ordered_fold_stage_reset_kernel!( - prepared.backend, - min(order_extent, _ORDERED_FOLD_BLOCK), order_extent - )( + _launch_1d!(_ordered_fold_stage_reset_kernel!, + prepared.backend, order_extent, Val(_ORDERED_FOLD_BLOCK), prepared.workspace.order, prepared.workspace.status, - prepared.workspace.validation, lease_index, Int32(order_extent); - ndrange = order_extent - ) + prepared.workspace.validation, lease_index, Int32(order_extent)) _launch_stage_relation_receipt!( prepared.backend, relation_guard, prepared.validation, program_validation, lease_index ) source_extent = max(Int(prepared.stage.source_count), 1) - _ordered_fold_stage_evaluate_kernel!( - prepared.backend, - min(source_extent, _ORDERED_FOLD_BLOCK), source_extent - )( + _launch_1d!(_ordered_fold_stage_evaluate_kernel!, + prepared.backend, source_extent, Val(_ORDERED_FOLD_BLOCK), qualified, prepared.workspace, lease_index, prefix, - prepared.stage.source_count; ndrange = source_extent - ) + prepared.stage.source_count) _ordered_fold_stage_launch_order!( prepared.backend, law.order, prepared.workspace ) state_extent = prepared.state_extent initialize_extent = max(source_extent, Int(state_extent), 1) - _ordered_fold_stage_validate_initialize_kernel!( - prepared.backend, - min(initialize_extent, _ORDERED_FOLD_BLOCK), initialize_extent - )( + _launch_1d!(_ordered_fold_stage_validate_initialize_kernel!, + prepared.backend, initialize_extent, Val(_ORDERED_FOLD_BLOCK), law.order, state, prepared.stage.fields, prepared.workspace, - lease_index, prefix, prepared.stage.source_count, state_extent; - ndrange = initialize_extent - ) + lease_index, prefix, prepared.stage.source_count, state_extent) _ordered_fold_stage_kernel!(prepared.backend)(recurrence; ndrange = 1) state_launch = max(Int(state_extent), 1) - _ordered_fold_stage_finalize_kernel!( - prepared.backend, - min(state_launch, _ORDERED_FOLD_BLOCK), state_launch - )( - finalization; - ndrange = state_launch - ) + _launch_1d!(_ordered_fold_stage_finalize_kernel!, + prepared.backend, state_launch, Val(_ORDERED_FOLD_BLOCK), finalization) return prepared end diff --git a/test/runtests.jl b/test/runtests.jl index bf275f1..eca528e 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -16,6 +16,7 @@ const LOCALMATH_INCLUDED_TESTS = ( "test_unique_stage.jl", "test_stage_program_lifecycle.jl", "test_execution_receipts.jl", + "test_launch_contract.jl", "test_reduce_stage.jl", "test_reduction_control.jl", "test_resolve_stage.jl", diff --git a/test/test_launch_contract.jl b/test/test_launch_contract.jl new file mode 100644 index 0000000..d5a045f --- /dev/null +++ b/test/test_launch_contract.jl @@ -0,0 +1,48 @@ +using Statistics: median + +function _baseline_scan_add!(backend, prefix, parent, extent::Int) + event = LocalMath._compacted_scan_add_kernel!( + backend, min(extent, 256), extent + )(prefix, parent, Int32(extent); ndrange = extent) + event === nothing || wait(event) + KernelAbstractions.synchronize(backend) + return prefix +end + +function _bounded_scan_add!(backend, prefix, parent, extent::Int) + event = LocalMath._launch_1d!( + LocalMath._compacted_scan_add_kernel!, backend, extent, Val(256), + prefix, parent, Int32(extent)) + event === nothing || wait(event) + KernelAbstractions.synchronize(backend) + return prefix +end + +function _warm_allocation_samples!(launch!, backend, prefix, parent, extent) + launch!(backend, prefix, parent, extent) + return map(1:7) do _ + GC.gc() + @allocated launch!(backend, prefix, parent, extent) + end +end + +@testset "bounded physical launch contract" begin + backend = KernelAbstractions.CPU() + for extent in (16, 24, 32, 33, 64, 65, 128, 129, 256, 300) + baseline = fill(Int32(1), extent) + bounded = copy(baseline) + parent = fill(Int32(2), max(cld(extent, 256), 1)) + _baseline_scan_add!(backend, baseline, parent, extent) + _bounded_scan_add!(backend, bounded, parent, extent) + @test bounded == baseline + end + + baseline = fill(Int32(1), 24) + bounded = copy(baseline) + parent = fill(Int32(0), 1) + baseline_bytes = _warm_allocation_samples!( + _baseline_scan_add!, backend, baseline, parent, 24) + bounded_bytes = _warm_allocation_samples!( + _bounded_scan_add!, backend, bounded, parent, 24) + @test median(bounded_bytes) <= median(baseline_bytes) +end