From 9d5b383fd324624aa6b91cf6e9ddece365d7b3a7 Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Mon, 14 Sep 2026 22:54:24 -0400 Subject: [PATCH] Add atomic keyed rebuild publication --- benchmark/keyed_reduce_compiler.jl | 108 ++++++++++++++++++++++-- docs/src/api/localmath.md | 35 +++++--- src/LocalMath.jl | 2 +- src/execution/keyed_reduce_stage.jl | 26 ++++-- src/execution/stage_preparation.jl | 10 ++- src/stage_model.jl | 38 ++++++--- test/fixtures/keyed_reduce_contracts.jl | 88 ++++++++++++++++++- test/test_keyed_reduce_stage.jl | 4 + test/test_public_api.jl | 2 +- 9 files changed, 268 insertions(+), 45 deletions(-) diff --git a/benchmark/keyed_reduce_compiler.jl b/benchmark/keyed_reduce_compiler.jl index fc1612d..9da63bd 100644 --- a/benchmark/keyed_reduce_compiler.jl +++ b/benchmark/keyed_reduce_compiler.jl @@ -19,7 +19,8 @@ struct KeyedCompilerSubtract end @inline (::KeyedCompilerSubtract)(left::Int32, right::Int32) = left - right function keyed_compiler_preparation(capacity::Int; - operation = +, retention = LocalMath.DropIdentityKeys()) + operation = +, retention = LocalMath.DropIdentityKeys(), + seed = LocalMath.NewKeyIdentity(Int32(0))) source = LocalMath.Space(KeyedCompilerNode, 3) key_type = Tuple{UInt32,UInt32} collection = LocalMath.Collection( @@ -27,7 +28,7 @@ function keyed_compiler_preparation(capacity::Int; stage = LocalMath.Stage(source, NamedTuple(), ( LocalMath.Publication(collection, LocalMath.KeyedReduce(key_type, Int32, operation; - seed = LocalMath.NewKeyIdentity(Int32(0)), retention); + seed, retention); value = :delta),), LocalMath.Evaluator(KeyedCompilerEvaluator()), LocalMath.Control(), LocalMath.SourceOrigin(:keyed_reduce_compiler, 1)) @@ -49,15 +50,20 @@ function collect_compiler_preparation(capacity::Int) backend = KernelAbstractions.CPU()) end -function typed_metrics(callable, signature) - info, return_type = only(Base.code_typed_by_type( - Tuple{typeof(callable),signature.parameters...}; optimize = true)) +function code_info_metrics(info, return_type, method_instance_count, + method_specialization_count) calls = count(statement -> statement isa Expr && statement.head in (:call, :invoke), info.code) any_indices = filter(index -> info.ssavaluetypes[index] === Any, eachindex(info.code)) control_flow = statement -> statement isa Union{ Core.GotoNode,Core.GotoIfNot,Core.ReturnNode} + statement_kinds = Dict{String,Int}() + for statement in info.code + kind = statement isa Expr ? string(statement.head) : + string(nameof(typeof(statement))) + statement_kinds[kind] = get(statement_kinds, kind, 0) + 1 + end return Dict( "statement_count" => length(info.code), "call_count" => calls, @@ -66,10 +72,47 @@ function typed_metrics(callable, signature) index -> control_flow(info.code[index]), any_indices), "any_value_count" => count( index -> !control_flow(info.code[index]), any_indices), + "ssa_count" => length(info.ssavaluetypes), + "slot_count" => length(info.slotnames), + "method_instance_count" => method_instance_count, + "method_specialization_count" => method_specialization_count, + "statement_kinds" => statement_kinds, "return_type" => string(return_type), ) end +function typed_metrics(callable, signature) + full_signature = Tuple{typeof(callable),signature.parameters...} + method = which(callable, signature) + info, return_type = only(Base.code_typed_by_type( + full_signature; optimize = true)) + return code_info_metrics(info, return_type, + length(Base.method_instances( + callable, signature, Base.get_world_counter())), + count(instance -> instance isa Core.MethodInstance, + Base.specializations(method))) +end + + +function kernel_typed_metrics(kernel, signature; + ndrange::Int, workgroupsize::Int) + launch_ndrange, launch_workgroupsize, iterspace, dynamic = + KernelAbstractions.launch_config( + kernel, ndrange, workgroupsize) + block = @inbounds KernelAbstractions.blocks(iterspace)[1] + context = KernelAbstractions.mkcontext(kernel, block, launch_ndrange, + iterspace, dynamic) + transformed_signature = Tuple{typeof(context),signature.parameters...} + method = which(kernel.f, transformed_signature) + info, return_type = only(KernelAbstractions.ka_code_typed( + kernel, signature; ndrange, workgroupsize, optimize = true)) + return code_info_metrics(info, return_type, + length(Base.method_instances(kernel.f, transformed_signature, + Base.get_world_counter())), + count(instance -> instance isa Core.MethodInstance, + Base.specializations(method))) +end + function warm_public_execution_allocations(prepared) wait(LocalMath.execute!(prepared)) return minimum(@allocated(wait(LocalMath.execute!(prepared))) for _ in 1:5) @@ -88,6 +131,17 @@ function keyed_compiler_metrics(prepared) states = getfield(execution, :states) emission = LocalMath.KeyedContribution( (UInt32(1), UInt32(1)), Int32(1)) + extent = max(Int(plan.bounds.candidate_count), 1) + reset_kernel = LocalMath._keyed_reduce_reset_kernel!( + KernelAbstractions.CPU(), min(extent, LocalMath._COMPACTED_BLOCK), + extent) + reset_signature = Tuple{typeof(plan.bounds),typeof(states.reset), + typeof(execution.storage),typeof(getfield(stage, :validation)),Int32} + segment_kernel = LocalMath._keyed_reduce_segments_kernel!( + KernelAbstractions.CPU(), min(extent, LocalMath._COMPACTED_BLOCK), + extent) + segment_signature = Tuple{typeof(plan.key_order),typeof(states.segment), + typeof(states.ordering.order_a),Int32,Int32} semantic_signature = Tuple{typeof(plan.emission),typeof(states.emission), typeof(emission),Int32} semantic = typed_metrics(LocalMath._keyed_reduce_materialize!, @@ -96,11 +150,17 @@ function keyed_compiler_metrics(prepared) return Dict( "host_orchestration" => host, "emission_boundary" => semantic, + "reset_kernel_boundary" => kernel_typed_metrics( + reset_kernel, reset_signature; ndrange = extent, + workgroupsize = min(extent, LocalMath._COMPACTED_BLOCK)), "sort_boundary" => typed_metrics(LocalMath._compacted_ordinal_less, Tuple{typeof(plan.key_order),typeof(states.ordering),Int32,Int32}), "fold_boundary" => typed_metrics(LocalMath._keyed_reduce_fold_segment!, Tuple{typeof(plan.fold),typeof(states.fold), typeof(states.ordering.order_a),Int32,Int32}), + "segment_kernel_boundary" => kernel_typed_metrics( + segment_kernel, segment_signature; ndrange = extent, + workgroupsize = min(extent, LocalMath._COMPACTED_BLOCK)), "publish_boundary" => typed_metrics(LocalMath._keyed_reduce_publish_record!, Tuple{typeof(plan.publication),typeof(states.publication), typeof(execution.storage),Int32,Int32}), @@ -140,6 +200,8 @@ variants = ( retention = LocalMath.RetainAllKeys()), ) metrics = keyed_compiler_metrics(first(preparations)) +metrics["measured_revision"] = get( + ENV, "LOCALMATH_COMPILER_REVISION", "working_tree") control = collect_compiler_preparation(4) metrics["collect_host_orchestration"] = collect_compiler_metrics(control) control_allocated = warm_public_execution_allocations(control) @@ -170,6 +232,40 @@ metrics["operation_retention_specializations"] = Dict( typeof(part.states.publication), typeof(part.storage)) for part in variant_parts)), ) -metrics["kaimon"] = "executed through a Kaimon persistent LocalMath project session; metrics are produced by these reproducible Base.code_typed_by_type probes" +if isdefined(LocalMath, :RebuildFromIdentity) + seed_variants = ( + keyed_compiler_preparation(4; + seed = LocalMath.NewKeyIdentity(Int32(0))), + keyed_compiler_preparation(4; + seed = LocalMath.RebuildFromIdentity(Int32(0))), + ) + seed_variant_parts = map(seed_variants) do prepared + stage = getfield(only(getfield( + getfield(prepared, :runtime), :launches)), :stage) + execution = getfield(stage, :execution) + (; stage, plan = execution.plan, states = execution.states, + storage = execution.storage) + end + metrics["seed_policy_specializations"] = Dict( + "prepared_stage" => length(unique(typeof(part.stage) + for part in seed_variant_parts)), + "bounds" => length(unique(typeof(part.plan.bounds) + for part in seed_variant_parts)), + "emission" => length(unique(typeof(part.plan.emission) + for part in seed_variant_parts)), + "sort" => length(unique((typeof(part.plan.key_order), + typeof(part.states.ordering)) for part in seed_variant_parts)), + "fold" => length(unique((typeof(part.plan.fold), + typeof(part.states.fold)) for part in seed_variant_parts)), + "publish" => length(unique((typeof(part.plan.publication), + typeof(part.states.publication), typeof(part.storage)) + for part in seed_variant_parts)), + ) + metrics["seed_policy_warm_public_execution_allocated_bytes"] = Dict( + "incremental" => warm_public_execution_allocations( + first(seed_variants)), + "rebuild" => warm_public_execution_allocations(last(seed_variants)), + ) +end TOML.print(stdout, Dict("keyed_reduce" => metrics); sorted = true) println() diff --git a/docs/src/api/localmath.md b/docs/src/api/localmath.md index 8727753..5fb6fd8 100644 --- a/docs/src/api/localmath.md +++ b/docs/src/api/localmath.md @@ -52,12 +52,15 @@ Duplicate canonical identities fail validation before changing the previously published count or records. The ordinary CPU and Metal collection-order tests exercise these behaviors across partial workgroups with bounds checks enabled. -### Incremental sparse keyed state +### Sparse keyed update and rebuild `KeyedReduce` updates one bounded keyed `Collection` without introducing a -second scheduler or storage authority. Prior records are intrinsic input state; -the required `seed` keyword supplies the initial value only for a key absent at -stage entry. +second scheduler or storage authority. With `NewKeyIdentity`, prior records are +intrinsic input state and the identity initializes only keys absent at stage +entry. With `RebuildFromIdentity`, stage-entry keys and values are ignored and +every emitted key begins at the identity, so successful execution replaces the +complete logical collection. Both policies use the same failure-atomic +publication path. ```julia import KernelAbstractions @@ -90,13 +93,21 @@ wait(LocalMath.execute!(prepared)) records = LocalMath.storage(prepared, counts) ``` -Every exact-key segment intrinsically folds an existing value first, then -participating tuple lanes in canonical `(source, lane)` order. Invalid prior counts, duplicate prior -keys, and final capacity overflow reject the whole publication, leaving its -records and logical count unchanged. `prepare` owns the bounded device -workspace; execution performs no device allocation. Public `execute!` and -`wait` still allocate shared host receipt/launch bookkeeping, which is tracked -separately rather than claimed as zero-allocation execution. +With `NewKeyIdentity`, every exact-key segment folds its stage-entry value first +when one exists, then participating tuple lanes in canonical `(source, lane)` +order. Invalid prior counts, duplicate prior keys, and final capacity overflow +reject the whole publication, leaving its records and logical count unchanged. +`prepare` owns the bounded device workspace; execution performs no device +allocation. Public `execute!` and `wait` still allocate shared host +receipt/launch bookkeeping, which is tracked separately rather than claimed as +zero-allocation execution. + +For `RebuildFromIdentity`, only participating contributions form the candidate: +every exact-key segment begins from the declared identity and then folds its +participating tuple lanes in the same canonical order. Stage-entry count, keys, +and values are not read or validated. Capacity, control, evaluator, and +predecessor-stage failure still suppress publication, leaving the complete +stage-entry collection observable until a successful replacement is published. ## Public surface @@ -118,7 +129,7 @@ equation namespace: | Lifecycle | `Plan`, `PreparedPlan`, `ExecutionReceipt`, `LocalMathValidationError`, `bind`, `plan`, `Allocate`, `Temporary`, `MutableRelationStorage`, `storage`, `inspect`, `compilation_report`, `execution_contract`, `lowering_identity` | | Explicit laws | `Stage`, `Publication`, `Access`, `Control`, `SourceOrigin`, `Parameter`, `ParameterSchema`, `Evaluator`, `FieldPublication`, `CollectionPublication`, `FoldPublication`, `PublicationValue`, `sequence` | | Collections | `CollectionAccess`, `CollectionCount`, `BoundedGroup`, `SourcePositionAccess`, `CompactedStorage`, `BoundedGroupView`, `KeyedValue`, `one_group`, `group_by`, `source_order`, `canonical_by`, `persistent_source_position` | -| Publication laws | `Unique`, `Reduce`, `Resolve`, `Collect`, `KeyedReduce`, `OrderedFold`, `TotalCoverage`, `PartialCoverage`, `UnreachableEmpty`, `PreserveEmpty`, `FillEmpty`, `IdentitySeed`, `ExistingSeed`, `NewKeyIdentity`, `RetainAllKeys`, `DropIdentityKeys`, `CanonicalLeftFold`, `RelaxedAtomic`, `ArgMin`, `ArgMax`, `CanonicalSourceLaneTie`, `TieMin`, `TieMax`, `RejectOverflow`, `EmptyCollection` | +| Publication laws | `Unique`, `Reduce`, `Resolve`, `Collect`, `KeyedReduce`, `OrderedFold`, `TotalCoverage`, `PartialCoverage`, `UnreachableEmpty`, `PreserveEmpty`, `FillEmpty`, `IdentitySeed`, `ExistingSeed`, `NewKeyIdentity`, `RebuildFromIdentity`, `RetainAllKeys`, `DropIdentityKeys`, `CanonicalLeftFold`, `RelaxedAtomic`, `ArgMin`, `ArgMax`, `CanonicalSourceLaneTie`, `TieMin`, `TieMax`, `RejectOverflow`, `EmptyCollection` | | Ordered state | `FoldComponent`, `InitializedState`, `initialized_state`, `BoundedWrites`, `FoldStep` | | Bounded scalar operations | `fold`, `BoundedFold`, `Where`, `RejectInvalid`, `SkipInvalid`, `FillInvalid`, `RejectEmpty`, `RelaxedAssociative`, `BoundedFoldOutcome`, `evaluate_bounded` | | Evaluator outputs | `UniqueValue`, `ConditionalUniqueValue`, `RoutedUniqueValue`, `ConditionalRoutedUniqueValue`, `Contribution`, `RoutedContribution`, `ResolutionValue`, `RoutedResolutionValue`, `CollectedValue`, `GroupedCollectedValue`, `KeyedContribution`, `FoldValue` | diff --git a/src/LocalMath.jl b/src/LocalMath.jl index 91ddf57..a73745c 100644 --- a/src/LocalMath.jl +++ b/src/LocalMath.jl @@ -47,7 +47,7 @@ public UniqueValue, ConditionalUniqueValue, RoutedUniqueValue public ConditionalRoutedUniqueValue, Contribution, RoutedContribution public ResolutionValue, RoutedResolutionValue, CollectedValue public GroupedCollectedValue, KeyedValue, KeyedContribution, FoldValue -public NewKeyIdentity, RetainAllKeys, DropIdentityKeys +public NewKeyIdentity, RebuildFromIdentity, RetainAllKeys, DropIdentityKeys import Adapt import Atomix import KernelAbstractions diff --git a/src/execution/keyed_reduce_stage.jl b/src/execution/keyed_reduce_stage.jl index ee32b0b..2d882a4 100644 --- a/src/execution/keyed_reduce_stage.jl +++ b/src/execution/keyed_reduce_stage.jl @@ -15,6 +15,7 @@ const _KEYED_REDUCE_STATUS_INVALID_CONTROL = Int32(5) struct _KeyedReduceBounds capacity::Int32 + includes_stage_entry::Bool candidate_count::Int32 merge_passes::Int32 end @@ -46,17 +47,20 @@ function _keyed_reduce_physical(stage, storage = only(publication.components).storage emitted = _candidate_record_capacity(Int(stage.source_count), W, :keyed_reduce_emission_capacity; int32_index = true) - total = _candidate_record_capacity(1, Int(length(storage.records)) + emitted, + law = publication.law + capacity = Int32(length(storage.records)) + includes_stage_entry = law.includes_stage_entry + prior_capacity = includes_stage_entry ? capacity : Int32(0) + total = _candidate_record_capacity(1, Int(prior_capacity) + emitted, :keyed_reduce_candidate_capacity; int32_index = true, terminal = true) merges = total > _COMPACTED_BLOCK ? ceil(Int, log2(cld(total, _COMPACTED_BLOCK))) : 0 - law = publication.law - capacity = Int32(length(storage.records)) - bounds = _KeyedReduceBounds(capacity, Int32(total), Int32(merges)) - emission = _KeyedReduceEmission{W}(capacity) + bounds = _KeyedReduceBounds(capacity, includes_stage_entry, + Int32(total), Int32(merges)) + emission = _KeyedReduceEmission{W}(prior_capacity) key_order = _KeyedReduceKeyOrder{K}() fold = _KeyedReduceFold{K,V,typeof(law.operation),typeof(law.retention)}( - capacity, law.operation, law.seed.value, law.retention) + prior_capacity, law.operation, law.identity, law.retention) publication = _KeyedReducePublication{K,V}() return _KeyedReducePhysical(bounds, emission, key_order, fold, publication) end @@ -268,8 +272,12 @@ end state.positions[candidate] = Int32(0) end end - live = @inbounds storage.count[1] - valid_live = Int32(0) <= live <= bounds.capacity + live = Int32(0) + valid_live = true + if bounds.includes_stage_entry + live = @inbounds storage.count[1] + valid_live = Int32(0) <= live <= bounds.capacity + end if candidate <= bounds.capacity && valid_live && candidate <= live record = _compacted_load_value(eltype(storage.records), _compacted_record_components(storage.records), candidate) @@ -576,7 +584,7 @@ function _execute_keyed_reduce_stage!(prepared::_KeyedReduceStagePreparation, states.ordering) order = _keyed_reduce_final_order(bounds, states.ordering) _keyed_reduce_segments_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( - plan.key_order, states.segment, order, bounds.capacity, + plan.key_order, states.segment, order, plan.fold.prior_capacity, bounds.candidate_count; ndrange = extent) _compacted_launch_prefix_scan!(backend, states.segment.item_counts, states.segment.prefix, states.ordering.sums) diff --git a/src/execution/stage_preparation.jl b/src/execution/stage_preparation.jl index cfafb75..7cf77f5 100644 --- a/src/execution/stage_preparation.jl +++ b/src/execution/stage_preparation.jl @@ -267,9 +267,10 @@ struct _PreparedCollectLaw{T,K,G,O,P} order::O projection::P end -struct _PreparedKeyedReduceLaw{K,V,W,F,S,R} +struct _PreparedKeyedReduceLaw{K,V,W,F,R} operation::F - seed::S + identity::V + includes_stage_entry::Bool retention::R end struct _PreparedOrderedFoldLaw{T,F,O} @@ -791,8 +792,9 @@ function _prepared_keyed_reduce_law( stage = :prepare, contract = :keyed_reduce_key_capability, expected = (K, :global_load_store), actual = typeof(backend))) return _PreparedKeyedReduceLaw{K,V,W,typeof(law.operation), - typeof(law.seed),typeof(law.retention)}( - law.operation, law.seed, law.retention) + typeof(law.retention)}( + law.operation, law.seed.value, + _keyed_reduce_includes_stage_entry(law.seed), law.retention) end function _prepared_fold_law(backend, law::OrderedFold{T}, diff --git a/src/stage_model.jl b/src/stage_model.jl index b5bc264..464821d 100644 --- a/src/stage_model.jl +++ b/src/stage_model.jl @@ -1007,18 +1007,34 @@ struct NewKeyIdentity{V} end end -"""Retain keys whose reduced value equals the new-key identity.""" +"""Rebuild every emitted key from an exact identity, ignoring stage-entry keys.""" +struct RebuildFromIdentity{V} + value::V + function RebuildFromIdentity(value::V) where {V} + _storage_value_type(V) || throw(LocalMathValidationError( + "a keyed rebuild identity requires an admitted storage value type"; + stage = :construct, contract = :keyed_reduce_identity_type, + actual = V)) + return new{V}(value) + end +end + +@inline _keyed_reduce_includes_stage_entry(::NewKeyIdentity) = true +@inline _keyed_reduce_includes_stage_entry(::RebuildFromIdentity) = false + +"""Retain keys whose reduced value equals the keyed-reduction identity.""" struct RetainAllKeys end -"""Remove keys whose reduced value equals the new-key identity.""" +"""Remove keys whose reduced value equals the keyed-reduction identity.""" struct DropIdentityKeys end """ KeyedReduce(K, V, operation; maximum=1, seed, retention=DropIdentityKeys()) -Update one bounded sparse `Collection{KeyedValue{K,V}}`. Stage-entry records -must have unique exact keys. Existing values are folded before participating -contributions, followed by contributions in canonical `(source, lane)` order. +Update or rebuild one bounded sparse `Collection{KeyedValue{K,V}}`. +`NewKeyIdentity` includes unique stage-entry records before participating +contributions. `RebuildFromIdentity` ignores stage-entry records. Contributions +always fold in canonical `(source, lane)` order. """ struct KeyedReduce{K,V,W,F,S,R} operation::F @@ -1039,10 +1055,11 @@ struct KeyedReduce{K,V,W,F,S,R} "KeyedReduce emission width must be a reviewed small static bound"; stage = :construct, contract = :keyed_reduce_emission_width, expected = 1:32, actual = W)) - seed isa NewKeyIdentity{V} || throw(LocalMathValidationError( - "KeyedReduce requires an exact-typed new-key identity"; - stage = :construct, contract = :keyed_reduce_seed, - expected = V, actual = typeof(seed))) + seed isa Union{NewKeyIdentity{V},RebuildFromIdentity{V}} || throw( + LocalMathValidationError( + "KeyedReduce requires an exact-typed incremental or rebuild identity"; + stage = :construct, contract = :keyed_reduce_seed, + expected = V, actual = typeof(seed))) retention isa Union{RetainAllKeys,DropIdentityKeys} || throw( LocalMathValidationError( "KeyedReduce requires an explicit key-retention law"; @@ -1930,7 +1947,8 @@ function _validate_stage_publication_domain( end if publication.law isa Union{Collect,KeyedReduce} width = _publication_width(publication.law) - prior = publication.law isa KeyedReduce ? + prior = publication.law isa KeyedReduce && + _keyed_reduce_includes_stage_entry(publication.law.seed) ? Int(only(publication.components).collection.capacity) : 0 length(source) <= div(Int(typemax(Int32) - 1) - prior, width) || throw( LocalMathValidationError( diff --git a/test/fixtures/keyed_reduce_contracts.jl b/test/fixtures/keyed_reduce_contracts.jl index 53ee355..be3573f 100644 --- a/test/fixtures/keyed_reduce_contracts.jl +++ b/test/fixtures/keyed_reduce_contracts.jl @@ -39,8 +39,13 @@ struct OverflowKeyedReduceEvaluator end @inline (::OverflowKeyedReduceEvaluator)(item::Int32, reads, parameters) = (; delta = LocalMath.KeyedContribution(Int32(item), Int32(1))) +struct RebuildOrderedKeyedReduceEvaluator end +@inline (::RebuildOrderedKeyedReduceEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution(Int32(3), item)) + function _shared_keyed_program(backend, source_count, key_type, capacity, evaluator; maximum = 1, operation = +, + seed = LocalMath.NewKeyIdentity(Int32(0)), retention = LocalMath.DropIdentityKeys(), parameters = (), control = LocalMath.Control(), storage = nothing) source = LocalMath.Space(KeyedReduceContractNode, source_count) @@ -48,8 +53,7 @@ function _shared_keyed_program(backend, source_count, key_type, capacity, stage = LocalMath.Stage(source, NamedTuple(), ( LocalMath.Publication(records, LocalMath.KeyedReduce(key_type, Int32, operation; - maximum, seed = LocalMath.NewKeyIdentity(Int32(0)), - retention); value = :delta),), + maximum, seed, retention); value = :delta),), LocalMath.Evaluator(evaluator, parameters), control, LocalMath.SourceOrigin(:keyed_reduce_contract, 2)) schema = isempty(parameters) ? LocalMath.ParameterSchema() : @@ -187,6 +191,86 @@ function keyed_reduce_contract(backend) @test collect(LocalMath.Adapt.adapt( Array, count_storage.records)) == count_before + rebuild_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 2), + Int32[4, 4], Int32[6, 7], 2) + rebuild, rebuild_storage = _shared_keyed_program(backend, 2, + Int32, 2, OverflowKeyedReduceEvaluator(); + seed = LocalMath.RebuildFromIdentity(Int32(0)), + storage = rebuild_storage) + wait(LocalMath.execute!(rebuild)) + @test Array(rebuild_storage.count) == Int32[2] + @test collect(LocalMath.Adapt.adapt( + Array, rebuild_storage.records))[1:2] == [ + LocalMath.KeyedValue(Int32(1), Int32(1)), + LocalMath.KeyedValue(Int32(2), Int32(1)), + ] + + ordered_rebuild_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 4), + Int32[8, 0, 0, 0], Int32[9, 0, 0, 0], 1) + ordered_rebuild, ordered_rebuild_storage = _shared_keyed_program( + backend, 2, Int32, 4, RebuildOrderedKeyedReduceEvaluator(); + operation = OrderedLaneDecimalFold(), + seed = LocalMath.RebuildFromIdentity(Int32(0)), + retention = LocalMath.RetainAllKeys(), + storage = ordered_rebuild_storage) + wait(LocalMath.execute!(ordered_rebuild)) + @test Array(ordered_rebuild_storage.count) == Int32[1] + @test collect(LocalMath.Adapt.adapt( + Array, ordered_rebuild_storage.records))[1] == + LocalMath.KeyedValue(Int32(3), Int32(12)) + incremental_peer, _ = _shared_keyed_program( + backend, 2, Int32, 4, RebuildOrderedKeyedReduceEvaluator(); + operation = OrderedLaneDecimalFold(), + seed = LocalMath.NewKeyIdentity(Int32(0)), + retention = LocalMath.RetainAllKeys()) + rebuild_facts = LocalMath.inspect(ordered_rebuild) + incremental_facts = LocalMath.inspect(incremental_peer) + @test only(only(rebuild_facts.stages).publications).details.law.seed isa + LocalMath.RebuildFromIdentity{Int32} + @test LocalMath.lowering_identity(ordered_rebuild.plan) === + LocalMath.lowering_identity(incremental_peer.plan) + @test only(rebuild_facts.planning.physical_segments).family === + only(incremental_facts.planning.physical_segments).family === + :sparse_keyed_sequence + wait(LocalMath.execute!(rebuild)) + @test collect(LocalMath.Adapt.adapt( + Array, rebuild_storage.records))[1:2] == [ + LocalMath.KeyedValue(Int32(1), Int32(1)), + LocalMath.KeyedValue(Int32(2), Int32(1)), + ] + + rebuild_overflow_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 1), + Int32[9], Int32[4], 1) + rebuild_overflow, rebuild_overflow_storage = _shared_keyed_program( + backend, 2, Int32, 1, OverflowKeyedReduceEvaluator(); + seed = LocalMath.RebuildFromIdentity(Int32(0)), + storage = rebuild_overflow_storage) + rebuild_overflow_before = collect(LocalMath.Adapt.adapt( + Array, rebuild_overflow_storage.records)) + failure = _keyed_failure(LocalMath.execute!(rebuild_overflow)) + @test failure isa LocalMath.LocalMathValidationError + @test failure.actual.failure_class === :capacity_overflow + @test Array(rebuild_overflow_storage.count) == Int32[1] + @test collect(LocalMath.Adapt.adapt(Array, + rebuild_overflow_storage.records)) == rebuild_overflow_before + + ignored_count_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 1), + Int32[8], Int32[5], -1) + ignored_count, ignored_count_storage = _shared_keyed_program(backend, + 0, Int32, 1, EmptyKeyedReduceEvaluator(); + seed = LocalMath.RebuildFromIdentity(Int32(0)), + storage = ignored_count_storage) + wait(LocalMath.execute!(ignored_count)) + @test Array(ignored_count_storage.count) == Int32[0] + limit = LocalMath.Parameter(:limit, Int32; bounds = (Int32(0), Int32(2))) control_storage = _initialize_scalar_keyed_storage!( diff --git a/test/test_keyed_reduce_stage.jl b/test/test_keyed_reduce_stage.jl index db00011..344209b 100644 --- a/test/test_keyed_reduce_stage.jl +++ b/test/test_keyed_reduce_stage.jl @@ -4,6 +4,10 @@ import LocalMath const LMKR = LocalMath include("fixtures/keyed_reduce_contracts.jl") +@test_throws LMKR.LocalMathValidationError LMKR.RebuildFromIdentity(:metadata) +@test_throws LMKR.LocalMathValidationError LMKR.KeyedReduce( + Int32, Int32, +; seed = LMKR.RebuildFromIdentity(0.0f0)) + keyed_reduce_contract(KernelAbstractions.CPU()) struct KeyedReduceNode end diff --git a/test/test_public_api.jl b/test/test_public_api.jl index 5be482a..90e8a32 100644 --- a/test/test_public_api.jl +++ b/test/test_public_api.jl @@ -45,7 +45,7 @@ :ConditionalRoutedUniqueValue, :Contribution, :RoutedContribution, :ResolutionValue, :RoutedResolutionValue, :CollectedValue, :GroupedCollectedValue, :KeyedValue, :KeyedContribution, :FoldValue, - :NewKeyIdentity, + :NewKeyIdentity, :RebuildFromIdentity, :RetainAllKeys, :DropIdentityKeys, )) public_qualified = Set(filter(