From ee8729b8db85d6f2b3410a576b0064506992b30b Mon Sep 17 00:00:00 2001 From: PraneethMerugu Date: Mon, 14 Sep 2026 19:54:30 -0400 Subject: [PATCH] Add exact keyed collection reduction --- benchmark/keyed_reduce_compiler.jl | 175 +++++++ docs/src/api/localmath.md | 52 +- spec/localmath.md | 23 + src/LocalMath.jl | 6 +- src/bound_law.jl | 36 +- src/execution/collect_physical_support.jl | 112 +++- src/execution/collect_stage.jl | 127 +---- src/execution/keyed_reduce_stage.jl | 602 ++++++++++++++++++++++ src/execution/program_inspection.jl | 37 +- src/execution/stage_preparation.jl | 89 ++++ src/execution/stage_program.jl | 50 ++ src/stage_model.jl | 164 +++++- src/stage_planning.jl | 19 +- test/fixtures/keyed_reduce_contracts.jl | 209 ++++++++ test/metal/keyed_reduce.jl | 8 + test/metal/runtests.jl | 1 + test/runtests.jl | 1 + test/test_keyed_reduce_stage.jl | 142 +++++ test/test_public_api.jl | 11 +- 19 files changed, 1712 insertions(+), 152 deletions(-) create mode 100644 benchmark/keyed_reduce_compiler.jl create mode 100644 src/execution/keyed_reduce_stage.jl create mode 100644 test/fixtures/keyed_reduce_contracts.jl create mode 100644 test/metal/keyed_reduce.jl create mode 100644 test/test_keyed_reduce_stage.jl diff --git a/benchmark/keyed_reduce_compiler.jl b/benchmark/keyed_reduce_compiler.jl new file mode 100644 index 0000000..fc1612d --- /dev/null +++ b/benchmark/keyed_reduce_compiler.jl @@ -0,0 +1,175 @@ +#!/usr/bin/env julia + +# Reproducible compiler/allocation evidence for the warm sparse-key boundary. +import KernelAbstractions +import LocalMath +import TOML + +struct KeyedCompilerNode end +struct KeyedCompilerEvaluator end +@inline (::KeyedCompilerEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution( + (UInt32(isodd(item)), UInt32(item)), Int32(1))) + +struct CollectCompilerEvaluator end +@inline (::CollectCompilerEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.CollectedValue(item)) + +struct KeyedCompilerSubtract end +@inline (::KeyedCompilerSubtract)(left::Int32, right::Int32) = left - right + +function keyed_compiler_preparation(capacity::Int; + operation = +, retention = LocalMath.DropIdentityKeys()) + source = LocalMath.Space(KeyedCompilerNode, 3) + key_type = Tuple{UInt32,UInt32} + collection = LocalMath.Collection( + LocalMath.KeyedValue{key_type,Int32}, capacity) + stage = LocalMath.Stage(source, NamedTuple(), ( + LocalMath.Publication(collection, + LocalMath.KeyedReduce(key_type, Int32, operation; + seed = LocalMath.NewKeyIdentity(Int32(0)), retention); + value = :delta),), + LocalMath.Evaluator(KeyedCompilerEvaluator()), LocalMath.Control(), + LocalMath.SourceOrigin(:keyed_reduce_compiler, 1)) + return LocalMath.prepare(LocalMath.LocalLaw(stage), + collection => LocalMath.Allocate(); + backend = KernelAbstractions.CPU()) +end + +function collect_compiler_preparation(capacity::Int) + source = LocalMath.Space(KeyedCompilerNode, 3) + collection = LocalMath.Collection(Int32, capacity) + stage = LocalMath.Stage(source, NamedTuple(), ( + LocalMath.Publication(collection, + LocalMath.Collect(Int32; maximum = 1); value = :delta),), + LocalMath.Evaluator(CollectCompilerEvaluator()), LocalMath.Control(), + LocalMath.SourceOrigin(:collect_compiler_control, 1)) + return LocalMath.prepare(LocalMath.LocalLaw(stage), + collection => LocalMath.Allocate(); + 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)) + 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} + return Dict( + "statement_count" => length(info.code), + "call_count" => calls, + "any_ssa_count" => length(any_indices), + "any_control_flow_count" => count( + index -> control_flow(info.code[index]), any_indices), + "any_value_count" => count( + index -> !control_flow(info.code[index]), any_indices), + "return_type" => string(return_type), + ) +end + +function warm_public_execution_allocations(prepared) + wait(LocalMath.execute!(prepared)) + return minimum(@allocated(wait(LocalMath.execute!(prepared))) for _ in 1:5) +end + +function keyed_compiler_metrics(prepared) + launch = only(getfield(getfield(prepared, :runtime), :launches)) + stage = getfield(launch, :stage) + validation = LocalMath._ProgramValidationTarget( + getfield(getfield(prepared, :runtime), :execution_gate), Int32(1)) + signature = Tuple{typeof(stage),Tuple{},Int32,Tuple{}, + typeof(getfield(launch, :guard)),typeof(validation)} + host = typed_metrics(LocalMath._execute_keyed_reduce_stage!, signature) + execution = getfield(stage, :execution) + plan = getfield(execution, :plan) + states = getfield(execution, :states) + emission = LocalMath.KeyedContribution( + (UInt32(1), UInt32(1)), Int32(1)) + semantic_signature = Tuple{typeof(plan.emission),typeof(states.emission), + typeof(emission),Int32} + semantic = typed_metrics(LocalMath._keyed_reduce_materialize!, + semantic_signature) + allocated = warm_public_execution_allocations(prepared) + return Dict( + "host_orchestration" => host, + "emission_boundary" => semantic, + "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}), + "publish_boundary" => typed_metrics(LocalMath._keyed_reduce_publish_record!, + Tuple{typeof(plan.publication),typeof(states.publication), + typeof(execution.storage),Int32,Int32}), + "warm_public_execution_allocated_bytes" => allocated, + "host_any_classes" => [ + "KernelAbstractions launch construction (_svec_ref and kwcall)", + "compacted scan/order launch calls whose host return is unused", + "host control-flow and return nodes represented as Any by CodeInfo", + ], + "allocation_root_classes" => [ + "shared ExecutionReceipt and validation-status grouping", + "KernelAbstractions Kernel, NDRange, keyword, and argument tuples per launch", + ], + "prepared_stage_type" => string(typeof(stage)), + ) +end + +function collect_compiler_metrics(prepared) + launch = only(getfield(getfield(prepared, :runtime), :launches)) + stage = getfield(launch, :stage) + validation = LocalMath._ProgramValidationTarget( + getfield(getfield(prepared, :runtime), :execution_gate), Int32(1)) + signature = Tuple{typeof(stage),Tuple{},Int32,Tuple{}, + typeof(getfield(launch, :guard)),typeof(validation)} + return typed_metrics(LocalMath._execute_collect_stage!, signature) +end + +preparations = map(capacity -> keyed_compiler_preparation(capacity), (4, 8)) +variants = ( + keyed_compiler_preparation(4; operation = +, + retention = LocalMath.DropIdentityKeys()), + keyed_compiler_preparation(4; operation = +, + retention = LocalMath.RetainAllKeys()), + keyed_compiler_preparation(4; operation = KeyedCompilerSubtract(), + retention = LocalMath.DropIdentityKeys()), + keyed_compiler_preparation(4; operation = KeyedCompilerSubtract(), + retention = LocalMath.RetainAllKeys()), +) +metrics = keyed_compiler_metrics(first(preparations)) +control = collect_compiler_preparation(4) +metrics["collect_host_orchestration"] = collect_compiler_metrics(control) +control_allocated = warm_public_execution_allocations(control) +metrics["collect_control_allocated_bytes"] = control_allocated +metrics["keyed_incremental_allocated_bytes"] = + metrics["warm_public_execution_allocated_bytes"] - control_allocated +metrics["capacity_specialization_count"] = length(unique(typeof( + getfield(only(getfield(getfield(prepared, :runtime), :launches)), :stage)) + for prepared in preparations)) +variant_parts = map(variants) do prepared + execution = getfield(getfield(only(getfield( + getfield(prepared, :runtime), :launches)), :stage), :execution) + (; plan = execution.plan, states = execution.states, + storage = execution.storage) +end +metrics["operation_retention_specializations"] = Dict( + "bounds" => length(unique(typeof(part.plan.bounds) for part in variant_parts)), + "emission" => length(unique(typeof(part.plan.emission) for part in variant_parts)), + "sort" => length(unique((typeof(part.plan.key_order), + typeof(part.states.ordering)) for part in variant_parts)), + "segment" => length(unique((typeof(part.plan.key_order), + typeof(part.states.segment)) for part in variant_parts)), + "fold" => length(unique((typeof(part.plan.fold), + typeof(part.states.fold)) for part in variant_parts)), + "finalize" => length(unique((typeof(part.plan.bounds), + typeof(part.states.final)) for part in variant_parts)), + "publish" => length(unique((typeof(part.plan.publication), + 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" +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 4e3ce6e..8727753 100644 --- a/docs/src/api/localmath.md +++ b/docs/src/api/localmath.md @@ -52,6 +52,52 @@ 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 + +`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. + +```julia +import KernelAbstractions +import LocalMath + +struct Contact end +struct ContactDeltas end + +@inline function (::ContactDeltas)(contact::Int32, reads, parameters) + owner = UInt32(isodd(contact) ? 1 : 2) + delta = isodd(contact) ? Int32(1) : Int32(-1) + return (; change = LocalMath.KeyedContribution(owner, delta)) +end + +contacts = LocalMath.Space(Contact, 4) +counts = LocalMath.Collection(LocalMath.KeyedValue{UInt32,Int32}, 8) +stage = LocalMath.Stage(contacts, NamedTuple(), ( + LocalMath.Publication(counts, + LocalMath.KeyedReduce(UInt32, Int32, +; + maximum = 1, + seed = LocalMath.NewKeyIdentity(Int32(0)), + retention = LocalMath.DropIdentityKeys()); + value = :change),), + LocalMath.Evaluator(ContactDeltas()), LocalMath.Control(), + LocalMath.SourceOrigin(:contact_counts, 1)) +law = LocalMath.LocalLaw(stage) +prepared = LocalMath.prepare(law, counts => LocalMath.Allocate(); + backend = KernelAbstractions.CPU()) +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. + ## Public surface Ordinary authoring exports only the mathematical and execution vocabulary: @@ -71,11 +117,11 @@ 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`, `one_group`, `group_by`, `source_order`, `canonical_by`, `persistent_source_position` | -| Publication laws | `Unique`, `Reduce`, `Resolve`, `Collect`, `OrderedFold`, `TotalCoverage`, `PartialCoverage`, `UnreachableEmpty`, `PreserveEmpty`, `FillEmpty`, `IdentitySeed`, `ExistingSeed`, `CanonicalLeftFold`, `RelaxedAtomic`, `ArgMin`, `ArgMax`, `CanonicalSourceLaneTie`, `TieMin`, `TieMax`, `RejectOverflow`, `EmptyCollection` | +| 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` | | 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`, `FoldValue` | +| Evaluator outputs | `UniqueValue`, `ConditionalUniqueValue`, `RoutedUniqueValue`, `ConditionalRoutedUniqueValue`, `Contribution`, `RoutedContribution`, `ResolutionValue`, `RoutedResolutionValue`, `CollectedValue`, `GroupedCollectedValue`, `KeyedContribution`, `FoldValue` | | Advanced execution | `allocate_workspace`, `submission_capacity`, `ispending`, `success_gate` | These qualified names are stable interfaces, not permission to access other diff --git a/spec/localmath.md b/spec/localmath.md index 917d24c..957a9bd 100644 --- a/spec/localmath.md +++ b/spec/localmath.md @@ -19,6 +19,7 @@ Publication laws define observable conflict behavior: - reduction applies its declared operation and ordering law; - resolution selects a score/payload pair with explicit tie and empty rules; - collection publishes bounded records with explicit grouping and overflow; +- keyed reduction updates bounded sparse exact-key state by a canonical fold; - ordered fold applies a bounded recurrence in its declared canonical order. In authored `resolve_to` expressions, a noncanonical tie is an explicit @@ -66,6 +67,28 @@ source-position lanes. Collection production distinguishes its global storage capacity from the per-source `maximum` emission width. These forms lower to the existing `CollectionAccess`, source-position, `Collect`, and control laws. +`KeyedReduce(K, V, operation; seed, ...)` is the narrow +incremental sparse-state companion to `Collect`. Its sole destination is a +`Collection{KeyedValue{K,V}}`. `K` is `Int32`, `UInt32`, or a bounded flat +tuple of those types. The stage-entry collection must contain unique keys. +Each exact-key segment folds its existing value first when present, then +participating `KeyedContribution`s in intrinsic `(source item, lane)` left-fold +order. It has no order selector, hash, registry, relaxed, or provider-specific +path. `DropIdentityKeys` frees capacity when the final +value equals the declared identity, while `RetainAllKeys` preserves it. +Invalid prior count, duplicate prior keys, invalid stage control, and final +capacity overflow reject the complete publication. Private workspace receives +all intermediate values; records and the device-resident count publish through +one validation gate, so failure leaves both unchanged. Workspace is +`O(capacity + source_count * maximum)` and exact-key sorting followed by +segmented folding is `O(n log n)` work plus linear scans. +Its bounded device workspace is allocated by `prepare`; execution performs no +device allocation. The public host `execute!`/`wait` path still allocates +receipt, launch, and event bookkeeping shared with other Stage executors, so +zero host allocation is not part of this contract. The reproducible compiler +benchmark reports both that shared baseline and the narrower keyed semantic +boundary rather than treating host orchestration `Any` values as device IR. + Ordered recurrence is authored by declaring a total event order and every evolving state component with its exact initial Field: diff --git a/src/LocalMath.jl b/src/LocalMath.jl index 000121c..91ddf57 100644 --- a/src/LocalMath.jl +++ b/src/LocalMath.jl @@ -32,7 +32,7 @@ public sequence, allocate_workspace, submission_capacity, ispending, success_gat public one_group, group_by, source_order, canonical_by public persistent_source_position public CompactedStorage, BoundedGroupView -public Unique, Reduce, Resolve, Collect, OrderedFold +public Unique, Reduce, Resolve, Collect, KeyedReduce, OrderedFold public TotalCoverage, PartialCoverage, UnreachableEmpty, PreserveEmpty, FillEmpty public IdentitySeed, ExistingSeed, CanonicalLeftFold, RelaxedAtomic public ArgMin, ArgMax, CanonicalSourceLaneTie, TieMin, TieMax @@ -46,7 +46,8 @@ public BoundedFoldOutcome, evaluate_bounded public UniqueValue, ConditionalUniqueValue, RoutedUniqueValue public ConditionalRoutedUniqueValue, Contribution, RoutedContribution public ResolutionValue, RoutedResolutionValue, CollectedValue -public GroupedCollectedValue, FoldValue +public GroupedCollectedValue, KeyedValue, KeyedContribution, FoldValue +public NewKeyIdentity, RetainAllKeys, DropIdentityKeys import Adapt import Atomix import KernelAbstractions @@ -89,6 +90,7 @@ include("execution/ordered_fold_stage.jl") include("execution/fixed_lane_support.jl") include("execution/collect_physical_support.jl") include("execution/collect_stage.jl") +include("execution/keyed_reduce_stage.jl") include("execution/stage_program_kernelabstractions.jl") include("execution/stage_program.jl") include("execution/program_inspection.jl") diff --git a/src/bound_law.jl b/src/bound_law.jl index 7f6a2d6..b29216d 100644 --- a/src/bound_law.jl +++ b/src/bound_law.jl @@ -265,28 +265,36 @@ function _require_definite_field_initialization(law::LocalLaw, field::Field) return nothing end -function _collect_allocation_schema(law::LocalLaw, collection::Collection) +function _collection_allocation_schema(law::LocalLaw, collection::Collection) schemas = Any[] for stage in law.stages, publication in stage.publications - publication.law isa Collect || continue + publication.law isa Union{Collect,KeyedReduce} || continue any(publication.components) do component component isa CollectionPublication && _same_descriptor(component.collection, collection) end || continue - law = publication.law - push!(schemas, ( - grouped = _is_grouped(law.groups), - groups = Int(_compacted_group_count(law.groups)), - persistent_source_positions = - law.projection isa _PersistentSourcePosition, - source_position_count = - length(stage.source) * _publication_width(law), - )) + publication_law = publication.law + if publication_law isa Collect + push!(schemas, ( + grouped = _is_grouped(publication_law.groups), + groups = Int(_compacted_group_count(publication_law.groups)), + persistent_source_positions = + publication_law.projection isa _PersistentSourcePosition, + source_position_count = + length(stage.source) * _publication_width(publication_law), + )) + else + push!(schemas, ( + grouped = false, groups = 1, + persistent_source_positions = false, + source_position_count = 0, + )) + end end isempty(schemas) && throw(LocalMathValidationError( - "an allocated Collection requires a producing Collect publication"; + "an allocated Collection requires a producing collection publication"; stage = :bind, contract = :collection_allocation_producer, - expected = :collect_publication, + expected = :collection_publication, actual = semantic_identity(collection), )) all(schema -> schema == first(schemas), schemas) || throw( @@ -306,7 +314,7 @@ function _collection_allocation( stage = :bind, contract = :collection_allocation_initialization, expected = :empty_collection, actual = request.initial, )) - schema = _collect_allocation_schema(law, collection) + schema = _collection_allocation_schema(law, collection) capacity = Int(collection.capacity) return CompactedStorage( backend, diff --git a/src/execution/collect_physical_support.jl b/src/execution/collect_physical_support.jl index bfa331f..b5afea8 100644 --- a/src/execution/collect_physical_support.jl +++ b/src/execution/collect_physical_support.jl @@ -1,10 +1,57 @@ -# Domain-neutral physical primitives shared by the sole Stage Collect executor. +# Domain-neutral physical primitives shared by compacted Stage executors. # This file owns no LocalLaw lowering, topology, phase graph, or alternate output -# declaration. Every kernel is launched only by `collect_stage.jl`. +# declaration. Every kernel is launched only by the sole StageProgram executor. const _COMPACTED_BLOCK = 256 -@inline function _collect_atomic_min!(array, index, value) +struct _CompactedScanLevel{A} + storage::A + offset::Int32 + count::Int32 +end +Adapt.@adapt_structure _CompactedScanLevel +Base.length(level::_CompactedScanLevel) = Int(level.count) +@inline Base.getindex(level::_CompactedScanLevel, index::Integer) = + @inbounds level.storage[Int(level.offset) + Int(index)] +@inline function Base.setindex!(level::_CompactedScanLevel, value, index::Integer) + @inbounds level.storage[Int(level.offset) + Int(index)] = value + return value +end + +@inline _compacted_scan_level(storage, offset::Int, count::Int) = + _CompactedScanLevel(storage, Int32(offset), Int32(count)) + +function _compacted_scan_storage_lengths(items::Int) + prefix_count = 0 + sums_count = 0 + current = items + while true + blocks = max(cld(current, _COMPACTED_BLOCK), 1) + prefix_count <= typemax(Int32) - current && + sums_count <= typemax(Int32) - blocks || throw( + LocalMathValidationError( + "compacted scan workspace exceeds Int32 device addressing"; + stage = :prepare, contract = :compacted_scan_capacity, + expected = 0:typemax(Int32), + actual = (prefix_count + current, sums_count + blocks))) + prefix_count += current + sums_count += blocks + current <= _COMPACTED_BLOCK && break + current = blocks + end + return prefix_count, sums_count +end + +@inline function _compacted_scan_level_count(items::Int) + levels = 1 + while items > _COMPACTED_BLOCK + items = max(cld(items, _COMPACTED_BLOCK), 1) + levels += 1 + end + return levels +end + +@inline function _compacted_atomic_min!(array, index, value) Atomix.@atomic min(array[index], value) return nothing end @@ -33,6 +80,8 @@ end return Expr(:block, expressions..., :(nothing)) end +@inline _compacted_reconstruct_value(::Type{T}, values...) where {T} = T(values...) + @generated function _compacted_load_value(::Type{T}, storage, index::Int) where {T} fieldcount(T) == 0 && return :(@inbounds storage[index]) values = map(1:fieldcount(T)) do field_index @@ -42,7 +91,7 @@ end end T <: Tuple && return Expr(:tuple, values...) T <: NamedTuple && return :($T(($(values...),))) - return :($T($(values...))) + return :(_compacted_reconstruct_value($T, $(values...))) end @inline _compacted_group(::_OneGroup, value) = Int32(1) @@ -71,15 +120,14 @@ end right = @inbounds order[position] same_group = workspace.groups === nothing || @inbounds(workspace.groups[left]) == @inbounds(workspace.groups[right]) - key_type = typeof(port).parameters[5] - identity_type = typeof(port).parameters[6] + key_type, identity_type = _compacted_order_types(port) same_order = _canonical_order_equal( _compacted_load_value(key_type, workspace.keys, Int(left)), _compacted_load_value(identity_type, workspace.identities, Int(left)), _compacted_load_value(key_type, workspace.keys, Int(right)), _compacted_load_value(identity_type, workspace.identities, Int(right))) same_group && same_order && - _collect_atomic_min!( + _compacted_atomic_min!( workspace.duplicate_position, 1, Int32(position - 1)) end end @@ -140,6 +188,53 @@ end end end +function _compacted_launch_prefix_scan!(backend, item_counts, prefix_storage, + sums_storage) + current = length(item_counts) + prefix_offset = 0 + sums_offset = 0 + blocks = max(cld(current, _COMPACTED_BLOCK), 1) + output = _compacted_scan_level(prefix_storage, prefix_offset, current) + sums = _compacted_scan_level(sums_storage, sums_offset, blocks) + extent = blocks * _COMPACTED_BLOCK + _compacted_scan_block_kernel!(backend, _COMPACTED_BLOCK, extent)( + item_counts, output, sums, Int32(current); ndrange = extent) + prefix_offset += current + sums_offset += blocks + current = blocks + input = sums + while current > 1 + blocks = max(cld(current, _COMPACTED_BLOCK), 1) + output = _compacted_scan_level(prefix_storage, prefix_offset, current) + sums = _compacted_scan_level(sums_storage, sums_offset, blocks) + extent = blocks * _COMPACTED_BLOCK + _compacted_scan_block_kernel!(backend, _COMPACTED_BLOCK, extent)( + input, output, sums, Int32(length(input)); ndrange = extent) + prefix_offset += current + sums_offset += blocks + current <= _COMPACTED_BLOCK && break + current = blocks + input = sums + end + levels = _compacted_scan_level_count(length(item_counts)) + for level in (levels - 1):-1:1 + size = length(item_counts) + child_offset = 0 + for prior in 1:(level - 1) + child_offset += size + size = max(cld(size, _COMPACTED_BLOCK), 1) + end + parent_size = max(cld(size, _COMPACTED_BLOCK), 1) + parent_offset = child_offset + size + prefix = _compacted_scan_level(prefix_storage, child_offset, size) + parent = _compacted_scan_level(prefix_storage, parent_offset, parent_size) + extent = max(length(prefix), 1) + _compacted_scan_add_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + prefix, parent, Int32(length(prefix)); ndrange = extent) + end + return nothing +end + @kernel function _compacted_scatter_kernel!( valid, item_counts, item_prefix, order, positions, count, ::Val{K}, nitems::Int32 @@ -175,8 +270,7 @@ end left_group != right_group && return left_group < right_group end if workspace.keys !== nothing - key_type = typeof(port).parameters[5] - identity_type = typeof(port).parameters[6] + key_type, identity_type = _compacted_order_types(port) comparison = _canonical_order_compare( _compacted_load_value(key_type, workspace.keys, Int(left)), _compacted_load_value(identity_type, workspace.identities, Int(left)), diff --git a/src/execution/collect_stage.jl b/src/execution/collect_stage.jl index 548403b..3c5b4d0 100644 --- a/src/execution/collect_stage.jl +++ b/src/execution/collect_stage.jl @@ -7,54 +7,6 @@ const _COLLECT_STATUS_SUCCESS = Int32(0) const _COLLECT_STATUS_INVALID_CONTROL = Int32(5) -const _COLLECT_BLOCK = _COMPACTED_BLOCK - -struct _CollectScanLevel{A} - storage::A - offset::Int32 - count::Int32 -end -Adapt.@adapt_structure _CollectScanLevel -Base.length(level::_CollectScanLevel) = Int(level.count) -@inline Base.getindex(level::_CollectScanLevel, index::Integer) = - @inbounds level.storage[Int(level.offset) + Int(index)] -@inline function Base.setindex!(level::_CollectScanLevel, value, index::Integer) - @inbounds level.storage[Int(level.offset) + Int(index)] = value - return value -end - -@inline _collect_scan_level(storage, offset::Int, count::Int) = - _CollectScanLevel(storage, Int32(offset), Int32(count)) - -function _collect_scan_storage_lengths(items::Int) - prefix_count = 0 - sums_count = 0 - current = items - while true - blocks = max(cld(current, _COLLECT_BLOCK), 1) - prefix_count <= typemax(Int32) - current && - sums_count <= typemax(Int32) - blocks || throw( - LocalMathValidationError( - "Collect scan workspace exceeds Int32 device addressing"; - stage = :prepare, contract = :collect_scan_capacity, - expected = 0:typemax(Int32), - actual = (prefix_count + current, sums_count + blocks))) - prefix_count += current - sums_count += blocks - current <= _COLLECT_BLOCK && break - current = blocks - end - return prefix_count, sums_count -end - -@inline function _collect_scan_level_count(items::Int) - levels = 1 - while items > _COLLECT_BLOCK - items = max(cld(items, _COLLECT_BLOCK), 1) - levels += 1 - end - return levels -end struct _CollectPortPhysical{K,T,G,O,KT,IT} capacity::Int32 @@ -65,6 +17,9 @@ struct _CollectPortPhysical{K,T,G,O,KT,IT} merge_passes::Int32 end +_compacted_order_types(::_CollectPortPhysical{K,T,G,O,KT,IT}) where { + K,T,G,O,KT,IT} = (KT, IT) + _collect_width(::_CollectPortPhysical{K}) where {K} = K _collect_grouped(::_OneGroup) = Val(false) _collect_grouped(::_GroupBy) = Val(true) @@ -90,8 +45,8 @@ function _collect_port_physical(stage, publication::_PreparedStagePublication{C, # Prepared ordering is already effect/type admitted; the type parameters # merely specialize recursive record scratch and comparison kernels. sort_required = _is_grouped(law.groups) || _is_canonical_order(order) - merges = sort_required && candidates > _COLLECT_BLOCK ? - ceil(Int, log2(cld(candidates, _COLLECT_BLOCK))) : 0 + merges = sort_required && candidates > _COMPACTED_BLOCK ? + ceil(Int, log2(cld(candidates, _COMPACTED_BLOCK))) : 0 return _CollectPortPhysical{K,T,typeof(law.groups),typeof(order), key_type,identity_type}( Int32(length(storage.records)), Int32(candidates), law.groups, order, @@ -169,7 +124,7 @@ function _collect_port_workspace_spec( _collect_component_workspace_spec(root, index, :identities, typeof(port).parameters[6], candidates, ())..., ) : () - prefix_count, sums_count = _collect_scan_storage_lengths(items) + prefix_count, sums_count = _compacted_scan_storage_lengths(items) scan = ( _workspace_leaf(Symbol(:collect_, index, :_prefix), (root..., :collect, :ports, index, :prefix), Int32, @@ -276,7 +231,7 @@ function _collect_require_port_workspace(port, workspace) LocalMathValidationError("Collect workspace does not match its physical law"; stage = :prepare, contract = :collect_workspace_specialization, expected = (candidates, items, T), actual = :mismatched)) - prefix_count, sums_count = _collect_scan_storage_lengths(items) + prefix_count, sums_count = _compacted_scan_storage_lengths(items) length(workspace.prefix) == prefix_count && eltype(workspace.prefix) === Int32 && length(workspace.sums) == sums_count && @@ -424,7 +379,7 @@ end if workspace.groups !== nothing group = @inbounds workspace.groups[candidate] if !(Int32(1) <= group <= port.groups.count) - _collect_atomic_min!(workspace.invalid_group, 1, Int32(candidate)) + _compacted_atomic_min!(workspace.invalid_group, 1, Int32(candidate)) end end return Int32(1) @@ -570,70 +525,28 @@ end end function _collect_launch_scan!(backend, workspace) - current = length(workspace.item_counts) - prefix_offset = 0 - sums_offset = 0 - blocks = max(cld(current, _COLLECT_BLOCK), 1) - output = _collect_scan_level(workspace.prefix, prefix_offset, current) - sums = _collect_scan_level(workspace.sums, sums_offset, blocks) - extent = blocks * _COLLECT_BLOCK - _compacted_scan_block_kernel!(backend, _COLLECT_BLOCK, extent)( - workspace.item_counts, output, sums, Int32(current); ndrange = extent) - prefix_offset += current - sums_offset += blocks - current = blocks - input = sums - while current > 1 - blocks = max(cld(current, _COLLECT_BLOCK), 1) - output = _collect_scan_level(workspace.prefix, prefix_offset, current) - sums = _collect_scan_level(workspace.sums, sums_offset, blocks) - extent = blocks * _COLLECT_BLOCK - _compacted_scan_block_kernel!(backend, _COLLECT_BLOCK, extent)( - input, output, sums, Int32(length(input)); ndrange = extent) - prefix_offset += current - sums_offset += blocks - current <= _COLLECT_BLOCK && break - current = blocks - input = sums - end - levels = _collect_scan_level_count(length(workspace.item_counts)) - for level in (levels - 1):-1:1 - size = length(workspace.item_counts) - child_offset = 0 - for prior in 1:(level - 1) - child_offset += size - size = max(cld(size, _COLLECT_BLOCK), 1) - end - parent_size = max(cld(size, _COLLECT_BLOCK), 1) - parent_offset = child_offset + size - prefix = _collect_scan_level(workspace.prefix, child_offset, size) - parent = _collect_scan_level( - workspace.prefix, parent_offset, parent_size) - extent = max(length(prefix), 1) - _compacted_scan_add_kernel!(backend, min(extent, _COLLECT_BLOCK), extent)( - prefix, parent, Int32(length(prefix)); ndrange = extent) - end - return nothing + _compacted_launch_prefix_scan!(backend, workspace.item_counts, + workspace.prefix, workspace.sums) end function _collect_launch_order!(backend, plan, workspace) candidates = Int(plan.candidate_count) items = div(candidates, _collect_width(plan)) - prefix = _collect_scan_level(workspace.prefix, 0, items) + prefix = _compacted_scan_level(workspace.prefix, 0, items) extent = max(items, 1) - _compacted_scatter_kernel!(backend, min(extent, _COLLECT_BLOCK), extent)( + _compacted_scatter_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( workspace.valid, workspace.item_counts, prefix, workspace.order_a, workspace.positions, workspace.count, Val(_collect_width(plan)), Int32(items); ndrange = extent) plan.sort_required || return nothing - local_extent = max(cld(candidates, _COLLECT_BLOCK), 1) * _COLLECT_BLOCK - _compacted_local_bitonic_kernel!(backend, _COLLECT_BLOCK, local_extent)( + local_extent = max(cld(candidates, _COMPACTED_BLOCK), 1) * _COMPACTED_BLOCK + _compacted_local_bitonic_kernel!(backend, _COMPACTED_BLOCK, local_extent)( plan, workspace, workspace.count, Int32(candidates); ndrange = local_extent) - width, to_b = _COLLECT_BLOCK, true + width, to_b = _COMPACTED_BLOCK, true while width < candidates source, destination = to_b ? (workspace.order_a, workspace.order_b) : (workspace.order_b, workspace.order_a) - _compacted_merge_kernel!(backend, min(candidates, _COLLECT_BLOCK), candidates)( + _compacted_merge_kernel!(backend, min(candidates, _COMPACTED_BLOCK), candidates)( plan, workspace, source, destination, workspace.count, Int32(width), Int32(candidates); ndrange = candidates) width *= 2 @@ -641,13 +554,13 @@ function _collect_launch_order!(backend, plan, workspace) end if _is_grouped(plan.groups) groups = Int(plan.groups.count) + 1 - _compacted_directory_kernel!(backend, min(groups, _COLLECT_BLOCK), groups)( + _compacted_directory_kernel!(backend, min(groups, _COMPACTED_BLOCK), groups)( workspace, _collect_final_order(plan, workspace), Int32(plan.groups.count); ndrange = groups) end if _is_canonical_order(plan.order) extent = max(candidates, 1) - _compacted_validate_order_kernel!(backend, min(extent, _COLLECT_BLOCK), extent)( + _compacted_validate_order_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( plan, workspace, _collect_final_order(plan, workspace); ndrange = extent) end return nothing @@ -662,7 +575,7 @@ function _collect_publish_chunk!(backend, plans::Tuple, workspaces::Tuple, groupeds = map(plan -> _collect_grouped(plan.groups), plans) extent = maximum(Int, extents) _compacted_publish_ports_kernel!(backend, - min(extent, _COLLECT_BLOCK), extent)(storages, workspaces, gate, + min(extent, _COMPACTED_BLOCK), extent)(storages, workspaces, gate, groupeds, extents; ndrange = extent) return nothing end @@ -712,7 +625,7 @@ function _execute_collect_stage!(prepared::_CollectStagePreparation, predecessor_statuses = (relation_guard, predecessors...) extent = max(Int(execution.stage.source_count), maximum((Int(plan.candidate_count) for plan in execution.plans); init = 0), 1) - _collect_stage_reset_kernel!(backend, min(extent, _COLLECT_BLOCK), extent)( + _collect_stage_reset_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( execution.workspaces, execution.status, execution.gate, prepared.validation, lease_index, Int32(extent); ndrange = extent) diff --git a/src/execution/keyed_reduce_stage.jl b/src/execution/keyed_reduce_stage.jl new file mode 100644 index 0000000..ee32b0b --- /dev/null +++ b/src/execution/keyed_reduce_stage.jl @@ -0,0 +1,602 @@ +# Exact sparse-key reduction within the sole StageProgram/KernelAbstractions +# execution path. The destination Collection is read only into private scratch; +# publication occurs once, after every count, key, control, and capacity check. + +const _KEYED_REDUCE_STATUS_SUCCESS = Int32(0) +const _KEYED_REDUCE_STATUS_CAPACITY = Int32(1) +const _KEYED_REDUCE_STATUS_PRIOR_COUNT = Int32(2) +const _KEYED_REDUCE_STATUS_DUPLICATE = Int32(3) +const _KEYED_REDUCE_STATUS_INVALID_CONTROL = Int32(5) + +# Storage was type-admitted during preparation; this sealed reconstruction keeps +# that device path independent of the validating public outer constructor. +@inline _compacted_reconstruct_value(::Type{KeyedValue{K,V}}, key::K, + value::V) where {K,V} = KeyedValue(_CONSTRUCTION_TOKEN, key, value) + +struct _KeyedReduceBounds + capacity::Int32 + candidate_count::Int32 + merge_passes::Int32 +end +struct _KeyedReduceEmission{W} + first_candidate::Int32 +end +struct _KeyedReduceKeyOrder{K} end +_compacted_order_types(::_KeyedReduceKeyOrder{K}) where {K} = (K, K) +struct _KeyedReduceFold{K,V,F,R} + prior_capacity::Int32 + operation::F + identity::V + retention::R +end +struct _KeyedReducePublication{K,V} end +struct _KeyedReducePhysical{B,E,O,F,P} + bounds::B + emission::E + key_order::O + fold::F + publication::P +end + +_keyed_reduce_width(::_KeyedReduceEmission{W}) where {W} = W + +function _keyed_reduce_physical(stage, + publication::_PreparedStagePublication{C,<:_PreparedKeyedReduceLaw{K,V,W}}, + ) where {C,K,V,W} + 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, + :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) + key_order = _KeyedReduceKeyOrder{K}() + fold = _KeyedReduceFold{K,V,typeof(law.operation),typeof(law.retention)}( + capacity, law.operation, law.seed.value, law.retention) + publication = _KeyedReducePublication{K,V}() + 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} + if fieldcount(T) == 0 + suffix = isempty(path) ? :scalar : Symbol(join(string.(path), :_)) + return (_workspace_leaf(Symbol(: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, + (path..., fieldname(T, field_index)))...) + end +end + +function _keyed_reduce_stage_workspace_spec(stage; path::Tuple = (), + name_prefix::Symbol = :keyed_reduce_stage) + publication = only(stage.publications) + plan = _keyed_reduce_physical(stage, publication) + K = typeof(plan.key_order).parameters[1] + V = typeof(plan.fold).parameters[2] + 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)..., + _workspace_leaf(Symbol(name_prefix, :_valid), + (path..., :keyed_reduce, :valid), UInt8, (candidates,); + role = :keyed_reduce_candidate_participation), + _workspace_leaf(Symbol(name_prefix, :_item_counts), + (path..., :keyed_reduce, :item_counts), Int32, (candidates,); + role = :keyed_reduce_scan_input), + _workspace_leaf(Symbol(name_prefix, :_segment_flags), + (path..., :keyed_reduce, :segment_flags), Int32, (candidates,); + role = :keyed_reduce_segment_starts), + _workspace_leaf(Symbol(name_prefix, :_order_a), + (path..., :keyed_reduce, :order_a), Int32, (candidates,); + role = :keyed_reduce_order_ping), + _workspace_leaf(Symbol(name_prefix, :_order_b), + (path..., :keyed_reduce, :order_b), Int32, (candidates,); + role = :keyed_reduce_order_pong), + _workspace_leaf(Symbol(name_prefix, :_positions), + (path..., :keyed_reduce, :positions), Int32, (candidates,); + role = :keyed_reduce_candidate_position), + _workspace_leaf(Symbol(name_prefix, :_prefix), + (path..., :keyed_reduce, :prefix), Int32, (prefix_count,); + role = :keyed_reduce_scan_prefix), + _workspace_leaf(Symbol(name_prefix, :_sums), + (path..., :keyed_reduce, :sums), Int32, (sums_count,); + role = :keyed_reduce_scan_block_sums), + _workspace_leaf(Symbol(name_prefix, :_count), + (path..., :keyed_reduce, :count), Int32, (1,); + role = :keyed_reduce_candidate_count), + _workspace_leaf(Symbol(name_prefix, :_unique_count), + (path..., :keyed_reduce, :unique_count), Int32, (1,); + role = :keyed_reduce_unique_count), + _workspace_leaf(Symbol(name_prefix, :_final_count), + (path..., :keyed_reduce, :final_count), Int32, (1,); + role = :keyed_reduce_final_count), + _workspace_leaf(Symbol(name_prefix, :_duplicate), + (path..., :keyed_reduce, :duplicate), Int32, (1,); + role = :keyed_reduce_duplicate_diagnostic), + _workspace_leaf(Symbol(name_prefix, :_gate), + (path..., :keyed_reduce, :gate), Bool, (1,); + role = :keyed_reduce_publication_gate), + _workspace_leaf(Symbol(name_prefix, :_status), + (path..., :keyed_reduce, :status), Int32, (1,); + role = :keyed_reduce_diagnostic), + _workspace_leaf(Symbol(name_prefix, :_validation), + (path..., :keyed_reduce, :validation), UInt32, + (_VALIDATION_STATUS_FIELDS, 1); role = :validation_status), + ) + local_leaves = Tuple(_workspace_leaf(leaf.name, + leaf.path[(length(path) + 1):end], leaf.element_type, leaf.size; + strides = leaf.strides, role = leaf.role) for leaf in leaves) + return (leaves, template = _workspace_template_from_leaves(local_leaves), + plan, path) +end + +struct _KeyedReduceStageWorkspace{T,A,X} + tree::T + authority::A + spec::X +end + +function _keyed_reduce_stage_workspace_from_tree(tree, spec) + local_leaves = Tuple(_workspace_leaf(leaf.name, + leaf.path[(length(spec.path) + 1):end], leaf.element_type, leaf.size; + strides = leaf.strides, role = leaf.role) for leaf in spec.leaves) + authority = _WorkspaceAuthority(local_leaves, spec.template) + return _KeyedReduceStageWorkspace(tree, authority, spec) +end + +function Adapt.adapt_structure(to, workspace::_KeyedReduceStageWorkspace) + _keyed_reduce_stage_workspace_from_tree(Adapt.adapt(to, workspace.tree), + workspace.spec) +end + +struct _KeyedReduceStageExecution{Q,P,W,S,G,R} + stage::Q + plan::P + states::W + status::S + gate::G + storage::R +end +struct _KeyedReduceStagePreparation{B,E,V} + backend::B + execution::E + validation::V +end +Adapt.@adapt_structure _KeyedReduceStageExecution + +function _keyed_reduce_state_views(stored) + ordering = ( + valid = stored.valid, item_counts = stored.item_counts, + prefix = stored.prefix, sums = stored.sums, + order_a = stored.order_a, order_b = stored.order_b, + positions = stored.positions, count = stored.count, + groups = nothing, keys = stored.keys, identities = stored.keys, + ) + return ( + reset = ( + valid = stored.valid, item_counts = stored.item_counts, + segment_flags = stored.segment_flags, + order_a = stored.order_a, order_b = stored.order_b, + positions = stored.positions, keys = stored.keys, + values = stored.values, status = stored.status, + gate = stored.gate, count = stored.count, + unique_count = stored.unique_count, + final_count = stored.final_count, duplicate = stored.duplicate, + ), + emission = (valid = stored.valid, item_counts = stored.item_counts, + keys = stored.keys, values = stored.values), + ordering, + segment = (count = stored.count, keys = stored.keys, + segment_flags = stored.segment_flags, + item_counts = stored.item_counts, prefix = stored.prefix, + duplicate = stored.duplicate, + unique_count = stored.unique_count), + fold = (count = stored.count, segment_flags = stored.segment_flags, + prefix = stored.prefix, keys = stored.keys, values = stored.values, + reduced_keys = stored.reduced_keys, + reduced_values = stored.reduced_values, + item_counts = stored.item_counts), + final = (unique_count = stored.unique_count, prefix = stored.prefix, + item_counts = stored.item_counts, final_count = stored.final_count, + duplicate = stored.duplicate), + publication = (gate = stored.gate, + unique_count = stored.unique_count, + item_counts = stored.item_counts, prefix = stored.prefix, + reduced_keys = stored.reduced_keys, + reduced_values = stored.reduced_values, + final_count = stored.final_count), + ) +end + +function _prepare_keyed_reduce_stage(admission::_StageAdmission, + raw::_KeyedReduceStageWorkspace) + stored = raw.tree.keyed_reduce + states = _keyed_reduce_state_views(stored) + plan = _keyed_reduce_physical(admission.stage, + only(admission.stage.publications)) + candidates = Int(plan.bounds.candidate_count) + ordering = states.ordering + length(ordering.valid) == candidates && + length(ordering.item_counts) == candidates && + length(ordering.order_a) == candidates && + length(ordering.order_b) == candidates && + length(ordering.positions) == candidates || throw(LocalMathValidationError( + "KeyedReduce workspace does not match its physical law"; + stage = :prepare, contract = :keyed_reduce_workspace_specialization)) + for leaf in raw.authority.leaves + _centrally_qualified_value_capability(admission.backend, + leaf.element_type, :load, :global) && + _centrally_qualified_value_capability(admission.backend, + leaf.element_type, :store, :global) || throw(LocalMathValidationError( + "KeyedReduce workspace leaf lacks reviewed memory operations"; + stage = :prepare, contract = :keyed_reduce_backend_capability, + workspace_leaf = leaf.name, actual = typeof(admission.backend))) + end + storage = only(only(admission.stage.publications).components).storage + return _KeyedReduceStagePreparation(admission.backend, + _KeyedReduceStageExecution(admission.stage, plan, states, + stored.status, stored.gate, storage), stored.validation) +end + +@inline _keyed_reduce_set_failure!(status, code::Int32) = begin + Atomix.@atomic max(status[1], code) + nothing +end + +@kernel function _keyed_reduce_reset_kernel!(bounds, state, storage, + validation, lease::Int32) + candidate = @index(Global, Linear) + if candidate <= bounds.candidate_count + @inbounds begin + state.valid[candidate] = UInt8(0) + state.item_counts[candidate] = Int32(0) + state.segment_flags[candidate] = Int32(0) + state.order_a[candidate] = -Int32(candidate) + state.order_b[candidate] = -Int32(candidate) + state.positions[candidate] = Int32(0) + end + end + live = @inbounds storage.count[1] + valid_live = Int32(0) <= live <= bounds.capacity + if candidate <= bounds.capacity && valid_live && candidate <= live + record = _compacted_load_value(eltype(storage.records), + _compacted_record_components(storage.records), candidate) + @inbounds begin + _compacted_store_value!(state.keys, candidate, record.key) + _compacted_store_value!(state.values, candidate, record.value) + state.valid[candidate] = UInt8(1) + state.item_counts[candidate] = Int32(1) + end + end + if candidate == 1 + @inbounds begin + state.status[1] = valid_live ? _KEYED_REDUCE_STATUS_SUCCESS : + _KEYED_REDUCE_STATUS_PRIOR_COUNT + state.gate[1] = false + state.count[1] = Int32(0) + state.unique_count[1] = Int32(0) + state.final_count[1] = Int32(0) + state.duplicate[1] = bounds.candidate_count + Int32(1) + _clear_validation_status!(validation, lease) + end + end +end + +@inline _keyed_reduce_lane(value::KeyedContribution, ::Val{1}, ::Val{1}) = value +@inline _keyed_reduce_lane(values::Tuple, width, lane) = + _emission_lane(values, width, lane) + +@inline function _keyed_reduce_materialize_lane!(emission_layout, state, + emission, item::Int32, ::Val{L}) where {L} + candidate = Int(emission_layout.first_candidate) + L + + _keyed_reduce_width(emission_layout) * (Int(item) - 1) + enabled = emission.participates + @inbounds begin + state.valid[candidate] = enabled ? UInt8(1) : UInt8(0) + state.item_counts[candidate] = enabled ? Int32(1) : Int32(0) + end + enabled || return nothing + @inbounds begin + _compacted_store_value!(state.keys, candidate, emission.key) + _compacted_store_value!(state.values, candidate, emission.value) + end + return nothing +end + +@generated function _keyed_reduce_materialize!( + emission_layout::_KeyedReduceEmission{W}, state, emissions, + item::Int32) where {W} + calls = [quote + emission = _keyed_reduce_lane(emissions, Val($W), Val($lane)) + _keyed_reduce_materialize_lane!(emission_layout, state, emission, item, + Val($lane)) + end for lane in 1:W] + return Expr(:block, calls..., :(nothing)) +end + +@kernel function _keyed_reduce_evaluate_kernel!(qualified, emission_layout, state, + predecessors, status, lease::Int32) + raw = @index(Global, Linear) + item = Int32(raw) + stage = qualified.stage + if _candidate_prefix_succeeded(predecessors, lease) && + _stage_gate_open(stage.control.gate, stage, qualified.parameters) + prefix = _stage_prefix_value(stage.control.prefix, stage, + qualified.parameters) + valid_prefix = prefix isa Integer && !(prefix isa Bool) && + 0 <= prefix <= stage.source_count + item == 1 && !valid_prefix && _keyed_reduce_set_failure!(status, + _KEYED_REDUCE_STATUS_INVALID_CONTROL) + if item <= stage.source_count && valid_prefix + _, active = _stage_control_state(stage, qualified.parameters, item) + access_valid = _stage_accesses_valid(stage.accesses, + stage.fields, item) + access_valid || _keyed_reduce_set_failure!(status, + _KEYED_REDUCE_STATUS_INVALID_CONTROL) + if active && access_valid + result = _call_stage_evaluator(qualified, item, + _stage_reads(stage, item), qualified.parameters) + _keyed_reduce_materialize!(emission_layout, state, + getfield(result, 1), item) + end + end + end +end + +function _keyed_reduce_launch_order!(backend, bounds, key_order, state) + candidates = Int(bounds.candidate_count) + _compacted_launch_prefix_scan!(backend, state.item_counts, + state.prefix, state.sums) + prefix = _compacted_scan_level(state.prefix, 0, candidates) + _compacted_scatter_kernel!(backend, + min(max(candidates, 1), _COMPACTED_BLOCK), max(candidates, 1))( + state.valid, state.item_counts, prefix, state.order_a, + state.positions, state.count, Val(1), Int32(candidates); + ndrange = max(candidates, 1)) + local_extent = max(cld(candidates, _COMPACTED_BLOCK), 1) * _COMPACTED_BLOCK + _compacted_local_bitonic_kernel!(backend, _COMPACTED_BLOCK, local_extent)( + key_order, state, state.count, Int32(candidates); + ndrange = local_extent) + width, to_b = _COMPACTED_BLOCK, true + while width < candidates + source, destination = to_b ? (state.order_a, state.order_b) : + (state.order_b, state.order_a) + _compacted_merge_kernel!(backend, min(candidates, _COMPACTED_BLOCK), candidates)( + key_order, state, source, destination, state.count, + Int32(width), Int32(candidates); ndrange = candidates) + width *= 2 + to_b = !to_b + end + return nothing +end + +@inline function _keyed_reduce_final_order(bounds, state) + isodd(bounds.merge_passes) ? state.order_b : state.order_a +end + +@inline _keyed_reduce_key(key_order::_KeyedReduceKeyOrder{K}, state, + candidate::Int32) where {K} = + _compacted_load_value(K, state.keys, + Int(candidate)) + +@kernel function _keyed_reduce_segments_kernel!(key_order, state, order, + prior_capacity::Int32, extent::Int32) + raw = @index(Global, Linear) + position = Int32(raw) + live = @inbounds state.count[1] + flag = Int32(0) + if position <= live + candidate = @inbounds order[position] + flag = if position == 1 + Int32(1) + else + prior = @inbounds order[position - Int32(1)] + _rank_equal(_keyed_reduce_key(key_order, state, prior), + _keyed_reduce_key(key_order, state, candidate)) ? Int32(0) : Int32(1) + end + if flag == 0 && candidate <= prior_capacity + prior = @inbounds order[position - Int32(1)] + prior <= prior_capacity && Atomix.@atomic min( + state.duplicate[1], position - Int32(1)) + end + end + if position <= extent + @inbounds begin + state.segment_flags[position] = flag + state.item_counts[position] = flag + end + end +end + +@kernel function _keyed_reduce_unique_count_kernel!(state) + index = @index(Global, Linear) + if index == 1 + live = @inbounds state.count[1] + @inbounds state.unique_count[1] = live == 0 ? Int32(0) : + state.prefix[live] + state.item_counts[live] + end +end + +@inline _keyed_reduce_retains(::RetainAllKeys, value, identity) = true +@inline _keyed_reduce_retains(::DropIdentityKeys, value, identity) = value != identity + +@kernel function _keyed_reduce_clear_counts_kernel!(state, extent::Int32) + index = @index(Global, Linear) + index <= extent && (@inbounds state.item_counts[index] = Int32(0)) +end + +@inline _keyed_reduce_fold_key(::_KeyedReduceFold{K}, state, + candidate::Int32) where {K} = + _compacted_load_value(K, state.keys, Int(candidate)) + +@inline function _keyed_reduce_fold_segment!(fold, state, order, + position::Int32, live::Int32) + if position <= live && @inbounds(state.segment_flags[position]) == 1 + group = @inbounds state.prefix[position] + Int32(1) + candidate = @inbounds order[position] + key = _keyed_reduce_fold_key(fold, state, candidate) + value_type = typeof(fold).parameters[2] + cursor = position + accumulator = if candidate <= fold.prior_capacity + cursor += Int32(1) + _compacted_load_value(value_type, state.values, Int(candidate)) + else + fold.identity + end + while cursor <= live + next_candidate = @inbounds order[cursor] + _rank_equal(key, + _keyed_reduce_fold_key(fold, state, next_candidate)) || break + contribution = _compacted_load_value(value_type, + state.values, Int(next_candidate)) + accumulator = fold.operation(accumulator, contribution) + cursor += Int32(1) + end + @inbounds begin + _compacted_store_value!(state.reduced_keys, Int(group), key) + _compacted_store_value!(state.reduced_values, Int(group), accumulator) + retained = _keyed_reduce_retains(fold.retention, accumulator, + fold.identity) ? Int32(1) : Int32(0) + state.item_counts[group] = retained + end + end + return nothing +end + +@kernel function _keyed_reduce_fold_kernel!(fold, state, order, + extent::Int32) + raw = @index(Global, Linear) + position = Int32(raw) + live = @inbounds state.count[1] + position <= extent && + _keyed_reduce_fold_segment!(fold, state, order, position, live) +end + +@kernel function _keyed_reduce_final_count_kernel!(bounds, state, status) + index = @index(Global, Linear) + if index == 1 + unique_count = @inbounds state.unique_count[1] + final_count = unique_count == 0 ? Int32(0) : + @inbounds(state.prefix[unique_count] + + state.item_counts[unique_count]) + @inbounds state.final_count[1] = final_count + final_count > bounds.capacity && _keyed_reduce_set_failure!(status, + _KEYED_REDUCE_STATUS_CAPACITY) + @inbounds(state.duplicate[1]) <= bounds.candidate_count && + _keyed_reduce_set_failure!(status, + _KEYED_REDUCE_STATUS_DUPLICATE) + end +end + +@kernel function _keyed_reduce_finalize_kernel!(gate, status, validation, + program_validation, predecessors, lease::Int32) + index = @index(Global, Linear) + if index == 1 && + _candidate_prefix_succeeded(predecessors, lease) + code = @inbounds status[1] + if code == _KEYED_REDUCE_STATUS_SUCCESS + @inbounds gate[1] = true + else + @inbounds gate[1] = false + _store_validation_status!(validation, lease, code, Int32(1), + Int32(0), Int32(0), UInt32(0)) + _store_program_validation_status!(program_validation, lease, code, + Int32(1), Int32(0), Int32(0), UInt32(0)) + end + end +end + +@inline _keyed_reduce_publication_key(::_KeyedReducePublication{K}, state, + group::Int32) where {K} = + _compacted_load_value(K, state.reduced_keys, Int(group)) +@inline _keyed_reduce_publication_value( + ::_KeyedReducePublication{K,V}, state, group::Int32) where {K,V} = + _compacted_load_value(V, state.reduced_values, Int(group)) + +@inline function _keyed_reduce_publish_record!(publication, state, storage, + group::Int32, unique_count::Int32) + if group <= unique_count && @inbounds(state.item_counts[group]) == 1 + output = @inbounds state.prefix[group] + Int32(1) + key = _keyed_reduce_publication_key(publication, state, group) + value = _keyed_reduce_publication_value(publication, state, group) + _compacted_store_value!(_compacted_record_components(storage.records), + Int(output), KeyedValue(_CONSTRUCTION_TOKEN, key, value)) + @inbounds begin + storage.source_item[output] = Int32(0) + storage.source_lane[output] = Int32(0) + end + end + return nothing +end + +@kernel function _keyed_reduce_publish_kernel!(publication, state, storage) + raw = @index(Global, Linear) + group = Int32(raw) + if @inbounds state.gate[1] + unique_count = @inbounds state.unique_count[1] + _keyed_reduce_publish_record!(publication, state, storage, group, + unique_count) + group == 1 && (@inbounds storage.count[1] = state.final_count[1]) + end +end + +function _execute_keyed_reduce_stage!(prepared::_KeyedReduceStagePreparation, + parameters::Tuple, lease_index::Int32, predecessors::Tuple, + relation_guard, program_validation) + execution = prepared.execution + backend = prepared.backend + states, plan = execution.states, execution.plan + bounds = plan.bounds + qualified = _QualifiedEvaluation(_stage_evaluation(execution.stage), + _stage_runtime_parameters(parameters, execution.stage)) + statuses = (relation_guard, predecessors...) + extent = max(Int(bounds.candidate_count), 1) + _keyed_reduce_reset_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + bounds, states.reset, execution.storage, prepared.validation, lease_index; + ndrange = extent) + _launch_stage_relation_receipt!(backend, relation_guard, + prepared.validation, program_validation, lease_index) + _keyed_reduce_evaluate_kernel!(backend)(qualified, plan.emission, + states.emission, + statuses, execution.status, lease_index; + ndrange = max(Int(execution.stage.source_count), 1)) + _keyed_reduce_launch_order!(backend, bounds, plan.key_order, + 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, + bounds.candidate_count; ndrange = extent) + _compacted_launch_prefix_scan!(backend, states.segment.item_counts, + states.segment.prefix, states.ordering.sums) + _keyed_reduce_unique_count_kernel!(backend, 1, 1)(states.segment; + ndrange = 1) + _keyed_reduce_clear_counts_kernel!(backend, + min(extent, _COMPACTED_BLOCK), extent)(states.fold, + bounds.candidate_count; + ndrange = extent) + _keyed_reduce_fold_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + plan.fold, states.fold, order, bounds.candidate_count; ndrange = extent) + _compacted_launch_prefix_scan!(backend, states.fold.item_counts, + states.fold.prefix, states.ordering.sums) + _keyed_reduce_final_count_kernel!(backend, 1, 1)( + bounds, states.final, execution.status; ndrange = 1) + _keyed_reduce_finalize_kernel!(backend, 1, 1)(execution.gate, + execution.status, + prepared.validation, program_validation, statuses, lease_index; + ndrange = 1) + _keyed_reduce_publish_kernel!(backend, min(extent, _COMPACTED_BLOCK), extent)( + plan.publication, states.publication, execution.storage; ndrange = extent) + return prepared +end diff --git a/src/execution/program_inspection.jl b/src/execution/program_inspection.jl index 52c4bd1..4893e86 100644 --- a/src/execution/program_inspection.jl +++ b/src/execution/program_inspection.jl @@ -371,7 +371,7 @@ function _planned_stage_phases(entry::_StageLoweringEntry{ push!(phases, _phase_fact(:collect_evaluate)) for port in entry.workspace.ports items = div(Int(port.candidate_count), _collect_width(port)) - levels = _collect_scan_level_count(items) + levels = _compacted_scan_level_count(items) push!(phases, _phase_fact(:collect_scan_block, levels)) levels == 1 || push!(phases, _phase_fact(:collect_scan_add, levels - 1)) @@ -393,6 +393,40 @@ function _planned_stage_phases(entry::_StageLoweringEntry{ _POINTWISE_SEGMENT_LIMIT)))) return Tuple(phases) end +function _planned_stage_phases(entry::_StageLoweringEntry{ + A,W,<:_KeyedReduceStageExecutor}) where {A,W} + plan = entry.workspace.plan + levels = _compacted_scan_level_count(Int(plan.bounds.candidate_count)) + phases = Any[_phase_fact(:keyed_reduce_reset)] + append!(phases, _planned_relation_phases(entry)) + append!(phases, ( + _phase_fact(:keyed_reduce_evaluate), + _phase_fact(:keyed_reduce_scan_block, levels), + )) + levels == 1 || push!(phases, + _phase_fact(:keyed_reduce_scan_add, levels - 1)) + append!(phases, ( + _phase_fact(:keyed_reduce_scatter), + _phase_fact(:keyed_reduce_local_bitonic), + )) + plan.bounds.merge_passes == 0 || push!(phases, + _phase_fact(:keyed_reduce_merge, plan.bounds.merge_passes)) + push!(phases, _phase_fact(:keyed_reduce_segment)) + push!(phases, _phase_fact(:keyed_reduce_segment_prefix_scan_block, levels)) + levels == 1 || push!(phases, + _phase_fact(:keyed_reduce_segment_prefix_scan_add, levels - 1)) + append!(phases, ( + _phase_fact(:keyed_reduce_unique_count), + _phase_fact(:keyed_reduce_clear_retention), + _phase_fact(:keyed_reduce_fold), + )) + push!(phases, _phase_fact(:keyed_reduce_retention_prefix_scan_block, levels)) + levels == 1 || push!(phases, + _phase_fact(:keyed_reduce_retention_prefix_scan_add, levels - 1)) + append!(phases, (_phase_fact(:keyed_reduce_final_count), + _phase_fact(:keyed_reduce_finalize), _phase_fact(:keyed_reduce_publish))) + return Tuple(phases) +end function _planned_stage_phases(entry::_StageLoweringEntry{ A,W,<:_OrderedFoldStageExecutor}) where {A,W} phases = Any[_phase_fact(:ordered_fold_reset)] @@ -424,6 +458,7 @@ _stage_layout_name(::_CandidateStageExecutor{<:_GroupedCandidateLayout}) = _stage_layout_name(::_CandidateStageExecutor{<:_DirectIdentityUniqueLayout}) = :direct_identity_unique _stage_layout_name(::_CollectStageExecutor) = :compacted_sequence +_stage_layout_name(::_KeyedReduceStageExecutor) = :sparse_keyed_sequence _stage_layout_name(::_OrderedFoldStageExecutor) = :ordered_recurrence function _segment_materializations(law::LocalLaw, indices) diff --git a/src/execution/stage_preparation.jl b/src/execution/stage_preparation.jl index 83e4e0d..cfafb75 100644 --- a/src/execution/stage_preparation.jl +++ b/src/execution/stage_preparation.jl @@ -267,12 +267,19 @@ struct _PreparedCollectLaw{T,K,G,O,P} order::O projection::P end +struct _PreparedKeyedReduceLaw{K,V,W,F,S,R} + operation::F + seed::S + retention::R +end struct _PreparedOrderedFoldLaw{T,F,O} transition::F order::O end _publication_value_type(::_PreparedCollectLaw{T}) where {T} = T _publication_width(::_PreparedCollectLaw{T,K}) where {T,K} = K +_publication_value_type(::_PreparedKeyedReduceLaw{K,V}) where {K,V} = V +_publication_width(::_PreparedKeyedReduceLaw{K,V,W}) where {K,V,W} = W _publication_value_type(::_PreparedOrderedFoldLaw{T}) where {T} = T _publication_width(::_PreparedOrderedFoldLaw) = 1 struct _PreparedStagePublication{C,L}; components::C; law::L; end @@ -307,6 +314,7 @@ Adapt.@adapt_structure _PreparedFoldAccumulatorView # would create a second semantic authority rather than adapting device state. Adapt.adapt_structure(to, law::_PreparedOrderedFoldLaw) = law Adapt.adapt_structure(to, law::_PreparedCollectLaw) = law +Adapt.adapt_structure(to, law::_PreparedKeyedReduceLaw) = law Adapt.adapt_structure(to, publication::_PreparedStagePublication) = _PreparedStagePublication( Adapt.adapt(to, publication.components), publication.law) @@ -588,6 +596,39 @@ function _prepare_collect_storage( return _PreparedStageCollection(storage) end +function _prepare_keyed_reduce_storage( + validated::_ValidatedStructuralBinding, stage::Stage, + component::CollectionPublication, use::_ProjectedCollectionUse, + law::KeyedReduce{K,V}, + ) where {K,V} + binding = _collection_binding(validated, use.slot) + binding.collection == component.collection || throw(LocalMathValidationError( + "Collection projection resolves a conflicting semantic descriptor"; + stage = :prepare, contract = :collection_projection_schema, + expected = component.collection, actual = binding.collection)) + storage = binding.storage + capacity = Int(component.collection.capacity) + try + _validate_compacted_record_storage( + storage.records, KeyedValue{K,V}, capacity) + eltype(storage.count) === Int32 && size(storage.count) == (1,) || + throw(ArgumentError("count")) + storage.segment_starts === nothing || throw(ArgumentError("directory")) + storage.source_position === nothing || throw(ArgumentError("source position")) + all(provenance -> eltype(provenance) === Int32 && + size(provenance) == (capacity,), + (storage.source_item, storage.source_lane)) || + throw(ArgumentError("provenance")) + catch error + throw(LocalMathValidationError( + "Collection storage does not exactly realize its KeyedReduce law"; + stage = :prepare, contract = :keyed_reduce_storage_schema, + expected = (record_type = KeyedValue{K,V}, capacity), + actual = sprint(showerror, error))) + end + return _PreparedStageCollection(storage) +end + function _prepare_fold_state( projected::_ProjectedFoldState, law::OrderedFold, ) @@ -742,6 +783,18 @@ function _prepared_collect_law(backend, law::Collect{T,K}, ) end +function _prepared_keyed_reduce_law( + backend, law::KeyedReduce{K,V,W}, analysis_cache, + ) where {K,V,W} + _centrally_qualified_rank_type(backend, K) || throw(LocalMathValidationError( + "KeyedReduce key type lacks centrally reviewed comparison operations"; + 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) +end + function _prepared_fold_law(backend, law::OrderedFold{T}, analysis_cache::Dict{Any,Any}) where {T} order = _prepare_stage_order( @@ -770,6 +823,16 @@ function _prepare_stage_publication( (_prepare_collect_storage(validated, stage, component, use, law),), _prepared_collect_law(backend, law, analysis_cache), ) + elseif law isa KeyedReduce + component = only(publication.components) + use = only(uses) + use isa _ProjectedCollectionUse || throw(LocalMathValidationError( + "KeyedReduce requires a positional Collection projection"; + stage = :prepare, contract = :keyed_reduce_projection)) + return _PreparedStagePublication( + (_prepare_keyed_reduce_storage( + validated, stage, component, use, law),), + _prepared_keyed_reduce_law(backend, law, analysis_cache)) elseif law isa OrderedFold projected = only(uses) projected isa _ProjectedFoldState || throw(LocalMathValidationError( @@ -824,6 +887,32 @@ _validate_stage_publication_operation( backend, publication::Publication{C,<:Collect}, ) where {C} = nothing +function _validate_stage_publication_operation( + backend, publication::Publication{C,<:KeyedReduce{K,V}}, + analysis_cache = nothing, + ) where {C,K,V} + law = publication.law + signature = Tuple{V,V} + analysis = analysis_cache === nothing ? _closed_callable_effect_analysis( + law.operation, signature, + method_signature -> length(method_signature) == 3) : + _cached_closed_callable_effect_analysis!(analysis_cache, + :closed_keyed_reduce, law.operation, signature, + method_signature -> length(method_signature) == 3) + analysis.qualified && analysis.return_type === V || throw( + LocalMathValidationError( + "KeyedReduce operation fails its exact closed typed-IR contract"; + stage = :plan, contract = :keyed_reduce_operation_effects, + expected = (signature, return_type = V), + actual = (callable_type = typeof(law.operation), + selected_method = analysis.method, signature, + qualified = analysis.qualified, + return_type = analysis.return_type, + reason = analysis.reason, operation = analysis.operation), + hint = analysis.hint)) + return nothing +end + _validate_stage_publication_operation( backend, publication::Publication{C,<:OrderedFold}, ) where {C} = nothing diff --git a/src/execution/stage_program.jl b/src/execution/stage_program.jl index 4183b9f..1ce7466 100644 --- a/src/execution/stage_program.jl +++ b/src/execution/stage_program.jl @@ -9,6 +9,7 @@ struct _CandidateStageExecutor{L} layout::L end struct _CollectStageExecutor end +struct _KeyedReduceStageExecutor end struct _OrderedFoldStageExecutor end struct _StageEntryContext @@ -102,11 +103,14 @@ function _collect_stage_executor(publications::Tuple{ }) where {C,L<:_PreparedCollectLaw} return _collect_stage_executor(Base.tail(publications)) end +_stage_executor(::Tuple{<:_PreparedStagePublication{C,L}}) where { + C,L<:_PreparedKeyedReduceLaw} = _KeyedReduceStageExecutor() _stage_executor(::Tuple{<:_PreparedStagePublication{C,L}}) where { C,L<:_PreparedOrderedFoldLaw} = _OrderedFoldStageExecutor() _stage_executor_name(::_CandidateStageExecutor) = :candidate _stage_executor_name(::_CollectStageExecutor) = :collect +_stage_executor_name(::_KeyedReduceStageExecutor) = :keyed_reduce _stage_executor_name(::_OrderedFoldStageExecutor) = :ordered_fold _direct_identity_unique_lane(::Type, ::Unique) = false @@ -257,6 +261,12 @@ function _stage_workspace_spec( return _collect_stage_workspace_spec(admission.stage; path = (:stages, index), name_prefix = Symbol(:stage_, index)) end +function _stage_workspace_spec( + ::_KeyedReduceStageExecutor, admission::_StageAdmission, index::Int, + ) + return _keyed_reduce_stage_workspace_spec(admission.stage; + path = (:stages, index), name_prefix = Symbol(:stage_, index)) +end function _stage_workspace_spec( ::_OrderedFoldStageExecutor, admission::_StageAdmission, index::Int, ) @@ -316,6 +326,12 @@ _publication_law_inspection(law::Collect{T}) where {T} = ( groups = law.groups, order = law.order, projection = law.projection, conflicts = :collect, overflow = law.overflow, onempty = law.onempty, ) +_publication_law_inspection(law::KeyedReduce{K,V}) where {K,V} = ( + kind = :keyed_reduce, key_type = K, value_type = V, + maximum = _publication_width(law), operation = law.operation, + seed = law.seed, retention = law.retention, + conflicts = :canonical_source_lane_left_fold, +) _publication_law_inspection(law::OrderedFold{T}) where {T} = ( kind = :ordered_fold, value_type = T, state = map(_fold_state_component_inspection, law.state.components), @@ -825,6 +841,10 @@ _stage_program_workspace(entry::_StageLoweringEntry{A,W,<:_CollectStageExecutor} tree, lease_capacity::Int) where {A,W} = _collect_stage_workspace_from_tree(tree, _stage_entry_workspace_spec(entry, lease_capacity)) +_stage_program_workspace(entry::_StageLoweringEntry{A,W,<:_KeyedReduceStageExecutor}, + tree, lease_capacity::Int) where {A,W} = + _keyed_reduce_stage_workspace_from_tree(tree, + _stage_entry_workspace_spec(entry, lease_capacity)) _stage_program_workspace(entry::_StageLoweringEntry{A,W,<:_OrderedFoldStageExecutor}, tree, lease_capacity::Int) where {A,W} = _ordered_fold_stage_workspace_from_tree(tree, @@ -856,11 +876,16 @@ end _prepare_stage_entry(entry::_StageLoweringEntry{A,W,<:_CollectStageExecutor}, raw) where {A,W} = _prepare_collect_stage(entry.admission, raw) +_prepare_stage_entry(entry::_StageLoweringEntry{A,W,<:_KeyedReduceStageExecutor}, + raw) where {A,W} = + _prepare_keyed_reduce_stage(entry.admission, raw) _prepare_stage_entry(entry::_StageLoweringEntry{A,W,<:_OrderedFoldStageExecutor}, raw) where {A,W} = _prepare_ordered_fold_stage(entry.admission, raw) _stage_entry_validation(raw, ::_StageLoweringEntry) = raw.validation +_stage_entry_validation(raw::_KeyedReduceStageWorkspace, + ::_StageLoweringEntry) = raw.tree.keyed_reduce.validation _stage_entry_validation(::Nothing, ::_StageLoweringEntry{ A,W,<:_CandidateStageExecutor{<:_DirectIdentityUniqueLayout}}) where {A,W} = nothing @@ -1019,6 +1044,14 @@ end return :invalid_failure_class end +@inline function _stage_keyed_reduce_failure(code::Int32) + code == _KEYED_REDUCE_STATUS_CAPACITY && return :capacity_overflow + code == _KEYED_REDUCE_STATUS_PRIOR_COUNT && return :invalid_prior_count + code == _KEYED_REDUCE_STATUS_DUPLICATE && return :duplicate_key + code == _KEYED_REDUCE_STATUS_INVALID_CONTROL && return :invalid_control + return :invalid_failure_class +end + function _validated_publication_error( status::_ValidatedPublicationStatus{D,H,C}, lease_index::Int ) where {D,H,C<:_StageEntryContext} @@ -1064,6 +1097,8 @@ function _validated_publication_error( failure_class = status.context.executor === :ordered_fold ? fold_failure : status.context.executor === :candidate ? _stage_candidate_failure(code) : status.context.executor === :collect ? _stage_collect_failure(code) : + status.context.executor === :keyed_reduce ? + _stage_keyed_reduce_failure(code) : :invalid_failure_class origin = publication === nothing || !_has_source_origin(publication.origin) ? @@ -1119,6 +1154,10 @@ _execute_stage_program_stage!(prepared::_CollectStagePreparation, parameters::Tuple, lease_index::Int32, predecessors::Tuple, guard, program_validation) = _execute_collect_stage!(prepared, parameters, lease_index, predecessors, guard, program_validation) +_execute_stage_program_stage!(prepared::_KeyedReduceStagePreparation, + parameters::Tuple, lease_index::Int32, predecessors::Tuple, guard, + program_validation) = _execute_keyed_reduce_stage!(prepared, parameters, + lease_index, predecessors, guard, program_validation) _execute_stage_program_stage!(prepared::_OrderedFoldStagePreparation, parameters::Tuple, lease_index::Int32, predecessors::Tuple, guard, program_validation) = _execute_ordered_fold_stage!(prepared, parameters, @@ -1317,6 +1356,17 @@ function _stage_publication_callable_admissions(stage, _stage_order_callable_admissions( law.order, T, :collect, analysis_cache)...) end +function _stage_publication_callable_admissions(stage, + publication::_PreparedStagePublication{C,<:_PreparedKeyedReduceLaw{K,V}}, + analysis_cache::Dict{Any,Any}, + ) where {C,K,V} + signature = Tuple{V,V} + analysis = _cached_closed_callable_effect_analysis!(analysis_cache, + :closed_keyed_reduce, publication.law.operation, signature, + method_signature -> length(method_signature) == 3) + return (_stage_callable_admission(publication.law.operation, signature, + :keyed_reduce_operation, :closed_keyed_reduce, analysis),) +end function _stage_publication_callable_admissions(stage, publication::_PreparedStagePublication{ C,<:_PreparedOrderedFoldLaw{T}}, diff --git a/src/stage_model.jl b/src/stage_model.jl index 5334992..b5bc264 100644 --- a/src/stage_model.jl +++ b/src/stage_model.jl @@ -854,6 +854,41 @@ struct GroupedCollectedValue{K, T} end end +"""One exact sparse key and its current reduced value.""" +struct KeyedValue{K,V} + key::K + value::V + function KeyedValue(::_ConstructionToken, key::K, value::V) where {K,V} + return new{K,V}(key, value) + end +end +function KeyedValue(key::K, value::V) where {K,V} + _qualified_rank_shape(K) || throw(LocalMathValidationError( + "a keyed value key must be Int32, UInt32, or a bounded flat tuple"; + stage = :construct, contract = :keyed_reduce_key_type, + expected = :bounded_total_key, actual = K)) + _storage_value_type(V) || throw(LocalMathValidationError( + "a keyed value requires an admitted storage value type"; + stage = :construct, contract = :keyed_reduce_value_type, actual = V)) + return KeyedValue(_CONSTRUCTION_TOKEN, key, value) +end +"""One optional sparse-key contribution emitted by a `KeyedReduce`.""" +struct KeyedContribution{K,V} + key::K + value::V + participates::Bool + function KeyedContribution(key::K, value::V, participates::Bool = true) where {K,V} + _qualified_rank_shape(K) || throw(LocalMathValidationError( + "a keyed contribution key must be Int32, UInt32, or a bounded flat tuple"; + stage = :construct, contract = :keyed_reduce_key_type, + expected = :bounded_total_key, actual = K)) + _storage_value_type(V) || throw(LocalMathValidationError( + "a keyed contribution requires an admitted storage value type"; + stage = :construct, contract = :keyed_reduce_value_type, actual = V)) + return new{K,V}(key, value, participates) + end +end + """`FoldValue(value, participates=true)` supplies one ordered recurrence value.""" struct FoldValue{T} value::T @@ -957,6 +992,82 @@ function Collect( ) end +"""Left-associated seed then participating candidates in canonical source/lane order.""" +struct CanonicalLeftFold end + +"""Seed a key absent from the stage-entry state with an exact identity.""" +struct NewKeyIdentity{V} + value::V + function NewKeyIdentity(value::V) where {V} + _storage_value_type(V) || throw(LocalMathValidationError( + "a keyed reduction identity requires an admitted storage value type"; + stage = :construct, contract = :keyed_reduce_identity_type, + actual = V)) + return new{V}(value) + end +end + +"""Retain keys whose reduced value equals the new-key identity.""" +struct RetainAllKeys end +"""Remove keys whose reduced value equals the new-key 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. +""" +struct KeyedReduce{K,V,W,F,S,R} + operation::F + seed::S + retention::R + function KeyedReduce(seal::_StageModelSeal, ::Type{K}, ::Type{V}, + ::Val{W}, operation::F, seed::S, + retention::R) where {K,V,W,F,S,R} + seal === _STAGE_MODEL_SEAL || error("invalid stage-model seal") + _qualified_rank_shape(K) || throw(LocalMathValidationError( + "KeyedReduce keys must be Int32, UInt32, or a bounded flat tuple"; + stage = :construct, contract = :keyed_reduce_key_type, + expected = :bounded_total_key, actual = K)) + _storage_value_type(V) || throw(LocalMathValidationError( + "KeyedReduce values require an admitted storage value type"; + stage = :construct, contract = :keyed_reduce_value_type, actual = V)) + 1 <= W <= 32 || throw(LocalMathValidationError( + "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))) + retention isa Union{RetainAllKeys,DropIdentityKeys} || throw( + LocalMathValidationError( + "KeyedReduce requires an explicit key-retention law"; + stage = :construct, contract = :keyed_reduce_retention, + actual = R)) + _device_law_callable(operation) || throw(LocalMathValidationError( + "KeyedReduce operation must be one concrete device-admissible callable"; + stage = :construct, contract = :keyed_reduce_operation, + actual = _device_callable_rejection(operation; + path = (:keyed_reduce, :operation)))) + return new{K,V,W,F,S,R}(operation, seed, retention) + end +end + +function KeyedReduce(::Type{K}, ::Type{V}, operation; + maximum::Integer = 1, seed, + retention = DropIdentityKeys()) where {K,V} + maximum isa Bool && throw(LocalMathValidationError( + "KeyedReduce emission width must be an integer"; + stage = :construct, contract = :keyed_reduce_emission_width, + actual = maximum)) + return KeyedReduce(_STAGE_MODEL_SEAL, K, V, Val(Int(maximum)), operation, + seed, retention) +end + """`OrderedFold(T, state, transition; order=source_order())` declares a finite ordered recurrence.""" struct OrderedFold{T, A, F, O} @@ -1198,8 +1309,6 @@ struct IdentitySeed{T} end """`ExistingSeed()` initializes a reduction from the stage-entry destination value.""" struct ExistingSeed end -"""Left-associated seed then participating candidates in (source item, Relation lane) order.""" -struct CanonicalLeftFold end """Explicitly permits the planner's centrally qualified atomic reassociation.""" struct RelaxedAtomic end @@ -1480,11 +1589,13 @@ _publication_value_type(::Unique{T}) where {T} = T _publication_value_type(::Reduce{T}) where {T} = T _publication_value_type(::Resolve{R, I, T}) where {R, I, T} = T _publication_value_type(::Collect{T}) where {T} = T +_publication_value_type(::KeyedReduce{K,V}) where {K,V} = V _publication_value_type(::OrderedFold{T}) where {T} = T _publication_width(::Unique{T, K}) where {T, K} = K _publication_width(::Reduce{T, K}) where {T, K} = K _publication_width(::Resolve{R, I, T, K}) where {R, I, T, K} = K _publication_width(::Collect{T, K}) where {T, K} = K +_publication_width(::KeyedReduce{K,V,W}) where {K,V,W} = W _publication_width(::OrderedFold) = 1 _unique_relation_admitted( @@ -1644,6 +1755,22 @@ function _validate_publication(components::Tuple, law::Collect) return nothing end +function _validate_publication(components::Tuple, law::KeyedReduce{K,V}) where {K,V} + length(components) == 1 && + only(components) isa CollectionPublication && + only(components).role isa PublicationValue || throw( + LocalMathValidationError( + "KeyedReduce owns exactly one evaluator-fed Collection component"; + stage = :construct, contract = :keyed_reduce_components)) + eltype(only(components).collection) === KeyedValue{K,V} || throw( + LocalMathValidationError( + "KeyedReduce key/value types must equal its Collection element type"; + stage = :construct, contract = :keyed_reduce_component_type, + expected = KeyedValue{K,V}, + actual = eltype(only(components).collection))) + return nothing +end + function _validate_publication(components::Tuple, law::OrderedFold) length(components) == 1 && only(components) isa FoldPublication && @@ -1801,19 +1928,33 @@ function _validate_stage_publication_domain( ) ) end - if publication.law isa Collect + if publication.law isa Union{Collect,KeyedReduce} width = _publication_width(publication.law) - length(source) <= div(Int(typemax(Int32) - 1), width) || throw( + prior = publication.law isa KeyedReduce ? + Int(only(publication.components).collection.capacity) : 0 + length(source) <= div(Int(typemax(Int32) - 1) - prior, width) || throw( LocalMathValidationError( - "Collect source/lane ordinals must fit below the reserved Int32 terminal"; - stage = :construct, contract = :collect_candidate_ordinal, - expected = :nonterminal_int32, actual = (length(source), width), + "Collection source/lane ordinals must fit below the reserved Int32 terminal"; + stage = :construct, contract = :collection_candidate_ordinal, + expected = :nonterminal_int32, + actual = (prior, length(source), width), ) ) end return nothing end +function _validate_keyed_reduce_stage_boundary(publications::Tuple) + position = findfirst(publication -> publication.law isa KeyedReduce, + publications) + position === nothing && return nothing + length(publications) == 1 || throw(LocalMathValidationError( + "KeyedReduce must be the sole publication of its Stage"; + stage = :construct, contract = :keyed_reduce_terminal_publication, + expected = 1, actual = length(publications))) + return nothing +end + function _validate_ordered_fold_stage_boundary( publications::Tuple, accesses::NamedTuple, control::Control @@ -1935,6 +2076,12 @@ function _collect_lane_type_valid(lane::Type, publication::Publication) lane <: CollectedValue || return false return lane.parameters[1] === _publication_value_type(publication.law) end +function _keyed_reduce_lane_type_valid(lane::Type, publication::Publication) + law = publication.law + lane <: KeyedContribution || return false + return lane.parameters[1] === typeof(law).parameters[1] && + lane.parameters[2] === typeof(law).parameters[2] +end function _ordered_fold_lane_type_valid(lane::Type, publication::Publication) lane <: FoldValue || return false return lane.parameters[1] === _publication_value_type(publication.law) @@ -1948,6 +2095,8 @@ _publication_lane_type_valid(lane::Type, publication::Publication{C, <:Resolve}) _resolve_lane_type_valid(lane, publication) _publication_lane_type_valid(lane::Type, publication::Publication{C, <:Collect}) where {C} = _collect_lane_type_valid(lane, publication) +_publication_lane_type_valid(lane::Type, publication::Publication{C, <:KeyedReduce}) where {C} = + _keyed_reduce_lane_type_valid(lane, publication) _publication_lane_type_valid(lane::Type, publication::Publication{C, <:OrderedFold}) where {C} = _ordered_fold_lane_type_valid(lane, publication) @@ -2121,6 +2270,7 @@ struct Stage{S <: Space, A, P, E <: Evaluator, C <: Control, O} _validate_stage_publication_fields(publications) _validate_stage_collection_uniqueness(publications) _validate_ordered_fold_stage_boundary(publications, accesses, control) + _validate_keyed_reduce_stage_boundary(publications) labels = _evaluator_port_names(publications) length(unique(labels)) == length(labels) || throw( LocalMathValidationError( diff --git a/src/stage_planning.jl b/src/stage_planning.jl index 8907efe..d8f9a42 100644 --- a/src/stage_planning.jl +++ b/src/stage_planning.jl @@ -217,7 +217,7 @@ function _project_publication( layout::_StageFieldLayout, publication::Publication, ) law = publication.law - if law isa Collect + if law isa Union{Collect,KeyedReduce} component = only(publication.components) return (_ProjectedCollectionUse( _local_collection_slot(layout, component.collection), @@ -370,9 +370,9 @@ function _control_field_dependency( ) end -function _stage_collect_publication(stage::Stage, collection::Collection) +function _stage_collection_publication(stage::Stage, collection::Collection) for publication in stage.publications - publication.law isa Collect || continue + publication.law isa Union{Collect,KeyedReduce} || continue component = only(publication.components) semantic_identity(component.collection) == semantic_identity(collection) || continue @@ -389,7 +389,7 @@ Base.@nospecializeinfer Base.@noinline function _nearest_preceding_collection_pu stages::Tuple, index::Int, collection::Collection) Base.@nospecialize stages for prior in (index - 1):-1:1 - publication = _stage_collect_publication(stages[prior], collection) + publication = _stage_collection_publication(stages[prior], collection) publication === nothing || return prior, publication end return nothing @@ -398,6 +398,10 @@ end function _validate_collection_access_law( access::CollectionAccess, publication::Publication, source::Space) producer_law = publication.law + producer_law isa KeyedReduce && throw(LocalMathValidationError( + "KeyedReduce state exposes its count but not dense-group or source-position access"; + stage = :plan, contract = :keyed_reduce_collection_access, + expected = :collection_count_only, actual = typeof(access.law))) if access.law isa _BoundedGroup _is_grouped(producer_law.groups) || throw(LocalMathValidationError( "a bounded-group Collection access requires a densely grouped producer"; @@ -429,7 +433,7 @@ function _resolve_collection_access_law( bound::_BoundLaw, access::CollectionAccess, dependency::_PrecedingCollectionDependency) access.law isa _SourcePositionsAccess || return access.law - publication = _stage_collect_publication( + publication = _stage_collection_publication( bound.law.stages[dependency.stage], access.collection) publication === nothing && error("validated Collection producer is missing") width = _publication_width(publication.law) @@ -462,9 +466,10 @@ Base.@nospecializeinfer Base.@noinline function _collection_dependency(bound::_B found = _nearest_preceding_collection_publication( bound.law.stages, index, collection) found === nothing && throw(LocalMathValidationError( - "a Collection consumer requires a preceding Collect publication in the same LocalLaw"; + "a Collection consumer requires a preceding Collection publication in the same LocalLaw"; stage = :plan, contract = :collection_producer, - expected = :preceding_collect, actual = semantic_identity(collection))) + expected = :preceding_collection_publication, + actual = semantic_identity(collection))) prior, publication = found access === nothing || _validate_collection_access_law( access, publication, bound.law.stages[index].source) diff --git a/test/fixtures/keyed_reduce_contracts.jl b/test/fixtures/keyed_reduce_contracts.jl new file mode 100644 index 0000000..53ee355 --- /dev/null +++ b/test/fixtures/keyed_reduce_contracts.jl @@ -0,0 +1,209 @@ +struct KeyedReduceContractNode end +struct KeyedReduceContractEvaluator end +@inline function (::KeyedReduceContractEvaluator)(item::Int32, reads, parameters) + direction = getfield(parameters, 1) + key = (UInt32(isodd(item) ? 1 : 2), UInt32(11), UInt32(3)) + return (; delta = LocalMath.KeyedContribution(key, direction)) +end + +struct EmptyKeyedReduceEvaluator end +@inline (::EmptyKeyedReduceEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution(Int32(0), Int32(0), false)) + +struct OrderedLaneKeyedReduceEvaluator end +@inline function (::OrderedLaneKeyedReduceEvaluator)(item::Int32, reads, parameters) + phase = getfield(parameters, 1) + if phase == Int32(0) + return (; delta = ( + LocalMath.KeyedContribution(Int32(7), Int32(9), item == Int32(1)), + LocalMath.KeyedContribution(Int32(7), Int32(0), false), + )) + end + first_value = Int32(2) * item - Int32(1) + return (; delta = ( + LocalMath.KeyedContribution(Int32(7), first_value), + LocalMath.KeyedContribution(Int32(7), first_value + Int32(1)), + )) +end + +struct OrderedLaneDecimalFold end +@inline (::OrderedLaneDecimalFold)(left::Int32, right::Int32) = + Int32(10) * left + right + +struct SignedExtremeKeyedReduceEvaluator end +@inline (::SignedExtremeKeyedReduceEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution( + item == Int32(1) ? typemin(Int32) : typemax(Int32), item)) + +struct OverflowKeyedReduceEvaluator end +@inline (::OverflowKeyedReduceEvaluator)(item::Int32, reads, parameters) = + (; delta = LocalMath.KeyedContribution(Int32(item), Int32(1))) + +function _shared_keyed_program(backend, source_count, key_type, capacity, + evaluator; maximum = 1, operation = +, + retention = LocalMath.DropIdentityKeys(), parameters = (), + control = LocalMath.Control(), storage = nothing) + source = LocalMath.Space(KeyedReduceContractNode, source_count) + records = LocalMath.Collection(LocalMath.KeyedValue{key_type,Int32}, capacity) + stage = LocalMath.Stage(source, NamedTuple(), ( + LocalMath.Publication(records, + LocalMath.KeyedReduce(key_type, Int32, operation; + maximum, seed = LocalMath.NewKeyIdentity(Int32(0)), + retention); value = :delta),), + LocalMath.Evaluator(evaluator, parameters), control, + LocalMath.SourceOrigin(:keyed_reduce_contract, 2)) + schema = isempty(parameters) ? LocalMath.ParameterSchema() : + LocalMath.ParameterSchema(parameters...) + binding = storage === nothing ? LocalMath.Allocate() : storage + prepared = LocalMath.prepare(LocalMath.LocalLaw(stage; parameters = schema), + records => binding; backend) + return prepared, LocalMath.storage(prepared, records) +end + +function _initialize_scalar_keyed_storage!(storage, keys, values, count) + copyto!(getproperty(storage.records, :key), keys) + copyto!(getproperty(storage.records, :value), values) + copyto!(storage.count, Int32[count]) + return storage +end + +function _keyed_failure(execution) + try + wait(execution) + return nothing + catch error + return error + end +end + +function keyed_reduce_contract(backend) + return @testset "keyed reduction CPU/Metal contract" begin + source = LocalMath.Space(KeyedReduceContractNode, 3) + key_type = Tuple{UInt32,UInt32,UInt32} + record_type = LocalMath.KeyedValue{key_type,Int32} + records = LocalMath.Collection(record_type, 4) + direction = LocalMath.Parameter(:direction, Int32; + bounds = (Int32(-1), Int32(1))) + stage = LocalMath.Stage(source, NamedTuple(), ( + LocalMath.Publication(records, + LocalMath.KeyedReduce(key_type, Int32, +; + maximum = 1, + seed = LocalMath.NewKeyIdentity(Int32(0)), + retention = LocalMath.DropIdentityKeys()); + value = :delta),), + LocalMath.Evaluator(KeyedReduceContractEvaluator(), (direction,)), + LocalMath.Control(), + LocalMath.SourceOrigin(:keyed_reduce_contract, 1)) + prepared = LocalMath.prepare(LocalMath.LocalLaw(stage; + parameters = LocalMath.ParameterSchema(direction)), + records => LocalMath.Allocate(); backend) + storage = LocalMath.storage(prepared, records) + + wait(LocalMath.execute!(prepared; + parameters = (; direction = Int32(1)))) + @test Array(storage.count) == Int32[2] + host = collect(LocalMath.Adapt.adapt(Array, storage.records))[1:2] + @test host == [ + LocalMath.KeyedValue((UInt32(1), UInt32(11), UInt32(3)), Int32(2)), + LocalMath.KeyedValue((UInt32(2), UInt32(11), UInt32(3)), Int32(1)), + ] + + wait(LocalMath.execute!(prepared; + parameters = (; direction = Int32(-1)))) + @test Array(storage.count) == Int32[0] + + empty, empty_storage = _shared_keyed_program(backend, 0, Int32, 0, + EmptyKeyedReduceEvaluator()) + wait(LocalMath.execute!(empty)) + @test Array(empty_storage.count) == Int32[0] + @test isempty(LocalMath.Adapt.adapt(Array, empty_storage.records)) + + phase = LocalMath.Parameter(:phase, Int32; + bounds = (Int32(0), Int32(1))) + ordered, ordered_storage = _shared_keyed_program(backend, 2, Int32, 1, + OrderedLaneKeyedReduceEvaluator(); maximum = 2, + operation = OrderedLaneDecimalFold(), + retention = LocalMath.RetainAllKeys(), parameters = (phase,)) + wait(LocalMath.execute!(ordered; parameters = (; phase = Int32(0)))) + wait(LocalMath.execute!(ordered; parameters = (; phase = Int32(1)))) + ordered_host = collect(LocalMath.Adapt.adapt( + Array, ordered_storage.records)) + @test Array(ordered_storage.count) == Int32[1] + @test ordered_host[1] == LocalMath.KeyedValue(Int32(7), Int32(91234)) + + extremes, extremes_storage = _shared_keyed_program(backend, 2, Int32, 2, + SignedExtremeKeyedReduceEvaluator(); + retention = LocalMath.RetainAllKeys()) + wait(LocalMath.execute!(extremes)) + extremes_host = collect(LocalMath.Adapt.adapt( + Array, extremes_storage.records))[1:2] + @test extremes_host == [ + LocalMath.KeyedValue(typemin(Int32), Int32(1)), + LocalMath.KeyedValue(typemax(Int32), Int32(2)), + ] + + overflow_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 1), + Int32[9], Int32[4], 1) + overflowing, overflow_storage = _shared_keyed_program(backend, 2, + Int32, 1, OverflowKeyedReduceEvaluator(); + storage = overflow_storage) + overflow_before = collect(LocalMath.Adapt.adapt( + Array, overflow_storage.records)) + failure = _keyed_failure(LocalMath.execute!(overflowing)) + @test failure isa LocalMath.LocalMathValidationError + @test failure.actual.failure_class === :capacity_overflow + @test Array(overflow_storage.count) == Int32[1] + @test collect(LocalMath.Adapt.adapt( + Array, overflow_storage.records)) == overflow_before + + duplicate_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 2), + Int32[4, 4], Int32[1, 2], 2) + duplicate, duplicate_storage = _shared_keyed_program(backend, 0, + Int32, 2, EmptyKeyedReduceEvaluator(); storage = duplicate_storage) + duplicate_before = collect(LocalMath.Adapt.adapt( + Array, duplicate_storage.records)) + failure = _keyed_failure(LocalMath.execute!(duplicate)) + @test failure isa LocalMath.LocalMathValidationError + @test failure.actual.failure_class === :duplicate_key + @test Array(duplicate_storage.count) == Int32[2] + @test collect(LocalMath.Adapt.adapt(Array, + duplicate_storage.records)) == duplicate_before + + count_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 1), + Int32[8], Int32[5], -1) + invalid_count, count_storage = _shared_keyed_program(backend, 0, + Int32, 1, EmptyKeyedReduceEvaluator(); storage = count_storage) + count_before = collect(LocalMath.Adapt.adapt(Array, count_storage.records)) + failure = _keyed_failure(LocalMath.execute!(invalid_count)) + @test failure isa LocalMath.LocalMathValidationError + @test failure.actual.failure_class === :invalid_prior_count + @test Array(count_storage.count) == Int32[-1] + @test collect(LocalMath.Adapt.adapt( + Array, count_storage.records)) == count_before + + limit = LocalMath.Parameter(:limit, Int32; + bounds = (Int32(0), Int32(2))) + control_storage = _initialize_scalar_keyed_storage!( + LocalMath.CompactedStorage(backend, + LocalMath.KeyedValue{Int32,Int32}, 1), + Int32[6], Int32[7], 1) + invalid_control, control_storage = _shared_keyed_program(backend, 1, + Int32, 1, OverflowKeyedReduceEvaluator(); parameters = (limit,), + control = LocalMath.Control(prefix = limit), storage = control_storage) + control_before = collect(LocalMath.Adapt.adapt( + Array, control_storage.records)) + failure = _keyed_failure(LocalMath.execute!(invalid_control; + parameters = (; limit = Int32(2)))) + @test failure isa LocalMath.LocalMathValidationError + @test failure.actual.failure_class === :invalid_control + @test Array(control_storage.count) == Int32[1] + @test collect(LocalMath.Adapt.adapt( + Array, control_storage.records)) == control_before + end +end diff --git a/test/metal/keyed_reduce.jl b/test/metal/keyed_reduce.jl new file mode 100644 index 0000000..63bbb00 --- /dev/null +++ b/test/metal/keyed_reduce.jl @@ -0,0 +1,8 @@ +using Test +import Metal +import LocalMath +include("../fixtures/keyed_reduce_contracts.jl") + +Metal.functional() || error("keyed reduction checks require real Metal") +Metal.allowscalar(false) +keyed_reduce_contract(Metal.MetalBackend()) diff --git a/test/metal/runtests.jl b/test/metal/runtests.jl index 4b11292..8313a57 100644 --- a/test/metal/runtests.jl +++ b/test/metal/runtests.jl @@ -19,6 +19,7 @@ const LOCALMATH_METAL_WITNESSES = ( "reduction_control.jl", "empty_pointwise_domains.jl", "collect_canonical_order.jl", + "keyed_reduce.jl", "trigonometric_stages.jl", ) diff --git a/test/runtests.jl b/test/runtests.jl index 8f46fd8..bf275f1 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -24,6 +24,7 @@ const LOCALMATH_INCLUDED_TESTS = ( "test_collect_stage_model.jl", "test_collect_stage_execution.jl", "test_collect_canonical_order.jl", + "test_keyed_reduce_stage.jl", "test_ordered_fold_stage_model.jl", "test_ordered_fold_stage_execution.jl", "test_ordered_fold_step_validation.jl", diff --git a/test/test_keyed_reduce_stage.jl b/test/test_keyed_reduce_stage.jl new file mode 100644 index 0000000..db00011 --- /dev/null +++ b/test/test_keyed_reduce_stage.jl @@ -0,0 +1,142 @@ +using Test +import KernelAbstractions +import LocalMath +const LMKR = LocalMath +include("fixtures/keyed_reduce_contracts.jl") + +keyed_reduce_contract(KernelAbstractions.CPU()) + +struct KeyedReduceNode end +struct KeyedReduceEvaluator end +@inline function (::KeyedReduceEvaluator)(item::Int32, reads, parameters) + remove = getfield(parameters, 1) + key = (UInt32(isodd(item) ? 1 : 2), UInt32(7)) + delta = remove ? Int32(-1) : Int32(1) + return (; delta = LMKR.KeyedContribution(key, delta)) +end + +struct DecimalKeyedFold end +@inline (::DecimalKeyedFold)(left::Int32, right::Int32) = + Int32(10) * left + right +struct DecimalKeyedEvaluator end +@inline (::DecimalKeyedEvaluator)(item::Int32, reads, parameters) = + (; delta = LMKR.KeyedContribution((UInt32(1), UInt32(9)), item)) + +struct DistinctKeyedEvaluator end +@inline (::DistinctKeyedEvaluator)(item::Int32, reads, parameters) = + (; delta = LMKR.KeyedContribution((UInt32(item), UInt32(3)), Int32(1))) + +struct InactiveKeyedEvaluator end +@inline (::InactiveKeyedEvaluator)(item::Int32, reads, parameters) = + (; delta = LMKR.KeyedContribution((UInt32(9), UInt32(9)), Int32(0), false)) + +function _keyed_reduce_program(source_count, collection, evaluator; + operation = +, retention = LMKR.DropIdentityKeys()) + source = LMKR.Space(KeyedReduceNode, source_count) + stage = LMKR.Stage(source, NamedTuple(), ( + LMKR.Publication(collection, + LMKR.KeyedReduce(Tuple{UInt32,UInt32}, Int32, operation; + maximum = 1, seed = LMKR.NewKeyIdentity(Int32(0)), retention); + value = :delta),), LMKR.Evaluator(evaluator), LMKR.Control(), + LMKR.SourceOrigin(:keyed_reduce_test, 2)) + return LMKR.LocalLaw(stage) +end + +function _keyed_reduce_fixture(; capacity = 4) + source = LMKR.Space(KeyedReduceNode, 3) + collection = LMKR.Collection( + LMKR.KeyedValue{Tuple{UInt32,UInt32},Int32}, capacity) + remove = LMKR.Parameter(:remove, Bool) + stage = LMKR.Stage(source, NamedTuple(), ( + LMKR.Publication(collection, + LMKR.KeyedReduce(Tuple{UInt32,UInt32}, Int32, +; + maximum = 1, seed = LMKR.NewKeyIdentity(Int32(0)), + retention = LMKR.DropIdentityKeys()); value = :delta),), + LMKR.Evaluator(KeyedReduceEvaluator(), (remove,)), LMKR.Control(), + LMKR.SourceOrigin(:keyed_reduce_test, 1)) + prepared = LMKR.prepare(LMKR.LocalLaw(stage; + parameters = LMKR.ParameterSchema(remove)), + collection => LMKR.Allocate(); backend = KernelAbstractions.CPU()) + return prepared, collection +end + + +@testset "keyed reduction canonical order and failure atomicity" begin + backend = KernelAbstractions.CPU() + record_type = LMKR.KeyedValue{Tuple{UInt32,UInt32},Int32} + + ordered_collection = LMKR.Collection(record_type, 3) + ordered = LMKR.prepare(_keyed_reduce_program(3, ordered_collection, + DecimalKeyedEvaluator(); operation = DecimalKeyedFold(), + retention = LMKR.RetainAllKeys()), + ordered_collection => LMKR.Allocate(); backend) + wait(LMKR.execute!(ordered)) + ordered_store = LMKR.storage(ordered, ordered_collection) + @test collect(LMKR.Adapt.adapt(Array, ordered_store.records))[1].value == 123 + wait(LMKR.execute!(ordered)) + @test collect(LMKR.Adapt.adapt(Array, ordered_store.records))[1].value == 123123 + + overflow_collection = LMKR.Collection(record_type, 1) + overflow_store = LMKR.CompactedStorage(backend, record_type, 1) + overflow_store.records[1] = LMKR.KeyedValue((UInt32(9), UInt32(9)), Int32(4)) + overflow_store.count[1] = Int32(1) + overflowing = LMKR.prepare(_keyed_reduce_program(2, overflow_collection, + DistinctKeyedEvaluator()), overflow_collection => overflow_store; + backend) + before = collect(LMKR.Adapt.adapt(Array, overflow_store.records)) + failure = try + wait(LMKR.execute!(overflowing)) + nothing + catch error + error + end + @test failure isa LMKR.LocalMathValidationError + @test failure.actual.failure_class === :capacity_overflow + @test Array(overflow_store.count) == Int32[1] + @test collect(LMKR.Adapt.adapt(Array, overflow_store.records)) == before + + duplicate_collection = LMKR.Collection(record_type, 2) + duplicate_store = LMKR.CompactedStorage(backend, record_type, 2) + duplicate_store.records[1] = LMKR.KeyedValue((UInt32(4), UInt32(4)), Int32(1)) + duplicate_store.records[2] = LMKR.KeyedValue((UInt32(4), UInt32(4)), Int32(2)) + duplicate_store.count[1] = Int32(2) + duplicate = LMKR.prepare(_keyed_reduce_program(1, duplicate_collection, + InactiveKeyedEvaluator()), duplicate_collection => duplicate_store; + backend) + before = collect(LMKR.Adapt.adapt(Array, duplicate_store.records)) + failure = try + wait(LMKR.execute!(duplicate)) + nothing + catch error + error + end + @test failure isa LMKR.LocalMathValidationError + @test failure.actual.failure_class === :duplicate_key + @test Array(duplicate_store.count) == Int32[2] + @test collect(LMKR.Adapt.adapt(Array, duplicate_store.records)) == before + + duplicate_store.count[1] = Int32(-1) + failure = try + wait(LMKR.execute!(duplicate)) + nothing + catch error + error + end + @test failure isa LMKR.LocalMathValidationError + @test failure.actual.failure_class === :invalid_prior_count + @test Array(duplicate_store.count) == Int32[-1] +end + +@testset "exact sparse keyed reduction" begin + prepared, collection = _keyed_reduce_fixture() + store = LMKR.storage(prepared, collection) + wait(LMKR.execute!(prepared; parameters = (; remove = false))) + @test Array(store.count) == Int32[2] + records = collect(LMKR.Adapt.adapt(Array, store.records))[1:2] + @test records == [ + LMKR.KeyedValue((UInt32(1), UInt32(7)), Int32(2)), + LMKR.KeyedValue((UInt32(2), UInt32(7)), Int32(1)), + ] + wait(LMKR.execute!(prepared; parameters = (; remove = true))) + @test Array(store.count) == Int32[0] +end diff --git a/test/test_public_api.jl b/test/test_public_api.jl index 47f7a07..5be482a 100644 --- a/test/test_public_api.jl +++ b/test/test_public_api.jl @@ -31,7 +31,7 @@ :one_group, :group_by, :source_order, :canonical_by, :persistent_source_position, :CompactedStorage, :BoundedGroupView, - :Unique, :Reduce, :Resolve, :Collect, :OrderedFold, + :Unique, :Reduce, :Resolve, :Collect, :KeyedReduce, :OrderedFold, :TotalCoverage, :PartialCoverage, :UnreachableEmpty, :PreserveEmpty, :FillEmpty, :IdentitySeed, :ExistingSeed, :CanonicalLeftFold, :RelaxedAtomic, :ArgMin, :ArgMax, :CanonicalSourceLaneTie, @@ -44,7 +44,9 @@ :UniqueValue, :ConditionalUniqueValue, :RoutedUniqueValue, :ConditionalRoutedUniqueValue, :Contribution, :RoutedContribution, :ResolutionValue, :RoutedResolutionValue, :CollectedValue, - :GroupedCollectedValue, :FoldValue, + :GroupedCollectedValue, :KeyedValue, :KeyedContribution, :FoldValue, + :NewKeyIdentity, + :RetainAllKeys, :DropIdentityKeys, )) public_qualified = Set(filter( name -> Base.ispublic(LocalMath, name) && @@ -88,6 +90,11 @@ LocalMath.GhostBoundary @test LocalMath.FoldValue(Int32(3)).value == Int32(3) @test LocalMath.Collect(Int32; maximum=1) isa LocalMath.Collect + keyed = LocalMath.KeyedValue(UInt32(2), Int32(3)) + @test (keyed.key, keyed.value) == (UInt32(2), Int32(3)) + @test !applicable(LocalMath.KeyedValue{Int64,Int32}, Int64(2), Int32(3)) + @test_throws LocalMath.LocalMathValidationError LocalMath.KeyedValue( + Int64(2), Int32(3)) collection = LocalMath.Collection(Int32, 2) @test LocalMath.SourcePositionAccess(collection) isa LocalMath.CollectionAccess @test LocalMath.SourcePositionAccess(collection, 2) isa LocalMath.CollectionAccess