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
10 changes: 10 additions & 0 deletions CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,16 @@ Ordered-fold participation and publication are owned by
retained destinations, and unchanged duplicate-order rejection through the
ordinary CPU and Metal inventories.

Field-derived control dependencies are owned by `src/stage_planning.jl`.
Its publication-totality predicate recognizes `Unique` with `TotalCoverage`
and `Reduce` with `IdentitySeed`, and rejects whole-stage-gated producers.
Prefix, mask, and subset controls filter contributions rather than bypassing
successful publication. An identity-seeded reduction initializes every
destination even when no source contributes; `ExistingSeed` cannot establish
a freshly produced control value. `test/fixtures/reduction_control_contracts.jl`
checks open, closed, and no-contribution gate production plus rejection of
existing-value seeds through the ordinary CPU and Metal inventories.

Pointwise traversal and its control checks are owned by
`src/execution/candidate_stage.jl`. The shared fixture
`test/fixtures/empty_pointwise_contracts.jl` checks empty-domain preparation,
Expand Down
8 changes: 8 additions & 0 deletions docs/src/api/localmath.md
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,14 @@ equation namespace:
These qualified names are stable interfaces, not permission to access other
underscored LocalMath implementation details.

A Field-derived gate or prefix requires a preceding total publication from a
Stage without a whole-stage gate. Prefix, mask, and subset controls may filter
contributions without weakening successful publication totality. `Unique`
proves publication totality with `TotalCoverage`; `Reduce` proves it with
`IdentitySeed`, which initializes destinations even when no source contributes.
`ExistingSeed` retains previous destination state and does not prove a freshly
produced control value.

`SourcePositionAccess(collection, lane=1)` is a scalar selected-lane access:
for a producer item it returns the compacted position of that exact emitted
lane. It is not an array-valued `SourcePositions` API. The producer must request
Expand Down
10 changes: 8 additions & 2 deletions src/stage_planning.jl
Original file line number Diff line number Diff line change
Expand Up @@ -280,15 +280,21 @@ end
_stage_slot_projection(bound::_BoundLaw, stage::Stage) =
_stage_slot_projection(bound.binding, bound.law.parameters, stage)

_publication_is_total(law) = false
_publication_is_total(law::Unique) = law.coverage isa TotalCoverage
_publication_is_total(law::Reduce) = law.seed isa IdentitySeed

