Skip to content

Add Mooncake forward and reverse rules - #225

Open
ParadaCarleton wants to merge 3 commits into
JuliaDecisionFocusedLearning:mainfrom
ParadaCarleton:mooncake-rules
Open

ParadaCarleton wants to merge 3 commits into
JuliaDecisionFocusedLearning:mainfrom
ParadaCarleton:mooncake-rules

Conversation

@ParadaCarleton

Copy link
Copy Markdown

This adds a Mooncake package extension with frule!! and rrule!! for calling an ImplicitFunction, prepared or unprepared, on any AbstractArray input. Mooncake can then differentiate implicit functions in forward mode (AutoMooncakeForward) and reverse mode (AutoMooncake) without tracing through the solver. It complements the Enzyme work in #220 / #221 and leaves that code alone.

Design

  • The rules reuse build_A, build_Aᵀ, build_B, build_Bᵀ and the linear solvers the same way the ForwardDiff and ChainRules extensions do. backends is respected, and by default the conditions are differentiated with Mooncake as well.

  • Array types. Tangents are converted to arrays shaped like the primal and back with Mooncake's own conversions (tangent_to_friendly!! in AsPrimal mode, primal_to_tangent!!). The conversion is skipped when the tangent type is the array type itself (Array, CuArray), so ComponentVector inputs also work under Mooncake.

  • Positional arguments beyond x, and data captured in the conditions. A Mooncake rule sees a tangent for every input, including DI Constants, so a rule that zeroed these tangents would silently return wrong derivatives whenever they are actually needed. Instead, the rules differentiate (implicit, args...) -> implicit.conditions(x, y, z, args...) with Mooncake at the solution, so these derivatives are exact:

    • forward: the right-hand side is B dx + ∂c/∂(implicit, args) · (dimplicit, dargs)
    • reverse: the cotangents for implicit and args are the pullback of that map at dc

    This extra pass is skipped whenever (implicit, args...) has no differentiable content (tangent_type is NoTangent), so the common case pays nothing. Data captured by the solver correctly gets zero derivative, as long as the conditions determine the solution.

Compat

  • Mooncake = "0.5" as a weak dependency.
  • ADTypes lower bound raised to 1.17, the first release with AutoMooncakeForward.

Tests

  • AutoMooncake and AutoMooncakeForward added to the default outer backends in test_implicit, including the prepared-call path. This covers every existing scenario: matrix and operator representations, all linear solvers, matrix-shaped x, ComponentVector, and add_arg_mult.
  • New testitem "Mooncake other inputs": Jacobians with respect to an argument beyond x, and with respect to data captured in the conditions, compared against the explicit function under both Mooncake modes. It also checks x with a Float64 constant argument.
  • Locally on Julia 1.13 with Mooncake 0.5.60: systematic.jl and formalities.jl give 821 passed and the existing 1 broken.

Related

With IterativeLeastSquaresSolver, KrylovKit's LSMR currently returns NaN for a zero right-hand side, so a zero tangent or cotangent gives NaN derivatives on every backend. The warnings in the test logs come from Mooncake's preparation pass, which runs a zero cotangent. The fix is in Jutho/KrylovKit.jl#169.

Mooncake frule!! and rrule!! for ImplicitFunction calls with any AbstractArray input, prepared or not. Tangents for the positional arguments beyond x and for data captured in the conditions are propagated exactly, by differentiating the conditions with respect to them.
@gdalle

gdalle commented Sep 24, 2026

Copy link
Copy Markdown
Member

Thanks! I authorized the tests and I'll take a look once they pass :)

@codecov

codecov Bot commented Sep 24, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 98.75000% with 1 line in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
ext/ImplicitDifferentiationMooncakeExt.jl 98.75% 1 Missing ⚠️
Files with missing lines Coverage Δ
ext/ImplicitDifferentiationMooncakeExt.jl 98.75% <98.75%> (ø)
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

Convert tangents through a buffer of the primal's own type and copy the
result to `zero(x)`, the tangent type the operators are prepared with, so
that e.g. `SubArray` inputs work. Convert the output of `Bᵀ` too, which is a
raw Mooncake tangent when the conditions are differentiated with Mooncake.

Test a `SubArray` input and a `z` with nonzero rdata, which covers the
remaining lines of the extension.
@ParadaCarleton

ParadaCarleton commented Sep 29, 2026 •

Copy link
Copy Markdown
Author

Fixed the formatting check (bc58946). The codecov gap turned out to be a real bug: inputs whose Mooncake tangent isn't an array, such as SubArray, crashed the conversions. 80feeaf fixes that and adds tests, including a z with nonzero rdata. One limitation remains, and it comes from DifferentiationInterface rather than this PR: AutoMooncakeForward pushforward rejects an array tangent for a SubArray input, so that one combination is skipped in the tests. The fix is in JuliaDiff/DifferentiationInterface.jl#1074. Could you approve the workflow runs?

