From 8c751531e21b01642b301509f2b8ad2b7b3b11ca Mon Sep 17 00:00:00 2001 From: Bagaev Dmitry Date: Tue, 15 Sep 2026 15:02:12 +0200 Subject: [PATCH] fix: correct the lowering of mixed positional/keyword broadcast arguments `combine_broadcast_args(args::Vector, kwargs::Vector)` built both halves of the `MixedArguments` wrongly, not just the keyword half as #309 reports: - the positional half spliced the original argument expressions, which refer to the outer, un-broadcast collections rather than the per-element slots; - the keyword half called `NamedTuple{keys}(args)` on the *whole* broadcast closure tuple, whose arity is `length(positional) + length(keywords)`. Inside the broadcasted expression `args` is the varargs tail of the closure, holding positional slices first and keyword values after, so both halves are now sliced out of it by index. They are emitted as a tuple literal and a named-tuple literal, mirroring the non-broadcast `combine_args`, so `proxy_args` wraps each element individually and the halves stay a `Tuple` and a `NamedTuple` as `MixedArguments{A <: Tuple, K <: NamedTuple}` requires. Updates the two tests that pinned the buggy generated shape. Closes #309. Co-Authored-By: Claude Opus 5 (1M context) --- src/model_macro.jl | 15 +++++++++++--- test/graph_construction_tests.jl | 35 ++++++++++++++++++++++++++++++++ test/model_macro_tests.jl | 20 ++++++++++++++---- 3 files changed, 63 insertions(+), 7 deletions(-) diff --git a/src/model_macro.jl b/src/model_macro.jl index 5a1388f8..3316ed4a 100644 --- a/src/model_macro.jl +++ b/src/model_macro.jl @@ -614,6 +614,11 @@ 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 @@ -621,9 +626,13 @@ function combine_broadcast_args(args::Vector, kwargs::Vector) 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 diff --git a/test/graph_construction_tests.jl b/test/graph_construction_tests.jl index eb6dd758..219942c1 100644 --- a/test/graph_construction_tests.jl +++ b/test/graph_construction_tests.jl @@ -2018,3 +2018,38 @@ end return (;) 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 diff --git a/test/model_macro_tests.jl b/test/model_macro_tests.jl index 30d53bd3..4eb90ebc 100644 --- a/test/model_macro_tests.jl +++ b/test/model_macro_tests.jl @@ -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 @@ -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