Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -951,7 +951,7 @@ function apply_constraints!(
applicable_nodes = unroll_nocreate(context[getvariables(marginal_constraint)])
for node in applicable_nodes
if hasextra(model[node], VariationalConstraintsMarginalFormConstraintKey)
@warn lazy"Node $node already has functional form constraint $(opt[:q]) applied, therefore $constraint_data will not be applied"
@warn lazy"Node $node already has functional form constraint $(getextra(model[node], VariationalConstraintsMarginalFormConstraintKey)) applied, therefore $(getconstraint(marginal_constraint)) will not be applied"
else
setextra!(model[node], VariationalConstraintsMarginalFormConstraintKey, getconstraint(marginal_constraint))
end
Expand All @@ -966,7 +966,7 @@ function apply_constraints!(model::Model, context::Context, message_constraint::
applicable_nodes = unroll_nocreate(context[getvariables(message_constraint)])
for node in applicable_nodes
if hasextra(model[node], VariationalConstraintsMessagesFormConstraintKey)
@warn lazy"Node $node already has functional form constraint $(opt[:q]) applied, therefore $constraint_data will not be applied"
@warn lazy"Node $node already has functional form constraint $(getextra(model[node], VariationalConstraintsMessagesFormConstraintKey)) applied, therefore $(getconstraint(message_constraint)) will not be applied"
else
setextra!(model[node], VariationalConstraintsMessagesFormConstraintKey, getconstraint(message_constraint))
end
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -613,6 +613,52 @@ end
end
end


@testitem "Applying a second form constraint warns and preserves the original" begin
import GraphPPL:
create_model,
MarginalFormConstraint,
MessageFormConstraint,
IndexedVariable,
apply_constraints!,
getextra,
VariationalConstraintsMarginalFormConstraintKey,
VariationalConstraintsMessagesFormConstraintKey

include("../../testutils.jl")

using .TestUtils.ModelZoo

struct FirstArbitraryFormConstraint end
struct SecondArbitraryFormConstraint end

# A node that already carries a marginal form constraint must warn (not throw) and keep the first one
model = create_model(simple_model())
context = GraphPPL.getcontext(model)
apply_constraints!(model, context, MarginalFormConstraint(IndexedVariable(:x, nothing), FirstArbitraryFormConstraint()))

@test_logs (:warn, r"already has functional form constraint") match_mode = :any apply_constraints!(
model, context, MarginalFormConstraint(IndexedVariable(:x, nothing), SecondArbitraryFormConstraint())
)

for node in filter(GraphPPL.as_variable(:x), model)
@test getextra(model[node], VariationalConstraintsMarginalFormConstraintKey) == FirstArbitraryFormConstraint()
end

# ... and the same for message form constraints
model = create_model(simple_model())
context = GraphPPL.getcontext(model)
apply_constraints!(model, context, MessageFormConstraint(IndexedVariable(:x, nothing), FirstArbitraryFormConstraint()))

@test_logs (:warn, r"already has functional form constraint") match_mode = :any apply_constraints!(
model, context, MessageFormConstraint(IndexedVariable(:x, nothing), SecondArbitraryFormConstraint())
)

for node in filter(GraphPPL.as_variable(:x), model)
@test getextra(model[node], VariationalConstraintsMessagesFormConstraintKey) == FirstArbitraryFormConstraint()
end
end

@testitem "save constraints with constants via `mean_field_constraint!`" begin
using BitSetTuples
import GraphPPL:
Expand Down
Loading