@gdalle

gdalle commented Sep 29, 2026

Copy link
Copy Markdown
Member

I'm not sure why we should accept a tangent of a different type in DI if Mooncake itself doesn't?

@gdalle-bot gdalle-bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Note

I'm a bot (Claude Code) reviewing on behalf of @gdalle. A manual review will follow, so please treat these comments as preliminary.

Thanks for this! The rules are well structured. I reviewed commit 80feeaf and ran some extra checks locally (Julia 1.12, Mooncake 0.5.59):

  • Things that work.
    • Mooncake.TestUtils.test_rule(rng, implicit, x; is_primitive=true, mode=...) passes in both ForwardMode and ReverseMode on a basic scenario.
    • Results are correct for matrix-shaped x, DirectLinearSolver + MatrixRepresentation, Float32, x passed again as an extra argument, and a solver whose y aliases x.
    • Once prepared, the overhead is reasonable: for a 3×3 Jacobian, about 55 µs in forward mode and 115 µs in reverse mode, compared with 7 µs for ForwardDiff.
  • Two bugs I could reproduce. Details are in the inline comments.
    1. With a SubArray input that doesn't cover its whole parent, reverse mode returns a wrong gradient without any error.
    2. backends is not respected on the extra-derivatives path, contrary to what the description says, and that path is triggered by something as simple as IterativeLinearSolver(; rtol=...).
  • Docs. The PR changes the semantics for args in Mooncake, but two places still say the old thing: the ImplicitFunction docstring (src/implicit_function.jl, "the following positional arguments args are considered constant") and docs/src/faq.md ("All of the positional arguments apart from x will get zero tangents"). Both need updating, including a note that ForwardDiff and ChainRules still treat args as constant. Whether this difference between backends is acceptable is for @gdalle to decide.
  • Minor. An integer x with a differentiable Float64 argument errors in both modes (DI.derivative(a -> first(implicit([1, 2], a)), AutoMooncake(; config=nothing), 2.0) throws "Cannot convert an object of type NoTangent to Int64"). This isn't necessarily a blocker, but it's worth an explicit error or a test, since the new args feature invites exactly this use.
  • Tests. Consider adding Mooncake.TestUtils.test_rule for the plain and prepared signatures, with non-trivial args and captured data. It checks fdata/rdata handling and compares against finite differences, which Jacobian comparisons alone don't do.

Generated by Claude Code

if a isa tangent_type(typeof(x))
return a
else
return primal_to_tangent!!(zero_tangent(x), copyto!(deepcopy(x), a))

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Wrong gradient without any error for views that don't cover their whole parent (reverse mode). deepcopy(x) copies the entire parent of a SubArray, so primal_to_tangent!! turns the parent entries outside the view (which hold primal values) into tangent entries. Reproducer:

implicit = ImplicitFunction(x -> (sqrt.(x), nothing), (x, y, z) -> y .^ 2 .- x)
p = [1.0, 2.0, 3.0]
DI.jacobian(p -> first(implicit(view(p, 1:2))), AutoMooncake(; config=nothing), p)
# [0.5 0.0 3.0; 0.0 0.354 3.0]   <- third column should be 0, it contains p[3]

Plain DI + Mooncake gives the right answer for p -> sqrt.(view(p, 1:2)). The existing test uses view(x, :), which covers the whole parent, so it doesn't catch this. The buffer needs to be zero outside the view before conversion. Please add a regression test with view(p, 1:2).

In forward mode, the same input fails with ArgumentError: Tangent types do not match primal types, even though plain DI AutoMooncakeForward handles p -> sqrt.(view(p, 1:2)). I think this comes from how the extension passes an array tangent for a SubArray primal to the inner DI call, not from DI itself. That may also answer @gdalle's question on DI#1074.


Generated by Claude Code

end

function has_other_tangents(implicit::ImplicitFunction, args::Tuple)
return tangent_type(typeof((implicit, args...))) !== NoTangent

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This check covers the whole ImplicitFunction, including linear_solver, so a plain tolerance is enough to switch the extra pass on:

Mooncake.tangent_type(typeof(IterativeLinearSolver(; rtol=1e-10)))
# Tangent{@NamedTuple{kwargs::Tangent{@NamedTuple{data::@NamedTuple{rtol::Float64}, itr::NoTangent}}}}

ConditionsAt only reads implicit.conditions, so the check (and the differentiated argument) could be narrowed to implicit.conditions. The result would stay correct, the unnecessary work would go away (about 45% extra in reverse mode on a small example), and most users would no longer hit the backends issue below.


Generated by Claude Code

dc = B(tangent_to_array(x0, tangent(x)))
if has_other_tangents(implicit, args0)
f = ConditionsAt(x0, y, z)
cache = prepare_derivative_cache(f, implicit, args0...)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

