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
4 changes: 4 additions & 0 deletions HISTORY.md
Original file line number Diff line number Diff line change
@@ -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
Expand Down
4 changes: 2 additions & 2 deletions Project.toml
Original file line number Diff line number Diff line change
@@ -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"
Expand Down Expand Up @@ -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.5.25, 0.6"
Reexport = "0.2, 1"
ReverseDiff = "1"
Roots = "1.3.15, 2, 3"
Expand Down
111 changes: 87 additions & 24 deletions ext/BijectorsMooncakeExt.jl
Original file line number Diff line number Diff line change
@@ -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)
Expand All @@ -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α))
Comment thread
shravanngoswamii marked this conversation as resolved.
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!!(
Expand All @@ -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!!(
Expand Down
2 changes: 1 addition & 1 deletion test/integration_tests/mooncake/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
8 changes: 1 addition & 7 deletions test/integration_tests/mooncake/main.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading