From 3dfc1e87729fb0d257776b3df373e470d9e228e9 Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Thu, 24 Sep 2026 00:21:40 -0400 Subject: [PATCH] Scope keyed reduction scratch leaves by stage --- src/execution/keyed_reduce_stage.jl | 24 +++++++----- test/fixtures/keyed_reduce_contracts.jl | 51 +++++++++++++++++++++++++ test/metal/keyed_reduce.jl | 1 + test/test_keyed_reduce_stage.jl | 1 + 4 files changed, 68 insertions(+), 9 deletions(-) diff --git a/src/execution/keyed_reduce_stage.jl b/src/execution/keyed_reduce_stage.jl index ac12cf6..def335f 100644 --- a/src/execution/keyed_reduce_stage.jl +++ b/src/execution/keyed_reduce_stage.jl @@ -65,17 +65,19 @@ function _keyed_reduce_physical(stage, return _KeyedReducePhysical(bounds, emission, key_order, fold, publication) end -function _keyed_reduce_component_workspace_spec(root::Tuple, label::Symbol, - ::Type{T}, count::Int, path::Tuple = ()) where {T} +function _keyed_reduce_component_workspace_spec(root::Tuple, + name_prefix::Symbol, label::Symbol, ::Type{T}, count::Int, + path::Tuple = ()) where {T} if fieldcount(T) == 0 suffix = isempty(path) ? :scalar : Symbol(join(string.(path), :_)) - return (_workspace_leaf(Symbol(:keyed_reduce_, label, :_, suffix), + return (_workspace_leaf( + Symbol(name_prefix, :_keyed_reduce_, label, :_, suffix), (root..., :keyed_reduce, label, path...), T, (count,); role = Symbol(:keyed_reduce_, label, :_component)),) end return reduce(1:fieldcount(T); init = ()) do leaves, field_index - (leaves..., _keyed_reduce_component_workspace_spec(root, label, - fieldtype(T, field_index), count, + (leaves..., _keyed_reduce_component_workspace_spec( + root, name_prefix, label, fieldtype(T, field_index), count, (path..., fieldname(T, field_index)))...) end end @@ -89,10 +91,14 @@ function _keyed_reduce_stage_workspace_spec(stage; path::Tuple = (), candidates = Int(plan.bounds.candidate_count) prefix_count, sums_count = _compacted_scan_storage_lengths(candidates) leaves = ( - _keyed_reduce_component_workspace_spec(path, :keys, K, candidates)..., - _keyed_reduce_component_workspace_spec(path, :values, V, candidates)..., - _keyed_reduce_component_workspace_spec(path, :reduced_keys, K, candidates)..., - _keyed_reduce_component_workspace_spec(path, :reduced_values, V, candidates)..., + _keyed_reduce_component_workspace_spec( + path, name_prefix, :keys, K, candidates)..., + _keyed_reduce_component_workspace_spec( + path, name_prefix, :values, V, candidates)..., + _keyed_reduce_component_workspace_spec( + path, name_prefix, :reduced_keys, K, candidates)..., + _keyed_reduce_component_workspace_spec( + path, name_prefix, :reduced_values, V, candidates)..., _workspace_leaf(Symbol(name_prefix, :_valid), (path..., :keyed_reduce, :valid), UInt8, (candidates,); role = :keyed_reduce_candidate_participation), diff --git a/test/fixtures/keyed_reduce_contracts.jl b/test/fixtures/keyed_reduce_contracts.jl index be3573f..f27e546 100644 --- a/test/fixtures/keyed_reduce_contracts.jl +++ b/test/fixtures/keyed_reduce_contracts.jl @@ -43,6 +43,57 @@ struct RebuildOrderedKeyedReduceEvaluator end @inline (::RebuildOrderedKeyedReduceEvaluator)(item::Int32, reads, parameters) = (; delta = LocalMath.KeyedContribution(Int32(3), item)) +struct FirstMixedKeyEvaluator end +@inline (::FirstMixedKeyEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution((item, UInt32(1)), Int32(1))) + +struct SecondMixedKeyEvaluator end +@inline (::SecondMixedKeyEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution((UInt32(1), -item), Int32(1))) + +function mixed_keyed_stage_sequence_contract(backend) + return @testset "successive keyed reductions retain distinct key layouts" begin + source = LocalMath.Space(KeyedReduceContractNode, 2) + first_key = Tuple{Int32,UInt32} + second_key = Tuple{UInt32,Int32} + first = LocalMath.Collection(LocalMath.KeyedValue{first_key,Int32}, 2) + second = LocalMath.Collection(LocalMath.KeyedValue{second_key,Int32}, 2) + function stage(destination, key_type, evaluator, label) + LocalMath.Stage(source, NamedTuple(), ( + LocalMath.Publication(destination, + LocalMath.KeyedReduce(key_type, Int32, +; + maximum = 1, + seed = LocalMath.RebuildFromIdentity(Int32(0)), + retention = LocalMath.DropIdentityKeys()); + value = :delta),), + LocalMath.Evaluator(evaluator), LocalMath.Control(), + LocalMath.SourceOrigin(:keyed_reduce_contract, label)) + end + law = LocalMath.sequence( + LocalMath.LocalLaw(stage( + first, first_key, FirstMixedKeyEvaluator(), 3)), + LocalMath.LocalLaw(stage( + second, second_key, SecondMixedKeyEvaluator(), 4)), + ) + prepared = LocalMath.prepare(law, + first => LocalMath.Allocate(), + second => LocalMath.Allocate(); backend) + wait(LocalMath.execute!(prepared)) + first_records = LocalMath.storage(prepared, first) + second_records = LocalMath.storage(prepared, second) + @test only(LocalMath.Adapt.adapt(Array, first_records.count)) == 2 + @test only(LocalMath.Adapt.adapt(Array, second_records.count)) == 2 + @test Set(LocalMath.Adapt.adapt(Array, first_records.records)) == Set([ + LocalMath.KeyedValue((Int32(1), UInt32(1)), Int32(1)), + LocalMath.KeyedValue((Int32(2), UInt32(1)), Int32(1)), + ]) + @test Set(LocalMath.Adapt.adapt(Array, second_records.records)) == Set([ + LocalMath.KeyedValue((UInt32(1), Int32(-1)), Int32(1)), + LocalMath.KeyedValue((UInt32(1), Int32(-2)), Int32(1)), + ]) + end +end + function _shared_keyed_program(backend, source_count, key_type, capacity, evaluator; maximum = 1, operation = +, seed = LocalMath.NewKeyIdentity(Int32(0)), diff --git a/test/metal/keyed_reduce.jl b/test/metal/keyed_reduce.jl index 63bbb00..f427adb 100644 --- a/test/metal/keyed_reduce.jl +++ b/test/metal/keyed_reduce.jl @@ -6,3 +6,4 @@ include("../fixtures/keyed_reduce_contracts.jl") Metal.functional() || error("keyed reduction checks require real Metal") Metal.allowscalar(false) keyed_reduce_contract(Metal.MetalBackend()) +mixed_keyed_stage_sequence_contract(Metal.MetalBackend()) diff --git a/test/test_keyed_reduce_stage.jl b/test/test_keyed_reduce_stage.jl index 344209b..10b63ea 100644 --- a/test/test_keyed_reduce_stage.jl +++ b/test/test_keyed_reduce_stage.jl @@ -9,6 +9,7 @@ include("fixtures/keyed_reduce_contracts.jl") Int32, Int32, +; seed = LMKR.RebuildFromIdentity(0.0f0)) keyed_reduce_contract(KernelAbstractions.CPU()) +mixed_keyed_stage_sequence_contract(KernelAbstractions.CPU()) struct KeyedReduceNode end struct KeyedReduceEvaluator end