backends is bypassed here, and at prepare_pullback_cache in the pullback. The PR description says backends is respected, but this pass always differentiates the conditions with Mooncake. Reproducer, where the conditions are ForwardDiff-compatible but not Mooncake-compatible:

bad(v) = v
Mooncake.@is_primitive Mooncake.DefaultCtx Tuple{typeof(bad),Vector{Float64}}  # no rule defined
backends = (; x=AutoForwardDiff(), y=AutoForwardDiff())
implicit = ImplicitFunction(x -> (sqrt.(x), nothing), (x, y, z) -> y .^ 2 .- bad(x);
    backends, linear_solver=IterativeLinearSolver(; rtol=1e-10))
DI.jacobian(x -> first(implicit(x)), AutoMooncake(; config=nothing), [1.0, 2.0, 3.0])
# MethodError: no method matching rrule!!(::CoDual{typeof(bad), NoFData}, ...)

Without rtol, this works in both modes. Users choose backends precisely when the outer backend can't differentiate the conditions, so this path should go through backends.x (for example, DI with the extra arguments as active inputs), or at least document the limitation clearly. The existing (; x=AutoForwardDiff(), y=AutoZygote()) × rtol=1e-8 test scenario only passes because its conditions happen to work with Mooncake.


Generated by Claude Code

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Comment thread Project.toml
ForwardDiff = "0.10.36, 1"
KrylovKit = "0.10.0"
LinearAlgebra = "1"
Mooncake = "0.5"

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The extension relies on several names that Mooncake doesn't declare @public: FriendlyTangentCache, AsPrimal, tangent_to_friendly!!, primal_to_tangent!!, zero_rdata, increment!!, and others. With a "0.5" bound, any 0.5.x release could break it. I checked that 0.5.45 still has these names, but I haven't verified the earliest 0.5.x versions. Could you raise the lower bound to the oldest version you've actually tested, and prefer public API where possible?


Generated by Claude Code

@gdalle gdalle left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thank you for getting this started! Here is a first round of feedback, from me and my trusted bot ;)

Comment thread test/Project.toml
JET = "0.9, 0.10, 0.11, 0.12"
KrylovKit = "0.10.2"
LinearAlgebra = "1"
Mooncake = "0.5"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A more precise lower bound is probably necessary

Comment thread test/utils.jl

@testset "Jacobian - $outer_backend" begin
if outer_backend isa AutoForwardDiff
if outer_backend isa Union{AutoForwardDiff,AutoMooncake,AutoMooncakeForward}

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'm not sure we can reliably implement the prepared version for Mooncake. This option was mostly meant as an internal for a research experiment, I haven't thought much about generalizing it beyond ForwardDiff

Comment thread test/utils.jl
Comment on lines +210 to +211
AutoMooncake(; config=nothing),
AutoMooncakeForward(; config=nothing),

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do we need to test friendly tangents?

Comment on lines +38 to +40
# The conditions are differentiated with Mooncake too, unless `backends` says otherwise.
const FORWARD_BACKEND = AutoMooncakeForward(; config=nothing)
const REVERSE_BACKEND = AutoMooncake(; config=nothing)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The default choice of inner backend should probably propagate the options of the outer backend (eg friendly tangents or debug mode)

Comment on lines +52 to +53
if t isa AbstractArray
return t

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is it always true that if a tangent has an array type, it is the right array type (same as the primal)?


function Mooncake.frule!!(
implicit::Dual{<:ImplicitFunction},
prep::Dual{<:ImplicitFunctionPreparation},

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't know how to handle this case, which is why I said the prepared version probably isn't ready for public consumption. In the ForwardDiff extension, the corresponding method is only defined when the prep doesn't contain duals

dc = B(tangent_to_array(x0, tangent(x)))
if has_other_tangents(implicit, args0)
f = ConditionsAt(x0, y, z)
cache = prepare_derivative_cache(f, implicit, args0...)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nothing
end
dc = B(tangent_to_array(x0, tangent(x)))
if has_other_tangents(implicit, args0)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This case isn't handled by the other backends and so we probably shouldn't handle it here either (it would require implicit to be differentiable, which goes against the whole point of the package). The issue is that we don't have an equivalent of ChainRulesCore.@not_implemented in Mooncake. All ideas welcome!

dc = linear_solver(Aᵀ, A, -dy, zero(c))
tx = array_to_tangent(x, copyto!(similar(x), tangent_to_array(x, Bᵀ(dc))))
increment!!(fx, fdata(tx))
if has_other_tangents(implicit, args)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same remark on the "other tangents" case


function Mooncake.rrule!!(
implicit::CoDual{<:ImplicitFunction},
prep::CoDual{<:ImplicitFunctionPreparation},

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We can't handle this case either

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants