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
3 changes: 3 additions & 0 deletions docs/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -6,5 +6,8 @@ GraphPPL = "b3f8163a-e979-4e85-b43e-1f63d8c8b42c"
GraphPlot = "a2cc645c-3eea-5389-862e-a155d0052231"
GraphViz = "f526b714-d49f-11e8-06ff-31ed36ee7ee0"

[sources]
GraphPPL = {path = ".."}

[compat]
Documenter = "1.0"
3 changes: 3 additions & 0 deletions docs/src/developers_guide.md
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,9 @@ GraphPPL.keyword_expressions_to_named_tuple
GraphPPL.convert_anonymous_variables
GraphPPL.is_kwargs_expression
GraphPPL.convert_to_kwargs_expression
GraphPPL.split_positional_and_keyword_args
GraphPPL.mixed_kwargs_rhs
GraphPPL.reconstruct_call
GraphPPL.convert_deterministic_statement
GraphPPL.proxy_args
GraphPPL.save_expression_in_tilde
Expand Down
5 changes: 4 additions & 1 deletion src/graph_engine.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2034,8 +2034,11 @@ make_node!(materialize::True, node_type::NodeType, behaviour::NodeBehaviour, mod
GraphPPL.default_parametrization(model, node_type, fform, rhs_interfaces)
)

# A node that gets materialized takes its arguments either all positionally or all by name, since the
# two have to be matched against the node's interfaces. Mixing them is reported here rather than
# further down, where the mismatch would surface as an unreadable dispatch failure.
make_node!(::True, node_type::NodeType, behaviour::NodeBehaviour, model::Model, ctx::Context, options::NodeCreationOptions, fform::F, lhs_interface::Union{NodeLabel, ProxyLabel, VariableRef}, rhs_interfaces::MixedArguments) where {F} = error(
"MixedArguments not supported for rhs_interfaces when node has to be materialized"
lazy"MixedArguments not supported for `$(fform)`: a node that has to be materialized cannot be called with both positional and keyword arguments. Got $(length(rhs_interfaces.args)) positional argument(s) and the keyword argument(s) $(keys(rhs_interfaces.kwargs)). Use either all positional or all keyword arguments."
)

