Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/LocalMath.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
5 changes: 2 additions & 3 deletions src/execution/candidate_grouping.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
12 changes: 6 additions & 6 deletions src/execution/collect_physical_support.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down
36 changes: 20 additions & 16 deletions src/execution/collect_stage.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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,
Expand Down
46 changes: 24 additions & 22 deletions src/execution/keyed_reduce_stage.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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)(
Expand All @@ -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
7 changes: 7 additions & 0 deletions src/execution/launch_support.jl
Original file line number Diff line number Diff line change
@@ -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
48 changes: 14 additions & 34 deletions src/execution/ordered_fold_stage.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
1 change: 1 addition & 0 deletions test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
48 changes: 48 additions & 0 deletions test/test_launch_contract.jl
Original file line number Diff line number Diff line change
@@ -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