From d44527fd07364e82be936cb2c2440485f05da4b4 Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Thu, 18 Jun 2026 13:10:41 +0100 Subject: [PATCH 1/4] Support Mooncake 0.6 forward mode in find_alpha rules Mooncake 0.6 reworks batched forward AD around `Lifted`/`NDual` and routes `prepare_derivative_cache` (used by `AutoMooncakeForward`) through `build_frule`. Add `frule!!(::Lifted, ...)` rules for `find_alpha` (float and integer third argument), gated behind `pkgversion(Mooncake) >= v"0.6"`, keeping the existing `Mooncake.Dual` rules for 0.5. Under 0.6 these `Lifted` rules intercept `find_alpha` as a primitive, so the batched pass never enters its body. Under 0.5 the batched (Nfwd) pass instead propagates `NDual`s straight through `find_alpha`, tripping on `ceil(::Int, ::NDual)` inside the root-finder; the `find_alpha(::NDual, ...)` short-circuit methods fix that and are therefore gated to `< v"0.6"` (verified redundant on 0.6, required on 0.5). This un-breaks the `AutoMooncakeForward` PlanarLayer-inverse integration case, so drop its `@test_broken` marker. Bump Mooncake compat to allow 0.6. Co-Authored-By: Claude Opus 4.8 (1M context) --- Project.toml | 2 +- ext/BijectorsMooncakeExt.jl | 111 +++++++++++++++---- test/integration_tests/mooncake/Project.toml | 2 +- test/integration_tests/mooncake/main.jl | 8 +- 4 files changed, 90 insertions(+), 33 deletions(-) diff --git a/Project.toml b/Project.toml index 74af0693..dadab71d 100644 --- a/Project.toml +++ b/Project.toml @@ -59,7 +59,7 @@ IrrationalConstants = "0.1, 0.2" LazyArrays = "2" LogExpFunctions = "0.3.3" MappedArrays = "0.2.2, 0.3, 0.4" -Mooncake = "0.4.95, 0.5" +Mooncake = "0.4.95, 0.5, 0.6" Reexport = "0.2, 1" ReverseDiff = "1" Roots = "1.3.15, 2, 3" diff --git a/ext/BijectorsMooncakeExt.jl b/ext/BijectorsMooncakeExt.jl index 503f077c..8b9bebde 100644 --- a/ext/BijectorsMooncakeExt.jl +++ b/ext/BijectorsMooncakeExt.jl @@ -1,7 +1,13 @@ module BijectorsMooncakeExt using Mooncake: @is_primitive, MinimalCtx, Mooncake, CoDual, primal, tangent_type -using Bijectors: find_alpha +import Bijectors: find_alpha + +@static if pkgversion(Mooncake) >= v"0.6" + using Mooncake: Lifted +else + using Mooncake.Nfwd: NDual +end # Closed-form partials of `find_alpha` derived from differentiating # wt_y == α + wt_u_hat * tanh(α + b) @@ -16,20 +22,61 @@ function _find_alpha_partials(wt_y::P, wt_u_hat::P, b::P) where {P<:Base.IEEEFlo return α, x, -tanh(α + b) * x, x - one(P) end +# Mooncake < 0.6 runs the batched forward (Nfwd) pass by propagating `NDual`s straight +# through `find_alpha`'s body, which trips on `ceil(::Int, ::NDual)` inside the root-finder. +# Short-circuit with the closed-form partials. From 0.6 the batched pass dispatches through +# the `frule!!(::Lifted, ...)` rules below instead, so these methods are unnecessary. +@static if pkgversion(Mooncake) < v"0.6" + function find_alpha( + wt_y::NDual{P,N}, wt_u_hat::NDual{P,N}, b::NDual{P,N} + ) where {P<:Base.IEEEFloat,N} + α, ∂y, ∂u, ∂b = _find_alpha_partials(wt_y.value, wt_u_hat.value, b.value) + dα = ntuple(Val(N)) do lane + ∂y * wt_y.partials[lane] + ∂u * wt_u_hat.partials[lane] + ∂b * b.partials[lane] + end + return NDual{P,N}(α, dα) + end + + function find_alpha( + wt_y::NDual{P,N}, wt_u_hat::NDual{P,N}, b::Real + ) where {P<:Base.IEEEFloat,N} + α, ∂y, ∂u, _ = _find_alpha_partials(wt_y.value, wt_u_hat.value, P(b)) + dα = ntuple(lane -> ∂y * wt_y.partials[lane] + ∂u * wt_u_hat.partials[lane], Val(N)) + return NDual{P,N}(α, dα) + end +end + # Floating-point third argument. @is_primitive(MinimalCtx, Tuple{typeof(find_alpha),P,P,P} where {P<:Base.IEEEFloat},) -function Mooncake.frule!!( - ::Mooncake.Dual{typeof(find_alpha)}, - x::Mooncake.Dual{P}, - y::Mooncake.Dual{P}, - z::Mooncake.Dual{P}, -) where {P<:Base.IEEEFloat} - α, ∂y, ∂u, ∂b = _find_alpha_partials( - Mooncake.primal(x), Mooncake.primal(y), Mooncake.primal(z) - ) - dα = ∂y * Mooncake.tangent(x) + ∂u * Mooncake.tangent(y) + ∂b * Mooncake.tangent(z) - return Mooncake.Dual(α, dα) +@static if pkgversion(Mooncake) >= v"0.6" + function Mooncake.frule!!( + ::Lifted{typeof(find_alpha),N}, + x::Lifted{P,N,Mooncake.NDual{P,N}}, + y::Lifted{P,N,Mooncake.NDual{P,N}}, + z::Lifted{P,N,Mooncake.NDual{P,N}}, + ) where {N,P<:Base.IEEEFloat} + α, ∂y, ∂u, ∂b = _find_alpha_partials(primal(x), primal(y), primal(z)) + dα = ntuple(Val(N)) do lane + ∂y * Mooncake.tangent(x).partials[lane] + + ∂u * Mooncake.tangent(y).partials[lane] + + ∂b * Mooncake.tangent(z).partials[lane] + end + return Lifted{P,N}(α, Mooncake.NDual{P,N}(α, dα)) + end +else + function Mooncake.frule!!( + ::Mooncake.Dual{typeof(find_alpha)}, + x::Mooncake.Dual{P}, + y::Mooncake.Dual{P}, + z::Mooncake.Dual{P}, + ) where {P<:Base.IEEEFloat} + α, ∂y, ∂u, ∂b = _find_alpha_partials( + Mooncake.primal(x), Mooncake.primal(y), Mooncake.primal(z) + ) + dα = ∂y * Mooncake.tangent(x) + ∂u * Mooncake.tangent(y) + ∂b * Mooncake.tangent(z) + return Mooncake.Dual(α, dα) + end end function Mooncake.rrule!!( @@ -55,18 +102,34 @@ function _assert_integer_nontangent(::Type{I}) where {I<:Integer} end end -function Mooncake.frule!!( - ::Mooncake.Dual{typeof(find_alpha)}, - x::Mooncake.Dual{P}, - y::Mooncake.Dual{P}, - z::Mooncake.Dual{I}, -) where {P<:Base.IEEEFloat,I<:Integer} - _assert_integer_nontangent(I) - α, ∂y, ∂u, _ = _find_alpha_partials( - Mooncake.primal(x), Mooncake.primal(y), P(Mooncake.primal(z)) - ) - dα = ∂y * Mooncake.tangent(x) + ∂u * Mooncake.tangent(y) - return Mooncake.Dual(α, dα) +@static if pkgversion(Mooncake) >= v"0.6" + function Mooncake.frule!!( + ::Lifted{typeof(find_alpha),N}, + x::Lifted{P,N,Mooncake.NDual{P,N}}, + y::Lifted{P,N,Mooncake.NDual{P,N}}, + z::Lifted{I,N,Mooncake.NoDual}, + ) where {N,P<:Base.IEEEFloat,I<:Integer} + _assert_integer_nontangent(I) + α, ∂y, ∂u, _ = _find_alpha_partials(primal(x), primal(y), P(primal(z))) + dα = ntuple(Val(N)) do lane + ∂y * Mooncake.tangent(x).partials[lane] + ∂u * Mooncake.tangent(y).partials[lane] + end + return Lifted{P,N}(α, Mooncake.NDual{P,N}(α, dα)) + end +else + function Mooncake.frule!!( + ::Mooncake.Dual{typeof(find_alpha)}, + x::Mooncake.Dual{P}, + y::Mooncake.Dual{P}, + z::Mooncake.Dual{I}, + ) where {P<:Base.IEEEFloat,I<:Integer} + _assert_integer_nontangent(I) + α, ∂y, ∂u, _ = _find_alpha_partials( + Mooncake.primal(x), Mooncake.primal(y), P(Mooncake.primal(z)) + ) + dα = ∂y * Mooncake.tangent(x) + ∂u * Mooncake.tangent(y) + return Mooncake.Dual(α, dα) + end end function Mooncake.rrule!!( diff --git a/test/integration_tests/mooncake/Project.toml b/test/integration_tests/mooncake/Project.toml index c036282e..b9a4238c 100644 --- a/test/integration_tests/mooncake/Project.toml +++ b/test/integration_tests/mooncake/Project.toml @@ -19,4 +19,4 @@ Bijectors = {path = "../../.."} [compat] AbstractPPL = "0.15" DifferentiationInterface = "0.6, 0.7" -Mooncake = "0.4, 0.5" +Mooncake = "0.4, 0.5, 0.6" diff --git a/test/integration_tests/mooncake/main.jl b/test/integration_tests/mooncake/main.jl index 88df74ee..d7347ef4 100644 --- a/test/integration_tests/mooncake/main.jl +++ b/test/integration_tests/mooncake/main.jl @@ -59,13 +59,7 @@ end @testset "Mooncake bijector AD" begin for c in generate_ad_testcases(), adtype in adtypes - # AbstractPPL's Mooncake extension routes `AutoMooncakeForward` through - # `Mooncake.prepare_derivative_cache` (batched NDual mode). `BijectorsMooncakeExt` - # only has `find_alpha` rules for `Mooncake.Dual`, so the PlanarLayer inverse path - # (which calls `find_alpha`) trips on `ceil(::Int, ::NDual)` inside Roots.ITP. - is_broken = - adtype isa AutoMooncakeForward && startswith(c.name, "PlanarLayer inverse") - run_ad_case(c, adtype; broken=is_broken) + run_ad_case(c, adtype) end end From c12c22f52e44bd33e35183173817bab06f92990a Mon Sep 17 00:00:00 2001 From: Hong Ge Date: Thu, 18 Jun 2026 14:01:59 +0100 Subject: [PATCH 2/4] Bump version to 0.16.1 Co-Authored-By: Claude Opus 4.8 (1M context) --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index dadab71d..f8fe39de 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "Bijectors" uuid = "76274a88-744f-5084-9051-94815aaf08c4" -version = "0.16.0" +version = "0.16.1" [deps] AbstractPPL = "7a57a42e-76ec-4ea3-a279-07e840d6d9cf" From 36c099bfd8572dcb7f20cea20fd87716252d06ae Mon Sep 17 00:00:00 2001 From: Hong Ge <3279477+yebai@users.noreply.github.com> Date: Fri, 19 Jun 2026 13:55:41 +0100 Subject: [PATCH 3/4] Update Project.toml --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Project.toml b/Project.toml index f8fe39de..60b4260c 100644 --- a/Project.toml +++ b/Project.toml @@ -59,7 +59,7 @@ IrrationalConstants = "0.1, 0.2" LazyArrays = "2" LogExpFunctions = "0.3.3" MappedArrays = "0.2.2, 0.3, 0.4" -Mooncake = "0.4.95, 0.5, 0.6" +Mooncake = "0.5, 0.6" Reexport = "0.2, 1" ReverseDiff = "1" Roots = "1.3.15, 2, 3" From d4380efdf4ea26b60b95c6e0aa09bd728a652c99 Mon Sep 17 00:00:00 2001 From: Shravan Goswami Date: Fri, 19 Jun 2026 19:15:16 +0530 Subject: [PATCH 4/4] Add HISTORY.md 0.16.1 entry and raise Mooncake floor to 0.5.25 --- HISTORY.md | 4 ++++ Project.toml | 2 +- 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/HISTORY.md b/HISTORY.md index 5044f07e..f819ed29 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,3 +1,7 @@ +# 0.16.1 + +Add Mooncake 0.6 forward-mode support for `find_alpha`, and widen Mooncake compat to include 0.6. + # 0.16.0 ## Breaking changes diff --git a/Project.toml b/Project.toml index 60b4260c..9776f2c8 100644 --- a/Project.toml +++ b/Project.toml @@ -59,7 +59,7 @@ IrrationalConstants = "0.1, 0.2" LazyArrays = "2" LogExpFunctions = "0.3.3" MappedArrays = "0.2.2, 0.3, 0.4" -Mooncake = "0.5, 0.6" +Mooncake = "0.5.25, 0.6" Reexport = "0.2, 1" ReverseDiff = "1" Roots = "1.3.15, 2, 3"