make_node!(materialize::True, node_type::Composite, behaviour::Stochastic, model::Model, ctx::Context, options::NodeCreationOptions, fform::F, lhs_interface::Union{NodeLabel, ProxyLabel, VariableRef}, rhs_interfaces::Tuple{}) where {F} = make_node!(
Expand Down
81 changes: 75 additions & 6 deletions src/model_macro.jl
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,66 @@ end

is_kwargs_expression(e) = false

"""
split_positional_and_keyword_args(args::Vector)

Split a call's argument list into its positional and keyword parts, returning them as a
`(positional, keywords)` tuple.

Julia writes an explicit `; ...` group into a single leading `:parameters` node, but keyword
arguments written in the comma form stay inline as `:kw` nodes among the positional arguments.
`f(a; b = c)` and `f(a, b = c)` are the same call, so both spellings are normalized here to the
same split. Everything downstream assumes the `:parameters` form, so the comma form has to be
rewritten into it before it reaches `combine_args`.
"""
function split_positional_and_keyword_args(args::Vector)
positional = Any[]
keywords = Any[]
for arg in args
if arg isa Expr && arg.head === :parameters
append!(keywords, arg.args)
elseif arg isa Expr && arg.head === :kw
push!(keywords, arg)
else
push!(positional, arg)
end
end
return positional, keywords
end

"""
mixed_kwargs_rhs(f, args::Vector, options::Vector)

Rebuild the right-hand side `f(a, b = c) where { ... }` in the canonical form
`f(a; b = c) where { ... }`, or return `nothing` if `args` holds no such mix and the expression
should be left alone.

Returns only the right-hand side, because the three operators this serves do not share an
expression head: `~` and `.~` are `:call`s, while `:=` has its own head.
"""
function mixed_kwargs_rhs(f, args::Vector, options::Vector)
positional, keywords = split_positional_and_keyword_args(args)
if isempty(positional) || isempty(keywords)
return nothing
end
return :($f($(positional...); $(keywords...)) where {$(options...)})
end

"""
reconstruct_call(f, args::Vector)

Rebuild the call `f(args...)` with its keyword arguments in the canonical `:parameters` form.

`convert_anonymous_variables` runs *after* `convert_to_kwargs_expression` in the pipeline, so the
tilde expressions it generates for nested calls never pass through that normalization. Splicing
the captured arguments back verbatim would reintroduce the comma form there, so anything that
generates a new call expression builds it through this function instead.
"""
function reconstruct_call(f, args::Vector)
positional, keywords = split_positional_and_keyword_args(args)
return isempty(keywords) ? :($f($(positional...))) : :($f($(positional...); $(keywords...)))
end

"""
convert_to_kwargs_expression(expr::Expr)

Expand All @@ -289,7 +349,10 @@ function convert_to_kwargs_expression(e::Expr)
if GraphPPL.is_kwargs_expression(args)
return :($lhs ~ $f(; $(args...)) where {$(options...)})
else
return e
# Mixed positional and keyword arguments, e.g. `f(a, b = c)`. Everything downstream
# expects the keywords in a `:parameters` group, so normalize before giving up.
rhs = GraphPPL.mixed_kwargs_rhs(f, args, options)
return isnothing(rhs) ? e : :($lhs ~ $rhs)
end
# Logic for .~ operator
elseif @capture(e, (lhs_ .~ f_(; kwargs__) where {options__}))
Expand All @@ -298,7 +361,10 @@ function convert_to_kwargs_expression(e::Expr)
if GraphPPL.is_kwargs_expression(args)
return :($lhs .~ $f(; $(args...)) where {$(options...)})
else
return e
# Mixed positional and keyword arguments, e.g. `f(a, b = c)`. Everything downstream
# expects the keywords in a `:parameters` group, so normalize before giving up.
rhs = GraphPPL.mixed_kwargs_rhs(f, args, options)
return isnothing(rhs) ? e : :($lhs .~ $rhs)
end
# Logic for := operator
elseif @capture(e, (lhs_ := f_(; kwargs__) where {options__}))
Expand All @@ -307,7 +373,10 @@ function convert_to_kwargs_expression(e::Expr)
if GraphPPL.is_kwargs_expression(args)
return :($lhs := $f(; $(args...)) where {$(options...)})
else
return e
# Mixed positional and keyword arguments, e.g. `f(a, b = c)`. Everything downstream
# expects the keywords in a `:parameters` group, so normalize before giving up.
rhs = GraphPPL.mixed_kwargs_rhs(f, args, options)
return isnothing(rhs) ? e : :($lhs := $rhs)
end
else
return e
Expand All @@ -329,21 +398,21 @@ function convert_to_anonymous(e::Expr, created_by)
f = Symbol(string(f)[2:end])
return quote
let $sym = GraphPPL.create_anonymous_variable!(__model__, __context__)
$sym .~ $f($(args...)) where {anonymous = true, created_by = $created_by}
$sym .~ $(reconstruct_call(f, args)) where {anonymous = true, created_by = $created_by}
end
end
end
sym = gensym(:anon)
return quote
let $sym = GraphPPL.create_anonymous_variable!(__model__, __context__)
$sym ~ $f($(args...)) where {anonymous = true, created_by = $created_by}
$sym ~ $(reconstruct_call(f, args)) where {anonymous = true, created_by = $created_by}
end
end
elseif @capture(e, f_.(args__))
sym = gensym(:anon)
return quote
let $sym = GraphPPL.create_anonymous_variable!(__model__, __context__)
$sym .~ $f($(args...)) where {anonymous = true, created_by = $created_by}
$sym .~ $(reconstruct_call(f, args)) where {anonymous = true, created_by = $created_by}
end
end
end
Expand Down
63 changes: 63 additions & 0 deletions test/graph_construction_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2088,3 +2088,66 @@ end
# The node property `name` is untouched, so `is_anonymous` keeps working
@test length(collect(filter(v -> is_anonymous(getproperties(model[v])), collect(variable_nodes(model))))) === 2
end

@testitem "Mixed positional and keyword arguments in the comma form `f(a, b = c)`" begin
using Distributions
import GraphPPL: create_model, getcontext, factor_nodes, variable_nodes

include("testutils.jl")

using .TestUtils.ModelZoo

# `f(a, b = c)` and `f(a; b = c)` are the same call in Julia, but the comma form used to leave
# the keyword among the positional arguments, so the generated code contained a tuple literal
# with a `:kw` node in it and the model failed to even macro-expand:
# syntax: invalid named tuple element ...
# The contract asserted here is that the two spellings are now indistinguishable.
mixed_det(a; s) = a + s
GraphPPL.NodeBehaviour(::TestUtils.TestGraphPPLBackend, ::typeof(mixed_det)) = GraphPPL.Deterministic()

# A nested deterministic call over constants is evaluated directly, so mixed arguments
# genuinely work here. This is the shape reported in the original issue.
@model function anon_comma()
x ~ NormalMeanVariance(mixed_det(1.0, s = 2.0), 1.0)
end

@model function anon_semicolon()
x ~ NormalMeanVariance(mixed_det(1.0; s = 2.0), 1.0)
end

model_comma = create_model(anon_comma())
model_semicolon = create_model(anon_semicolon())
@test model_comma isa GraphPPL.Model
@test length(collect(factor_nodes(model_comma))) === length(collect(factor_nodes(model_semicolon)))
@test length(collect(variable_nodes(model_comma))) === length(collect(variable_nodes(model_semicolon)))

# A node that has to be materialized still cannot take both, but the comma form now reaches
# that stated limitation instead of failing as invalid syntax, exactly like the semicolon form
@model function materialized_comma()
x ~ NormalMeanVariance(0, var = 1)
end

@model function materialized_semicolon()
x ~ NormalMeanVariance(0; var = 1)
end

@test_throws "MixedArguments not supported" create_model(materialized_comma())
@test_throws "cannot be called with both positional and keyword arguments" create_model(materialized_comma())
@test_throws "MixedArguments not supported" create_model(materialized_semicolon())

# `:=` takes the same path and must agree as well
@model function deterministic_comma()
y ~ NormalMeanVariance(0, 1)
z := mixed_det(y, s = 3.0)
x ~ NormalMeanVariance(z, 1.0)
end

@model function deterministic_semicolon()
y ~ NormalMeanVariance(0, 1)
z := mixed_det(y; s = 3.0)
x ~ NormalMeanVariance(z, 1.0)
end

@test_throws "MixedArguments not supported" create_model(deterministic_comma())
@test_throws "MixedArguments not supported" create_model(deterministic_semicolon())
end
73 changes: 73 additions & 0 deletions test/model_macro_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -757,6 +757,79 @@ end
x ~ Normal(; μ = μ, σ = σ) where {created_by = (x ~ Normal(μ = μ, σ = σ) where {q = MeanField()}), q = MeanField()}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output

# Test 28: mixed positional and keyword arguments written in the comma form are split into
# the keyword form, exactly as if they had been written with an explicit `;`. Previously the
# `var = v` stayed among the positional arguments and `combine_args` emitted a tuple literal
# containing a `:kw` node, which is not valid syntax.
input = quote
x ~ Normal(m, var = v) where {created_by = (x ~ Normal(m, var = v))}
end
output = quote
x ~ Normal(m; var = v) where {created_by = (x ~ Normal(m, var = v))}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output

# Test 29: the same for `.~`
input = quote
x .~ Normal(m, var = v) where {created_by = (x .~ Normal(m, var = v))}
end
output = quote
x .~ Normal(m; var = v) where {created_by = (x .~ Normal(m, var = v))}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output

# Test 30: ... and for `:=`
input = quote
x := f(m, s = v) where {created_by = (x := f(m, s = v))}
end
output = quote
x := f(m; s = v) where {created_by = (x := f(m, s = v))}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output

# Test 31: both spellings at once collapse into a single keyword group. Julia puts the explicit
# `;` group first in the argument list, so the keywords from it lead
input = quote
x ~ Normal(m, var = v; mean = q) where {created_by = (x ~ Normal(m, var = v; mean = q))}
end
output = quote
x ~ Normal(m; mean = q, var = v) where {created_by = (x ~ Normal(m, var = v; mean = q))}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output

# Test 32: a call with several positional arguments and several comma-form keywords
input = quote
x ~ Normal(μ, σ, a = τ, b = θ) where {created_by = (x ~ Normal(μ, σ, a = τ, b = θ))}
end
output = quote
x ~ Normal(μ, σ; a = τ, b = θ) where {created_by = (x ~ Normal(μ, σ, a = τ, b = θ))}
end
@test_expression_generating apply_pipeline(input, convert_to_kwargs_expression) output
end

@testitem "split_positional_and_keyword_args" begin
import GraphPPL: split_positional_and_keyword_args
import MacroTools: @capture

include("testutils.jl")

split_of(s) = (@capture(s, f_(args__)); split_positional_and_keyword_args(args))

# Comma form: the keyword sits inline among the positional arguments as a `:kw` node
@test split_of(:(foo(a, b = c))) == (Any[:a], Any[Expr(:kw, :b, :c)])

# Semicolon form: Julia collects it into a leading `:parameters` node instead. Both spellings
# mean the same call, so both must produce the same split
@test split_of(:(foo(a; b = c))) == (Any[:a], Any[Expr(:kw, :b, :c)])

# Both at once -- the `:parameters` group comes first in the argument list
@test split_of(:(foo(a, b = c; d = e))) == (Any[:a], Any[Expr(:kw, :d, :e), Expr(:kw, :b, :c)])

# Degenerate cases: nothing to split
@test split_of(:(foo(a, b))) == (Any[:a, :b], Any[])
@test split_of(:(foo(a = 1, b = 2))) == (Any[], Any[Expr(:kw, :a, 1), Expr(:kw, :b, 2)])
@test split_of(:(foo())) == (Any[], Any[])
end

@testitem "convert_to_anonymous" begin
Expand Down
Loading