Repository navigation
Add Mooncake forward and reverse rules - #225
ParadaCarleton wants to merge 3 commits into
Conversation
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.
|
Thanks! I authorized the tests and I'll take a look once they pass :) |
Codecov Report❌ Patch coverage is
🚀 New features to boost your workflow:
|
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.
|
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 |
|
I'm not sure why we should accept a tangent of a different type in DI if Mooncake itself doesn't? |
gdalle-bot
left a comment
There was a problem hiding this comment.
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 bothForwardModeandReverseModeon a basic scenario.- Results are correct for matrix-shaped
x,DirectLinearSolver+MatrixRepresentation,Float32,xpassed again as an extra argument, and a solver whoseyaliasesx. - 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.
- With a
SubArrayinput that doesn't cover its whole parent, reverse mode returns a wrong gradient without any error. backendsis not respected on the extra-derivatives path, contrary to what the description says, and that path is triggered by something as simple asIterativeLinearSolver(; rtol=...).
- With a
- Docs. The PR changes the semantics for
argsin Mooncake, but two places still say the old thing: theImplicitFunctiondocstring (src/implicit_function.jl, "the following positional argumentsargsare considered constant") anddocs/src/faq.md("All of the positional arguments apart fromxwill get zero tangents"). Both need updating, including a note that ForwardDiff and ChainRules still treatargsas constant. Whether this difference between backends is acceptable is for @gdalle to decide. - Minor. An integer
xwith a differentiableFloat64argument errors in both modes (DI.derivative(a -> first(implicit([1, 2], a)), AutoMooncake(; config=nothing), 2.0)throws "Cannotconvertan object of type NoTangent to Int64"). This isn't necessarily a blocker, but it's worth an explicit error or a test, since the newargsfeature invites exactly this use. - Tests. Consider adding
Mooncake.TestUtils.test_rulefor the plain and prepared signatures, with non-trivialargsand 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)) |
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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...) |
There was a problem hiding this comment.
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
There was a problem hiding this comment.
| ForwardDiff = "0.10.36, 1" | ||
| KrylovKit = "0.10.0" | ||
| LinearAlgebra = "1" | ||
| Mooncake = "0.5" |
There was a problem hiding this comment.
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
left a comment
There was a problem hiding this comment.
Thank you for getting this started! Here is a first round of feedback, from me and my trusted bot ;)
| JET = "0.9, 0.10, 0.11, 0.12" | ||
| KrylovKit = "0.10.2" | ||
| LinearAlgebra = "1" | ||
| Mooncake = "0.5" |
There was a problem hiding this comment.
A more precise lower bound is probably necessary
|
|
||
| @testset "Jacobian - $outer_backend" begin | ||
| if outer_backend isa AutoForwardDiff | ||
| if outer_backend isa Union{AutoForwardDiff,AutoMooncake,AutoMooncakeForward} |
There was a problem hiding this comment.
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
| AutoMooncake(; config=nothing), | ||
| AutoMooncakeForward(; config=nothing), |
There was a problem hiding this comment.
Do we need to test friendly tangents?
| # The conditions are differentiated with Mooncake too, unless `backends` says otherwise. | ||
| const FORWARD_BACKEND = AutoMooncakeForward(; config=nothing) | ||
| const REVERSE_BACKEND = AutoMooncake(; config=nothing) |
There was a problem hiding this comment.
The default choice of inner backend should probably propagate the options of the outer backend (eg friendly tangents or debug mode)
| if t isa AbstractArray | ||
| return t |
There was a problem hiding this comment.
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}, |
There was a problem hiding this comment.
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...) |
There was a problem hiding this comment.
| nothing | ||
| end | ||
| dc = B(tangent_to_array(x0, tangent(x))) | ||
| if has_other_tangents(implicit, args0) |
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
Same remark on the "other tangents" case
|
|
||
| function Mooncake.rrule!!( | ||
| implicit::CoDual{<:ImplicitFunction}, | ||
| prep::CoDual{<:ImplicitFunctionPreparation}, |
This adds a
Mooncakepackage extension withfrule!!andrrule!!for calling anImplicitFunction, prepared or unprepared, on anyAbstractArrayinput. 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.backendsis 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!!inAsPrimalmode,primal_to_tangent!!). The conversion is skipped when the tangent type is the array type itself (Array,CuArray), soComponentVectorinputs also work under Mooncake.Positional arguments beyond
x, and data captured in the conditions. A Mooncake rule sees a tangent for every input, including DIConstants, 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:B dx + ∂c/∂(implicit, args) · (dimplicit, dargs)implicitandargsare the pullback of that map atdcThis extra pass is skipped whenever
(implicit, args...)has no differentiable content (tangent_typeisNoTangent), 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.ADTypeslower bound raised to 1.17, the first release withAutoMooncakeForward.Tests
AutoMooncakeandAutoMooncakeForwardadded to the default outer backends intest_implicit, including the prepared-call path. This covers every existing scenario: matrix and operator representations, all linear solvers, matrix-shapedx,ComponentVector, andadd_arg_mult.x, and with respect to data captured in the conditions, compared against the explicit function under both Mooncake modes. It also checksxwith aFloat64constant argument.systematic.jlandformalities.jlgive 821 passed and the existing 1 broken.Related
With
IterativeLeastSquaresSolver, KrylovKit's LSMR currently returnsNaNfor a zero right-hand side, so a zero tangent or cotangent givesNaNderivatives 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.