Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion src/graph_engine.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

"""
Expand Down
13 changes: 7 additions & 6 deletions src/plugins/meta/meta_engine.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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!(
Expand All @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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}) =
Expand All @@ -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(
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
19 changes: 19 additions & 0 deletions src/resizable_array.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

25 changes: 25 additions & 0 deletions test/resizable_array_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading