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
15 changes: 12 additions & 3 deletions src/model_macro.jl
Original file line number Diff line number Diff line change
Expand Up @@ -614,16 +614,25 @@ combine_broadcast_args(args::Vector, kwargs::Nothing) = quote
args
end

# Inside the broadcasted expression `args` is the varargs tail of the broadcast closure, holding the
# per-element slice of every combinable argument: the positional ones first, then the keyword values, in
# declaration order. Both halves therefore have to be sliced out of that runtime tuple. Splicing the
# original positional expressions here instead would capture the outer, un-broadcast collections, and
# building the keyword `NamedTuple` out of the whole tuple mismatches its arity.
function combine_broadcast_args(args::Vector, kwargs::Vector)
kwargs_keys = [arg.args[1] for arg in kwargs]
if length(args) == 0
return quote
NamedTuple{$(Tuple(kwargs_keys))}(args)
end
else
return quote
GraphPPL.MixedArguments($(Expr(:tuple, args...)), NamedTuple{$(Tuple(kwargs_keys))}(args))
end
npositional = length(args)
# A tuple literal and a named-tuple literal, mirroring the non-broadcast `combine_args`, so that
# `proxy_args` wraps each element in its own `proxylabel` and the two halves stay a `Tuple` and a
# `NamedTuple` as `MixedArguments{A <: Tuple, K <: NamedTuple}` requires
positional_args = Expr(:tuple, [:(args[$i]) for i in 1:npositional]...)
keyword_args = Expr(:tuple, [Expr(:(=), key, :(args[$(npositional + j)])) for (j, key) in enumerate(kwargs_keys)]...)
return :(GraphPPL.MixedArguments($positional_args, $keyword_args))
end
end

Expand Down
35 changes: 35 additions & 0 deletions test/graph_construction_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2019,6 +2019,41 @@ end
end
end

@testitem "Broadcasting with mixed positional and keyword arguments" begin
using Distributions
import GraphPPL: create_model

include("testutils.jl")

@model function bc_mixed_sub(out, x, y)
out ~ Normal(x, y)
end

@model function bc_mixed_main()
local mu
local sg
for i in 1:5
mu[i] ~ Normal(0, 1)
sg[i] ~ Gamma(1, 1)
end
z .~ bc_mixed_sub(mu; y = sg)
out ~ Normal(z[5], 1)
end

# A broadcast mixing positional and keyword arguments now lowers to a well-formed `MixedArguments`,
# built by slicing the broadcast closure's `args` tuple. Materializing a node from `MixedArguments`
# is still unsupported (same as for the non-broadcast `~`), but the user gets that stated limitation
# instead of an opaque `MethodError: no method matching tuple(...)` from the broken lowering.
@test_throws "MixedArguments not supported" create_model(bc_mixed_main())

# Keyword-only and positional-only broadcasts are unaffected
@model function bc_kwargs_only()
y .~ Normal(fill(0.0, 5), 1.0)
z .~ Normal(mean = y, var = fill(1.0, 5))
end
@test create_model(bc_kwargs_only()) isa GraphPPL.Model
end

@testitem "Multiple anonymous variables in one context should not collapse to a single key" begin
using Distributions
import GraphPPL:
Expand Down
20 changes: 16 additions & 4 deletions test/model_macro_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -1456,8 +1456,14 @@ end
@test_expression_generating combine_broadcast_args([], [Expr(:kw, :μ, :μ), Expr(:kw, :σ, :σ)]) quote
NamedTuple{$(:μ, :σ)}(args)
end
@test_expression_generating combine_broadcast_args([:μ, :σ], [Expr(:kw, :μ, :μ), Expr(:kw, :σ, :σ)]) quote
GraphPPL.MixedArguments((μ, σ), NamedTuple{$(:μ, :σ)}(args))
# Both halves are sliced out of the broadcast closure's `args` tuple: the positional arguments come
# first, the keyword values after them. Splicing the original expressions (`(μ, σ)`) would capture the
# outer, un-broadcast collections instead of the per-element slots.
@test_expression_generating combine_broadcast_args([:μ, :σ], [Expr(:kw, :τ, :τ), Expr(:kw, :θ, :θ)]) quote
GraphPPL.MixedArguments((args[1], args[2]), (τ = args[3], θ = args[4]))
end
@test_expression_generating combine_broadcast_args([:μ], [Expr(:kw, :σ, :σ)]) quote
GraphPPL.MixedArguments((args[1],), (σ = args[2],))
end
end

Expand Down Expand Up @@ -1792,8 +1798,14 @@ end
some_node,
ilhs,
GraphPPL.MixedArguments(
(GraphPPL.proxylabel(:a, a, nothing, GraphPPL.False()), GraphPPL.proxylabel(:b, b, nothing, GraphPPL.False())),
GraphPPL.proxylabel(:anonymous, NamedTuple{$(:μ, :σ)}(args), nothing, GraphPPL.False())
(
GraphPPL.proxylabel(:args, args, (1,), GraphPPL.False()),
GraphPPL.proxylabel(:args, args, (2,), GraphPPL.False())
),
(
μ = GraphPPL.proxylabel(:args, args, (3,), GraphPPL.False()),
σ = GraphPPL.proxylabel(:args, args, (4,), GraphPPL.False())
)
)
)
end
Expand Down
Loading