From f41a37da99c31781d4d2b23678a8e729845306a1 Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Wed, 9 Sep 2026 18:56:45 -0400 Subject: [PATCH] Admit identity-seeded reduction controls with total publication --- CONTRIBUTING.md | 10 + docs/src/api/localmath.md | 8 + src/stage_planning.jl | 10 +- test/fixtures/reduction_control_contracts.jl | 265 +++++++++++++++++++ test/metal/reduction_control.jl | 8 + test/metal/runtests.jl | 1 + test/runtests.jl | 1 + test/test_reduction_control.jl | 5 + 8 files changed, 306 insertions(+), 2 deletions(-) create mode 100644 test/fixtures/reduction_control_contracts.jl create mode 100644 test/metal/reduction_control.jl create mode 100644 test/test_reduction_control.jl diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index f64e02b..a31eb8a 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -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, diff --git a/docs/src/api/localmath.md b/docs/src/api/localmath.md index aaca8c1..c1dc8d8 100644 --- a/docs/src/api/localmath.md +++ b/docs/src/api/localmath.md @@ -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 diff --git a/src/stage_planning.jl b/src/stage_planning.jl index 28b9a45..8907efe 100644 --- a/src/stage_planning.jl +++ b/src/stage_planning.jl @@ -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, @@ -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 diff --git a/test/fixtures/reduction_control_contracts.jl b/test/fixtures/reduction_control_contracts.jl new file mode 100644 index 0000000..cd3bc79 --- /dev/null +++ b/test/fixtures/reduction_control_contracts.jl @@ -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 diff --git a/test/metal/reduction_control.jl b/test/metal/reduction_control.jl new file mode 100644 index 0000000..9273acf --- /dev/null +++ b/test/metal/reduction_control.jl @@ -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) diff --git a/test/metal/runtests.jl b/test/metal/runtests.jl index 3d070aa..ece18ab 100644 --- a/test/metal/runtests.jl +++ b/test/metal/runtests.jl @@ -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", diff --git a/test/runtests.jl b/test/runtests.jl index 53b8389..7809a9b 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -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", diff --git a/test/test_reduction_control.jl b/test/test_reduction_control.jl new file mode 100644 index 0000000..c7d5afb --- /dev/null +++ b/test/test_reduction_control.jl @@ -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)