From 36ff6aab67c453fbef61d069be7fbd8d15df223e Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Thu, 17 Sep 2026 01:31:38 -0400 Subject: [PATCH] Execute source-order folds without sorting --- src/execution/ordered_fold_stage.jl | 94 ++++++++--- src/execution/program_inspection.jl | 25 +-- .../ordered_fold_source_order_contracts.jl | 153 ++++++++++++++++++ test/metal/ordered_fold_control.jl | 2 + test/test_ordered_fold_stage_execution.jl | 4 + 5 files changed, 241 insertions(+), 37 deletions(-) create mode 100644 test/fixtures/ordered_fold_source_order_contracts.jl diff --git a/src/execution/ordered_fold_stage.jl b/src/execution/ordered_fold_stage.jl index 0e7bcbd..3d2886c 100644 --- a/src/execution/ordered_fold_stage.jl +++ b/src/execution/ordered_fold_stage.jl @@ -8,11 +8,20 @@ const _ORDERED_FOLD_BLOCK = 256 struct _OrderedFoldStageWorkspace{V, O, S, X, R, T} values::V; order::O; status::S; validation::X; state::R; tree::T end -struct _OrderedFoldRecurrenceStage{F, A, P, G} +struct _DirectSparseSourceTraversal end +struct _CompactedPrefixTraversal end + +_ordered_fold_recurrence_traversal(::_SourceOrder) = + _DirectSparseSourceTraversal() +_ordered_fold_recurrence_traversal(::_PreparedCanonicalBy) = + _CompactedPrefixTraversal() + +struct _OrderedFoldRecurrenceStage{F, A, P, G, T} fields::F accesses::A prefix::P gate::G + traversal::T source_count::Int32 end struct _OrderedFoldStagePreparation{B, S, W, V} @@ -347,6 +356,8 @@ end end end +_ordered_fold_stage_launch_order!(backend, ::_SourceOrder, workspace) = nothing + function _ordered_fold_stage_launch_order!(backend, order_law, workspace) extent = length(workspace.order) width = 2 @@ -365,8 +376,61 @@ function _ordered_fold_stage_launch_order!(backend, order_law, workspace) return nothing end -@inline function _ordered_fold_stage_execute!(run) +@inline function _ordered_fold_stage_apply_item!( + run, accumulator, item::Int32, position::Int32 + ) workspace = run.workspace + step = run.transition( + accumulator, + @inbounds(workspace.values[item]), item, + _stage_reads( + run.stage, item, + _OrderedFoldEvaluationValidation(workspace.status) + ) + ) + _ordered_fold_stage_success(run) || return false + code, component_index, witness = _ordered_fold_stage_validate_step!( + run.state, workspace.state, step + ) + if code != 0 + _ordered_fold_stage_fail!( + run, code, component_index, item, position, witness + ) + return false + end + _ordered_fold_stage_apply_step!(run.state, workspace.state, step) + return !step.halt +end + +@inline function _ordered_fold_stage_recur!( + run, ::_DirectSparseSourceTraversal, accumulator + ) + position = Int32(0) + for source_item in Int32(1):run.stage.source_count + item = @inbounds run.workspace.order[source_item] + item == 0 && continue + position += Int32(1) + _ordered_fold_stage_apply_item!( + run, accumulator, item, position + ) || return + end + return +end + +@inline function _ordered_fold_stage_recur!( + run, ::_CompactedPrefixTraversal, accumulator + ) + for position in Int32(1):run.stage.source_count + item = @inbounds run.workspace.order[position] + item == 0 && break + _ordered_fold_stage_apply_item!( + run, accumulator, item, position + ) || return + end + return +end + +@inline function _ordered_fold_stage_execute!(run) _ordered_fold_stage_prefix_ok(run) || return gate = _stage_gate_open( run.stage.gate, run.stage, @@ -381,30 +445,7 @@ end return _ordered_fold_stage_fail!(run, Int32(_CANDIDATE_STATUS_INVALID_CONTROL)) _ordered_fold_stage_success(run) || return accumulator = _prepared_fold_accumulator(run.workspace.state) - for position in Int32(1):run.stage.source_count - item = @inbounds workspace.order[position] - item == 0 && break - step = run.transition( - accumulator, - @inbounds(workspace.values[item]), item, - _stage_reads( - run.stage, item, - _OrderedFoldEvaluationValidation(workspace.status) - ) - ) - _ordered_fold_stage_success(run) || return - code, component_index, witness = - _ordered_fold_stage_validate_step!( - run.state, run.workspace.state, step - ) - code == 0 || return _ordered_fold_stage_fail!( - run, code, - component_index, item, position, witness - ) - _ordered_fold_stage_apply_step!(run.state, run.workspace.state, step) - step.halt && return - end - return + return _ordered_fold_stage_recur!(run, run.stage.traversal, accumulator) end @generated function _ordered_fold_stage_commit_item!( @@ -568,6 +609,7 @@ function _execute_ordered_fold_stage!( recurrence_stage = _OrderedFoldRecurrenceStage( prepared.stage.fields, prepared.stage.accesses, prepared.stage.control.prefix, prepared.stage.control.gate, + _ordered_fold_recurrence_traversal(law.order), prepared.stage.source_count ) boundary_parameters = (; diff --git a/src/execution/program_inspection.jl b/src/execution/program_inspection.jl index 4893e86..c54985b 100644 --- a/src/execution/program_inspection.jl +++ b/src/execution/program_inspection.jl @@ -432,19 +432,22 @@ function _planned_stage_phases(entry::_StageLoweringEntry{ phases = Any[_phase_fact(:ordered_fold_reset)] append!(phases, _planned_relation_phases(entry)) push!(phases, _phase_fact(:ordered_fold_evaluate)) - extent = nextpow(2, max(Int(entry.admission.stage.source_count), 1)) - bitonic = 0 - width = 2 - while width <= extent - distance = width >>> 1 - while distance >= 1 - bitonic += 1 - distance >>>= 1 + order = only(entry.admission.stage.publications).law.order + if !(order isa _SourceOrder) + extent = nextpow(2, max(Int(entry.admission.stage.source_count), 1)) + bitonic = 0 + width = 2 + while width <= extent + distance = width >>> 1 + while distance >= 1 + bitonic += 1 + distance >>>= 1 + end + width <<= 1 end - width <<= 1 + bitonic == 0 || push!(phases, + _phase_fact(:ordered_fold_bitonic, bitonic)) end - bitonic == 0 || push!(phases, - _phase_fact(:ordered_fold_bitonic, bitonic)) append!(phases, (_phase_fact(:ordered_fold_validate_initialize), _phase_fact(:ordered_fold_apply), _phase_fact(:ordered_fold_finalize))) diff --git a/test/fixtures/ordered_fold_source_order_contracts.jl b/test/fixtures/ordered_fold_source_order_contracts.jl new file mode 100644 index 0000000..55b16b6 --- /dev/null +++ b/test/fixtures/ordered_fold_source_order_contracts.jl @@ -0,0 +1,153 @@ +using Test +import KernelAbstractions +import LocalMath + +struct OrderedFoldSourceOrderDomain end +struct OrderedFoldSourceOrderEvaluator end + +struct OrderedFoldCanonicalKey end +@inline (::OrderedFoldCanonicalKey)(value::Int32) = value +struct OrderedFoldCanonicalIdentity end +@inline (::OrderedFoldCanonicalIdentity)(value::Int32) = value +struct AlternateOrderedFoldCanonicalKey end +@inline (::AlternateOrderedFoldCanonicalKey)(value::Int32) = value +struct AlternateOrderedFoldCanonicalIdentity end +@inline (::AlternateOrderedFoldCanonicalIdentity)(value::Int32) = value + +@inline function (::OrderedFoldSourceOrderEvaluator)(item::Int32, reads, parameters) + return (event = LocalMath.FoldValue( + item, something(reads[1][1].value)),) +end + +struct OrderedFoldSourceOrderTrace + halt_at::Int32 + invalid_at::Int32 +end + +@inline function (transition::OrderedFoldSourceOrderTrace)( + state, value, item, reads) + destination = item == transition.invalid_at ? Int32(2) : Int32(1) + next = state.result[Int32(1)] * Int32(10) + value + return LocalMath.FoldStep((result = LocalMath.BoundedWrites( + (destination,), (next,), Int32(1)),); + halt = item == transition.halt_at) +end + +function _ordered_fold_source_order_prepared( + array_type, selected_values, emitted_values; + selection::Symbol, halt_at::Int32 = Int32(0), + invalid_at::Int32 = Int32(0), order = LocalMath.source_order()) + n = length(selected_values) + source = LocalMath.Space(OrderedFoldSourceOrderDomain, n) + singleton = LocalMath.Space(OrderedFoldSourceOrderDomain, 1) + selected = LocalMath.Field(source, Bool) + emitted = LocalMath.Field(source, Bool) + initial = LocalMath.Field(singleton, Int32) + result = LocalMath.Field(singleton, Int32) + identity_relation = LocalMath.IdentityRelation(source) + control = selection === :mask ? LocalMath.Control(; mask = selected) : + selection === :subset ? LocalMath.Control(; subset = + LocalMath.MaskedRelation(identity_relation, selected)) : + error("unknown source-order selection") + state = LocalMath.initialized_state(; + result = LocalMath.FoldComponent(result; from = initial)) + publication = LocalMath.Publication((LocalMath.FoldPublication( + LocalMath.PublicationValue(:event)),), LocalMath.OrderedFold( + Int32, state, + OrderedFoldSourceOrderTrace(halt_at, invalid_at); order)) + stage = LocalMath.Stage( + source, + (emitted = LocalMath.Access(emitted, identity_relation; + required = true),), + (publication,), LocalMath.Evaluator(OrderedFoldSourceOrderEvaluator()), + control, + LocalMath.SourceOrigin( + @__FILE__, @__LINE__; label = :source_order_direct_traversal)) + selected_storage = array_type(selected_values) + emitted_storage = array_type(emitted_values) + destination = array_type(Int32[91]) + prepared = LocalMath.prepare( + LocalMath.LocalLaw(stage), + selected => selected_storage, + emitted => emitted_storage, + initial => array_type(Int32[0]), result => destination; + backend = KernelAbstractions.get_backend(destination)) + return (; prepared, destination) +end + +function _ordered_fold_source_order_execute(prepared) + return try + wait(LocalMath.execute!(prepared)) + nothing + catch error + error + end +end + +function ordered_fold_source_order_contracts(array_type) + return @testset "source-order direct traversal preserves sparse semantics" begin + selected = Bool[false, true, true, false, true, false, true] + emitted = Bool[true, true, false, true, true, true, true] + for selection in (:mask, :subset) + witness = _ordered_fold_source_order_prepared( + array_type, selected, emitted; selection) + @test _ordered_fold_source_order_execute(witness.prepared) === nothing + @test Array(witness.destination) == Int32[257] + facts = LocalMath.inspect(witness.prepared) + phases = map(phase -> phase.kind, + only(facts.stages).planning.phases) + @test :ordered_fold_bitonic ∉ phases + @test facts.planning.base_provider_launch_count == 6 + + halted = _ordered_fold_source_order_prepared( + array_type, selected, emitted; selection, + halt_at = Int32(5)) + @test _ordered_fold_source_order_execute(halted.prepared) === nothing + @test Array(halted.destination) == Int32[25] + + rejected = _ordered_fold_source_order_prepared( + array_type, selected, emitted; selection, + invalid_at = Int32(7)) + failure = _ordered_fold_source_order_execute(rejected.prepared) + @test failure isa LocalMath.LocalMathValidationError + @test failure.contract === :runtime_ordered_fold_validation + @test failure.actual.failure_class === :invalid_destination + @test failure.actual.source_item == Int32(7) + @test failure.actual.canonical_position == Int32(3) + @test failure.actual.witness == Int32(2) + @test Array(rejected.destination) == Int32[91] + end + + empty = _ordered_fold_source_order_prepared( + array_type, Bool[], Bool[]; selection = :mask) + @test _ordered_fold_source_order_execute(empty.prepared) === nothing + @test Array(empty.destination) == Int32[0] + @test LocalMath.inspect(empty.prepared).planning.base_provider_launch_count == 6 + + canonical = _ordered_fold_source_order_prepared( + array_type, trues(7), trues(7); selection = :mask, + order = LocalMath.canonical_by( + OrderedFoldCanonicalKey(), OrderedFoldCanonicalIdentity())) + @test _ordered_fold_source_order_execute(canonical.prepared) === nothing + @test Array(canonical.destination) == Int32[1234567] + canonical_facts = LocalMath.inspect(canonical.prepared) + canonical_phases = map(phase -> phase.kind, + only(canonical_facts.stages).planning.phases) + @test :ordered_fold_bitonic in canonical_phases + @test canonical_facts.planning.base_provider_launch_count == 12 + + alternate_canonical = _ordered_fold_source_order_prepared( + array_type, trues(7), trues(7); selection = :mask, + order = LocalMath.canonical_by( + AlternateOrderedFoldCanonicalKey(), + AlternateOrderedFoldCanonicalIdentity())) + @test _ordered_fold_source_order_execute( + alternate_canonical.prepared) === nothing + @test Array(alternate_canonical.destination) == Int32[1234567] + alternate_facts = LocalMath.inspect(alternate_canonical.prepared) + @test only(alternate_facts.stages).planning.phases == + only(canonical_facts.stages).planning.phases + @test alternate_facts.planning.base_provider_launch_count == + canonical_facts.planning.base_provider_launch_count + end +end diff --git a/test/metal/ordered_fold_control.jl b/test/metal/ordered_fold_control.jl index d212171..4193b9c 100644 --- a/test/metal/ordered_fold_control.jl +++ b/test/metal/ordered_fold_control.jl @@ -1,5 +1,7 @@ using Metal include(joinpath(@__DIR__, "..", "fixtures", "ordered_fold_control_contracts.jl")) +include(joinpath(@__DIR__, "..", "fixtures", "ordered_fold_source_order_contracts.jl")) Metal.functional() || error("ordered-fold control tests require functional Metal") Metal.allowscalar(false) ordered_fold_control_contracts(Metal.MtlArray) +ordered_fold_source_order_contracts(Metal.MtlArray) diff --git a/test/test_ordered_fold_stage_execution.jl b/test/test_ordered_fold_stage_execution.jl index 930314b..e6aca04 100644 --- a/test/test_ordered_fold_stage_execution.jl +++ b/test/test_ordered_fold_stage_execution.jl @@ -2,6 +2,8 @@ using Test import LocalMath const LWFSE = LocalMath +include(joinpath(@__DIR__, "fixtures", "ordered_fold_source_order_contracts.jl")) + struct OrderedFoldExecutionNode end struct OrderedFoldExecutionEvaluator end @inline (::OrderedFoldExecutionEvaluator)(item::Int32, reads, parameters) = @@ -285,3 +287,5 @@ end @test prepared.runtime.launches[1].stage.stage.parameter_slots == (LWFSE._ParameterSlot{1}(),) end + +ordered_fold_source_order_contracts(Array)