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
94 changes: 68 additions & 26 deletions src/execution/ordered_fold_stage.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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!(
Expand Down Expand Up @@ -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 = (;
Expand Down
25 changes: 14 additions & 11 deletions src/execution/program_inspection.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)))
Expand Down
153 changes: 153 additions & 0 deletions test/fixtures/ordered_fold_source_order_contracts.jl
Original file line number Diff line number Diff line change
@@ -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
2 changes: 2 additions & 0 deletions test/metal/ordered_fold_control.jl
Original file line number Diff line number Diff line change
@@ -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)
4 changes: 4 additions & 0 deletions test/test_ordered_fold_stage_execution.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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) =
Expand Down Expand Up @@ -285,3 +287,5 @@ end
@test prepared.runtime.launches[1].stage.stage.parameter_slots ==
(LWFSE._ParameterSlot{1}(),)
end

ordered_fold_source_order_contracts(Array)