function _stage_publishes_field(
stage::Stage, field::Field; total::Bool = false,
)
total && !(stage.control.gate isa _NoGate) && return false
for publication in stage.publications
if publication.law isa OrderedFold
for component in values(publication.law.state.components)
semantic_identity(component.target) == semantic_identity(field) ||
continue
component.target == field || throw(LocalMathValidationError(
component.target == field || throw(
LocalMathValidationError(
"a fold target Field identity has conflicting schema";
stage = :plan, contract = :field_dependency_schema,
expected = field, actual = component.target,
Expand All @@ -307,7 +313,7 @@ function _stage_publishes_field(
stage = :plan, contract = :field_dependency_schema,
expected = field, actual = component.field,
))
!total || publication.law.coverage isa TotalCoverage || return false
!total || _publication_is_total(publication.law) || return false
return true
end
end
Expand Down
265 changes: 265 additions & 0 deletions test/fixtures/reduction_control_contracts.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,265 @@
using Test
import LocalMath
import KernelAbstractions

struct ReductionControlProducer{SkipClosed} end
@inline function (::ReductionControlProducer{SkipClosed})(item::Int32, reads, parameters) where {SkipClosed}
enabled = something(reads[1][1].value)
return (value = LocalMath.RoutedContribution(Int32(1), enabled, !SkipClosed || enabled),)
end
struct ReductionControlSum end
@inline (::ReductionControlSum)(item::Int32, reads, parameters) =
(value = LocalMath.RoutedContribution(Int32(1), something(reads[1][1].value)),)

function reduction_control_filtered_producers(array_type)
return @testset "filtered reductions still produce total control" begin
for selection in (:prefix, :mask, :subset)
sites, singleton = LocalMath.Space(3), LocalMath.Space(1)
enabled = LocalMath.Field(sites, Bool)
selected = LocalMath.Field(sites, Bool)
input = LocalMath.Field(sites, Float32)
gate, output = LocalMath.Field(singleton, Bool), LocalMath.Field(singleton, Float32)
count = LocalMath.Parameter(:source_count, Int32)
route = LocalMath.RuntimeRelation(sites => singleton; degree_bound = 1, key_type = Int32)
control = selection === :prefix ? LocalMath.Control(; prefix = count) :
selection === :mask ? LocalMath.Control(; mask = selected) :
LocalMath.Control(; subset = LocalMath.MaskedRelation(LocalMath.IdentityRelation(sites), selected))
producer = LocalMath.Stage(
sites,
(enabled = LocalMath.Access(enabled, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(gate, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Bool, |; maximum = 1, seed = LocalMath.IdentitySeed(false), order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlProducer{false}(), selection === :prefix ? (count,) : ()), control,
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :filtered_control_producer)
)
consumer = LocalMath.Stage(
sites,
(input = LocalMath.Access(input, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(output, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Float32, +; maximum = 1, seed = LocalMath.IdentitySeed(0.0f0), order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlSum()), LocalMath.Control(; gate),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :filtered_control_consumer)
)
selected_values = array_type(Bool[true, false, true])
gate_values, destination = array_type(Bool[true]), array_type(Float32[-99])
bindings = (
enabled => array_type(trues(3)), input => array_type(Float32[1, 2, 3]),
gate => gate_values, output => destination,
(selection === :prefix ? () : (selected => selected_values,))...,
)
prepared = LocalMath.prepare(
LocalMath.sequence(LocalMath.LocalLaw(producer), LocalMath.LocalLaw(consumer)),
bindings...; backend = KernelAbstractions.get_backend(destination)
)
for active in (true, false, true)
copyto!(selected_values, array_type(fill(active, 3)))
parameters = selection === :prefix ? (; source_count = active ? Int32(2) : Int32(0)) : NamedTuple()
wait(LocalMath.execute!(prepared; parameters))
@test Array(gate_values) == Bool[active]
@test Array(destination) == Float32[6]
end
end
end
end

struct ReductionControlCount end
@inline (::ReductionControlCount)(item::Int32, reads, parameters) =
(value = LocalMath.RoutedContribution(Int32(1), Int32(1), something(reads[1][1].value)),)

function reduction_control_atomic_prefix(array_type)
return @testset "atomic identity-seeded reduction produces a prefix" begin
sites, singleton = LocalMath.Space(3), LocalMath.Space(1)
enabled = LocalMath.Field(sites, Bool)
input = LocalMath.Field(sites, Float32)
prefix, output = LocalMath.Field(singleton, Int32), LocalMath.Field(singleton, Float32)
route = LocalMath.RuntimeRelation(sites => singleton; degree_bound = 1, key_type = Int32)
producer = LocalMath.Stage(
sites,
(enabled = LocalMath.Access(enabled, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(prefix, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Int32, +; maximum = 1, seed = LocalMath.IdentitySeed(Int32(0)), order = LocalMath.RelaxedAtomic())
),
),
LocalMath.Evaluator(ReductionControlCount()), LocalMath.Control(),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :atomic_prefix_producer)
)
consumer = LocalMath.Stage(
sites,
(input = LocalMath.Access(input, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(output, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Float32, +; maximum = 1, seed = LocalMath.IdentitySeed(0.0f0), order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlSum()), LocalMath.Control(; prefix),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :atomic_prefix_consumer)
)
enabled_values = array_type(trues(3))
prefix_values, destination = array_type(Int32[3]), array_type(Float32[-99])
prepared = LocalMath.prepare(
LocalMath.sequence(LocalMath.LocalLaw(producer), LocalMath.LocalLaw(consumer)),
enabled => enabled_values, input => array_type(Float32[1, 2, 3]), prefix => prefix_values, output => destination;
backend = KernelAbstractions.get_backend(destination)
)
for (flags, expected_count, expected_sum) in ((trues(3), 3, 6), (falses(3), 0, 0), (Bool[true, false, true], 2, 3))
copyto!(enabled_values, array_type(flags))
wait(LocalMath.execute!(prepared))
@test Array(prefix_values) == Int32[expected_count]
@test Array(destination) == Float32[expected_sum]
end
end
end

struct ReductionControlUnique end
@inline (::ReductionControlUnique)(item::Int32, reads, parameters) =
(value = LocalMath.UniqueValue(something(reads[1][1].value)),)

function reduction_control_producer_totality(array_type)
return @testset "whole-stage participation and control totality" begin
for reduction in (false, true), gated in (false, true)
singleton = LocalMath.Space(1)
external = LocalMath.Field(singleton, Bool)
gate = LocalMath.Field(singleton, Bool)
input = LocalMath.Field(singleton, Float32)
output = LocalMath.Field(singleton, Float32)
producer_enabled = LocalMath.Parameter(:producer_enabled, Bool)
route = LocalMath.RuntimeRelation(singleton => singleton; degree_bound = 1, key_type = Int32)
law = reduction ? LocalMath.Reduce(
Bool, |; maximum = 1,
seed = LocalMath.IdentitySeed(false), order = LocalMath.CanonicalLeftFold()
) : LocalMath.Unique(Bool)
relation = reduction ? route : LocalMath.IdentityRelation(singleton)
evaluator = reduction ? ReductionControlProducer{false}() : ReductionControlUnique()
producer = LocalMath.Stage(
singleton,
(enabled = LocalMath.Access(external, LocalMath.IdentityRelation(singleton); required = true),),
(LocalMath.Publication((LocalMath.FieldPublication(gate, relation, LocalMath.PublicationValue(:value)),), law),),
LocalMath.Evaluator(evaluator, gated ? (producer_enabled,) : ()),
gated ? LocalMath.Control(; gate = producer_enabled) : LocalMath.Control(),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :conditional_control_producer)
)
consumer = LocalMath.Stage(
singleton,
(input = LocalMath.Access(input, LocalMath.IdentityRelation(singleton); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(output, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Float32, +; maximum = 1, seed = LocalMath.IdentitySeed(0.0f0), order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlSum()), LocalMath.Control(; gate),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :conditional_control_consumer)
)
gate_values, destination = array_type(Bool[true]), array_type(Float32[-99])
enabled_values = array_type(Bool[false])
bindings = (external => enabled_values, input => array_type(Float32[7]), gate => gate_values, output => destination)
work = LocalMath.sequence(LocalMath.LocalLaw(producer), LocalMath.LocalLaw(consumer))
backend = KernelAbstractions.get_backend(destination)
if !gated
prepared = LocalMath.prepare(work, bindings...; backend)
wait(LocalMath.execute!(prepared))
@test Array(gate_values) == Bool[false]
@test Array(destination) == Float32[-99]
copyto!(enabled_values, array_type(Bool[true]))
wait(LocalMath.execute!(prepared))
@test Array(gate_values) == Bool[true]
@test Array(destination) == Float32[7]
continue
end
failure = try
LocalMath.prepare(work, bindings...; backend)
nothing
catch error
error
end
@test failure isa LocalMath.LocalMathValidationError
@test failure.contract === :control_field_totality
@test Array(gate_values) == Bool[true]
@test Array(destination) == Float32[-99]
end
end
end

function reduction_control_contracts(array_type)
return @testset "identity-seeded reduction controls" begin
for skip_closed in (false, true), seed in (LocalMath.IdentitySeed(false), LocalMath.ExistingSeed())
@testset "skip closed=$skip_closed seed=$(nameof(typeof(seed)))" begin
sites, singleton = LocalMath.Space(3), LocalMath.Space(1)
enabled = LocalMath.Field(sites, Bool)
input = LocalMath.Field(sites, Float32)
gate = LocalMath.Field(singleton, Bool)
output = LocalMath.Field(singleton, Float32)
route = LocalMath.RuntimeRelation(sites => singleton; degree_bound = 1, key_type = Int32)
producer = LocalMath.Stage(
sites,
(enabled = LocalMath.Access(enabled, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(gate, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Bool, |; maximum = 1, seed, order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlProducer{skip_closed}()), LocalMath.Control(),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :reduced_control_producer)
)
consumer = LocalMath.Stage(
sites,
(input = LocalMath.Access(input, LocalMath.IdentityRelation(sites); required = true),),
(
LocalMath.Publication(
(LocalMath.FieldPublication(output, route, LocalMath.PublicationValue(:value)),),
LocalMath.Reduce(Float32, +; maximum = 1, seed = LocalMath.IdentitySeed(0.0f0), order = LocalMath.CanonicalLeftFold())
),
),
LocalMath.Evaluator(ReductionControlSum()), LocalMath.Control(; gate),
LocalMath.SourceOrigin(@__FILE__, @__LINE__; label = :reduced_control_consumer)
)
enabled_values = array_type(Bool[true, false, true])
input_values = array_type(Float32[1, 2, 3])
gate_values = array_type(Bool[true])
destination = array_type(Float32[-99])
work = LocalMath.sequence(LocalMath.LocalLaw(producer), LocalMath.LocalLaw(consumer))
bindings = (enabled => enabled_values, input => input_values, gate => gate_values, output => destination)
backend = KernelAbstractions.get_backend(destination)
if seed isa LocalMath.ExistingSeed
failure = try
LocalMath.prepare(work, bindings...; backend)
nothing
catch error
error
end
@test failure isa LocalMath.LocalMathValidationError
@test failure.contract === :control_field_totality
@test Array(gate_values) == Bool[true]
@test Array(destination) == Float32[-99]
continue
end
prepared = LocalMath.prepare(work, bindings...; backend)
wait(LocalMath.execute!(prepared))
@test Array(gate_values) == Bool[true]
@test Array(destination) == Float32[6]
copyto!(enabled_values, array_type(falses(3)))
copyto!(input_values, array_type(Float32[4, 5, 6]))
wait(LocalMath.execute!(prepared))
@test Array(gate_values) == Bool[false]
@test Array(destination) == Float32[6]
copyto!(enabled_values, array_type(Bool[false, true, false]))
wait(LocalMath.execute!(prepared))
@test Array(gate_values) == Bool[true]
@test Array(destination) == Float32[15]
@test Array(input_values) == Float32[4, 5, 6]
end
end
end
end
8 changes: 8 additions & 0 deletions test/metal/reduction_control.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
using Metal
include(joinpath(@__DIR__, "..", "fixtures", "reduction_control_contracts.jl"))
Metal.functional() || error("reduction control tests require functional Metal")
Metal.allowscalar(false)
reduction_control_contracts(Metal.MtlArray)
reduction_control_producer_totality(Metal.MtlArray)
reduction_control_filtered_producers(Metal.MtlArray)
reduction_control_atomic_prefix(Metal.MtlArray)
1 change: 1 addition & 0 deletions test/metal/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ const LOCALMATH_METAL_WITNESSES = (
"destination_grouping.jl",
"product_values.jl",
"ordered_fold_control.jl",
"reduction_control.jl",
"empty_pointwise_domains.jl",
"collect_canonical_order.jl",
"trigonometric_stages.jl",
Expand Down
1 change: 1 addition & 0 deletions test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ const LOCALMATH_INCLUDED_TESTS = (
"test_stage_program_lifecycle.jl",
"test_execution_receipts.jl",
"test_reduce_stage.jl",
"test_reduction_control.jl",
"test_resolve_stage.jl",
"test_runtime_routed_stage.jl",
"test_candidate_grouping.jl",
Expand Down
5 changes: 5 additions & 0 deletions test/test_reduction_control.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
include(joinpath(@__DIR__, "fixtures", "reduction_control_contracts.jl"))
reduction_control_contracts(Array)
reduction_control_producer_totality(Array)
reduction_control_filtered_producers(Array)
reduction_control_atomic_prefix(Array)