From 0a3ec0a445998c23910561f2a957249207e32d32 Mon Sep 17 00:00:00 2001 From: Bagaev Dmitry Date: Mon, 28 Sep 2026 16:27:47 +0200 Subject: [PATCH] perf: two quadratic passes in model creation, and three smaller costs - apply_meta! walked every node of the model for every meta entry, and kept those in the context: quadratic in the model's size with @meta or @algorithm. It walks the context's own factor nodes. - Resolving a constraint on a matrix variable summed the lengths of every slice before an index for each element, quadratic in the variable's size. While constraints are applied, the plugin binds a cache of those prefix sums per array (with_flattened_index_cache, cached_flattened_index). - NodeLabel == compares the counter first, an Int unique per node, before the untyped name. - The factorisation bitsets start as the full joint, built directly, and a node with no factorisation at all skips the partition step. - ConstraintStack keeps its constraints in a Vector read from the top, where DataStructures' Stack allocated a 1024-element block for every model; it gains length and eltype. Tests: cached_flattened_index against flattened_index on a ragged array, with and without the cache; ConstraintStack's order, counts and pops. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01Re8DzeAcZpjeENCE44QJjF --- src/graph_engine.jl | 3 ++- src/plugins/meta/meta_engine.jl | 13 ++++----- .../variational_constraints.jl | 8 +++--- .../variational_constraints_engine.jl | 24 ++++++++++++----- src/resizable_array.jl | 19 +++++++++++++ .../variational_constraints_engine_tests.jl | 27 +++++++++++++++++++ test/resizable_array_tests.jl | 25 +++++++++++++++++ 7 files changed, 102 insertions(+), 17 deletions(-) diff --git a/src/graph_engine.jl b/src/graph_engine.jl index bf8c4c89..562811a2 100644 --- a/src/graph_engine.jl +++ b/src/graph_engine.jl @@ -266,7 +266,8 @@ to_symbol(label::NodeLabel) = to_symbol(label.name, label.global_counter) to_symbol(name::Any, index::Int) = Symbol(string(name, "_", index)) Base.show(io::IO, label::NodeLabel) = print(io, label.name, "_", label.global_counter) -Base.:(==)(label1::NodeLabel, label2::NodeLabel) = label1.name == label2.name && label1.global_counter == label2.global_counter +# The counter first: it is unique per node and an `Int`, where `name` is untyped and compared by a dynamic call. +Base.:(==)(label1::NodeLabel, label2::NodeLabel) = label1.global_counter == label2.global_counter && label1.name == label2.name Base.hash(label::NodeLabel, h::UInt) = hash(label.global_counter, h) """ diff --git a/src/plugins/meta/meta_engine.jl b/src/plugins/meta/meta_engine.jl index f6b3a870..fa2fe60d 100644 --- a/src/plugins/meta/meta_engine.jl +++ b/src/plugins/meta/meta_engine.jl @@ -149,10 +149,13 @@ function apply_meta!( end end +# The context's own factor nodes, those the descriptor's function names: a pass over the context, +# not over every node of the model. +context_factor_nodes(model::Model, context::Context, predicate) = + Iterators.filter(node -> apply(predicate, model, node), values(factor_nodes(context))) + function apply_meta!(model::Model, context::Context, meta::MetaObject{S, T} where {S <: FactorMetaDescriptor{<:Tuple}, T}) - applicable_nodes = Iterators.filter( - node -> node ∈ values(factor_nodes(context)), filter(as_node(fform(getnodedescriptor(meta))), model) - ) + applicable_nodes = context_factor_nodes(model, context, as_node(fform(getnodedescriptor(meta)))) for node in applicable_nodes neighborhood = neighbors(model, node) save = true @@ -168,9 +171,7 @@ function apply_meta!(model::Model, context::Context, meta::MetaObject{S, T} wher end function apply_meta!(model::Model, context::Context, meta::MetaObject{S, T} where {S <: FactorMetaDescriptor{Nothing}, T}) - applicable_nodes = Iterators.filter( - node -> node ∈ values(factor_nodes(context)), filter(as_node(fform(getnodedescriptor(meta))), model) - ) + applicable_nodes = context_factor_nodes(model, context, as_node(fform(getnodedescriptor(meta)))) for node in applicable_nodes save_meta!(model, node, meta) end diff --git a/src/plugins/variational_constraints/variational_constraints.jl b/src/plugins/variational_constraints/variational_constraints.jl index 94d48240..abd4235c 100644 --- a/src/plugins/variational_constraints/variational_constraints.jl +++ b/src/plugins/variational_constraints/variational_constraints.jl @@ -81,7 +81,7 @@ function postprocess_plugin(plugin::VariationalConstraintsPlugin{NoConstraints}, nodedata = model[flabel] nodeproperties = getproperties(nodedata) number_of_neighbours = length(neighbors(nodeproperties)) - setextra!(nodedata, VariationalConstraintsFactorizationBitSetKey, BoundedBitSetTuple(number_of_neighbours)) + setextra!(nodedata, VariationalConstraintsFactorizationBitSetKey, BoundedBitSetTuple(trues(number_of_neighbours, number_of_neighbours))) end apply_constraints!( @@ -97,8 +97,10 @@ function postprocess_plugin(plugin::VariationalConstraintsPlugin, model::Model) nodedata = model[flabel] nodeproperties = getproperties(nodedata) number_of_neighbours = length(neighbors(nodeproperties)) - setextra!(nodedata, VariationalConstraintsFactorizationBitSetKey, BoundedBitSetTuple(number_of_neighbours)) + setextra!(nodedata, VariationalConstraintsFactorizationBitSetKey, BoundedBitSetTuple(trues(number_of_neighbours, number_of_neighbours))) + end + with_flattened_index_cache() do + apply_constraints!(model, GraphPPL.get_principal_submodel(model), plugin.constraints) end - apply_constraints!(model, GraphPPL.get_principal_submodel(model), plugin.constraints) materialize_constraints!(model) end diff --git a/src/plugins/variational_constraints/variational_constraints_engine.jl b/src/plugins/variational_constraints/variational_constraints_engine.jl index f9f532f1..3d0326aa 100644 --- a/src/plugins/variational_constraints/variational_constraints_engine.jl +++ b/src/plugins/variational_constraints/variational_constraints_engine.jl @@ -457,7 +457,7 @@ Base.in( i::NTuple{M, Int} where {M} ) = (getname(properties) == getname(var)) && - (flattened_index(getcontext(var)[getname(var)], i) ∈ index(var)) && + (cached_flattened_index(getcontext(var)[getname(var)], i) ∈ index(var)) && (getcontext(var) == getcontext(nodedata)) Base.in(nodedata::NodeData, properties::VariableNodeProperties, var::ResolvedIndexedVariable{T} where {T <: Nothing}) = @@ -474,7 +474,7 @@ Base.in( i::NTuple{N, Int} where {N} ) = (getname(properties) == getname(var)) && - (flattened_index(getcontext(var)[getname(var)], i) ∈ index(var)) && + (cached_flattened_index(getcontext(var)[getname(var)], i) ∈ index(var)) && (getcontext(var) == getcontext(nodedata)) Base.in( @@ -540,16 +540,18 @@ rhs(constraint::ResolvedFunctionalFormConstraint) = constraint.rhs const ResolvedConstraint = Union{ResolvedFactorizationConstraint, ResolvedFunctionalFormConstraint} +# A `Vector` used as a stack, read from the top as `DataStructures.Stack` iterates: a `Stack` +# allocates a 1024-element block on creation, once per model even without constraints. struct ConstraintStack - constraints::Stack{ResolvedConstraint} + constraints::Vector{ResolvedConstraint} context_counts::Dict{Context, Int} end -constraints(stack::ConstraintStack) = stack.constraints +constraints(stack::ConstraintStack) = Iterators.reverse(stack.constraints) context_counts(stack::ConstraintStack) = stack.context_counts Base.getindex(stack::ConstraintStack, context::Context) = context_counts(stack)[context] -ConstraintStack() = ConstraintStack(Stack{ResolvedConstraint}(), Dict{Context, Int}()) +ConstraintStack() = ConstraintStack(ResolvedConstraint[], Dict{Context, Int}()) function Base.push!(stack::ConstraintStack, constraint::Any, context::Context) push!(stack.constraints, constraint) @@ -566,13 +568,15 @@ function Base.pop!(stack::ConstraintStack, context::Context) return false end context_counts(stack)[context] -= 1 - pop!(constraints(stack)) + pop!(stack.constraints) return true end return false end -Base.iterate(stack::ConstraintStack, state = 1) = iterate(constraints(stack), state) +Base.iterate(stack::ConstraintStack, state...) = iterate(constraints(stack), state...) +Base.length(stack::ConstraintStack) = length(stack.constraints) +Base.eltype(::Type{ConstraintStack}) = ResolvedConstraint function mean_field_constraint!(constraint::BoundedBitSetTuple) fill!(contents(constraint), false) @@ -632,6 +636,12 @@ function materialize_constraints!(model::Model, node_label::NodeLabel, node_data # Factorize out `neighbors` for which `is_factorized` is `true` materialize_is_factorized_neighbors!(constraint_bitset, neighbor_data(properties)) + # No factorisation at all, the common case: one cluster of every interface. + if all(contents(constraint_bitset)) + setextra!(node_data, VariationalConstraintsFactorizationIndicesKey, (collect(1:size(contents(constraint_bitset), 1)),)) + return nothing + end + constraint_set = unique(eachcol(contents(constraint_bitset))) if !is_valid_partition(constraint_set) diff --git a/src/resizable_array.jl b/src/resizable_array.jl index 7866bdbc..5abbfc0b 100644 --- a/src/resizable_array.jl +++ b/src/resizable_array.jl @@ -245,6 +245,25 @@ function __flattened_index(::Val{N}, array::Vector{V}, findex, index...) where { end end +# `flattened_index` sums the lengths of every slice before `index`, so resolving the index of each +# element of an array is quadratic in its size. While constraints are applied the model is complete, +# so the plugin binds a cache of those prefix sums per array (see `with_flattened_index_cache`). +const FLATTENED_INDEX_CACHE_KEY = :graphppl_flattened_index_cache + +with_flattened_index_cache(f) = task_local_storage(f, FLATTENED_INDEX_CACHE_KEY, IdDict{Any, Vector{Int}}()) + +function cached_flattened_index(array::ResizableArray{T, V, N}, index::NTuple{N, Int}) where {T, V, N} + cache = get(task_local_storage(), FLATTENED_INDEX_CACHE_KEY, nothing) + cache === nothing && return flattened_index(array, index) + prefix = get!(cache::IdDict{Any, Vector{Int}}, array) do + counts = map(slice -> __recursive_length(Val(N - 1), slice), array.data) + return pushfirst!(cumsum(counts), 0) + end + findex = first(index) + return prefix[findex] + __flattened_index(Val(N - 1), array.data[findex], Base.tail(index)...) +end +cached_flattened_index(array::ResizableArray{T, V, 1}, index::NTuple{1, Int}) where {T, V} = flattened_index(array, first(index)) + function Base.first(array::ResizableArray{T, V, N}) where {T, V, N} for index in CartesianIndices(size(array)) #TODO improve performance of this function since it uses splatting if isassigned(array, index.I...)::Bool diff --git a/test/plugins/variational_constraints/variational_constraints_engine_tests.jl b/test/plugins/variational_constraints/variational_constraints_engine_tests.jl index f64f37b0..3efaf513 100644 --- a/test/plugins/variational_constraints/variational_constraints_engine_tests.jl +++ b/test/plugins/variational_constraints/variational_constraints_engine_tests.jl @@ -1569,3 +1569,30 @@ end @test occursin(r"q\(x, y\) = q\(x\)q\(y\)", repr(constraint)) @test occursin(r"μ\(x\) ::(.*?)PointMass", repr(constraint)) end + +@testitem "ConstraintStack" begin + import GraphPPL: + ConstraintStack, constraints, Context, ResolvedFunctionalFormConstraint, ResolvedConstraintLHS, ResolvedIndexedVariable, rhs + + context = Context() + other = Context() + constraint(form) = ResolvedFunctionalFormConstraint(ResolvedConstraintLHS((ResolvedIndexedVariable(:x, nothing, context),)), form) + + stack = ConstraintStack() + push!(stack, constraint(:a), context) + push!(stack, constraint(:b), context) + push!(stack, constraint(:c), other) + # read from the top, the constraint pushed last first + @test length(stack) == 3 + @test map(rhs, collect(stack)) == [:c, :b, :a] + @test map(rhs, collect(constraints(stack))) == [:c, :b, :a] + @test stack[context] == 2 && stack[other] == 1 + + # each context pops as many as it pushed, from the top + @test pop!(stack, other) === true + @test map(rhs, collect(stack)) == [:b, :a] + @test pop!(stack, other) === false + @test pop!(stack, context) === true + @test map(rhs, collect(stack)) == [:a] +end + diff --git a/test/resizable_array_tests.jl b/test/resizable_array_tests.jl index 3ee249b3..d02084f0 100644 --- a/test/resizable_array_tests.jl +++ b/test/resizable_array_tests.jl @@ -318,6 +318,31 @@ end @test flattened_index(s, (2, 1, 1)) == 4 end +@testitem "cached_flattened_index agrees with flattened_index" begin + import GraphPPL: ResizableArray, flattened_index, cached_flattened_index, with_flattened_index_cache + + # a ragged array: the slices have different lengths, so the prefix sums differ from a product + s = ResizableArray(Ref, Val(3)) + for (i, j, k) in ((1, 1, 1), (1, 1, 2), (1, 2, 3), (2, 1, 1), (3, 2, 2), (3, 3, 1), (3, 3, 4)) + s[i, j, k] = Ref(i + j + k) + end + indices = [(i, j, k) for i in 1:3, j in 1:3, k in 1:4 if isassigned(s, i, j, k)] + + # without a cache bound, it computes the index as `flattened_index` does + @test all(i -> cached_flattened_index(s, i) == flattened_index(s, i), indices) + # with one, from prefix sums computed once per array + with_flattened_index_cache() do + @test all(i -> cached_flattened_index(s, i) == flattened_index(s, i), indices) + end + + v = ResizableArray(Ref, Val(1)) + v[1] = Ref(1) + v[3] = Ref(3) + with_flattened_index_cache() do + @test cached_flattened_index(v, (3,)) == flattened_index(v, 3) + end +end + @testitem "iterate" begin import GraphPPL: ResizableArray