fix: convert input tangents for non-Array inputs in forward-mode Mooncake - #1074
ParadaCarleton wants to merge 3 commits into
Conversation
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #1074 +/- ##
=======================================
Coverage 97.40% 97.41%
=======================================
Files 143 143
Lines 8300 8314 +14
=======================================
+ Hits 8085 8099 +14
Misses 215 215
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
gdalle
left a comment
There was a problem hiding this comment.
Hi @ParadaCarleton, thank you for this PR! I understand where you're coming from, but I'm really not convinced that this fix belongs in DI. I'm willing to be persuaded so I'll wait for your answer to my comments, but I suggest your prepare yourself emotionally for this PR to be closed ;)
|
|
||
| Convert the tangent `dx` provided by DI for an input `x` into what Mooncake expects. | ||
|
|
||
| DI builds tangents of an array `x` with `similar(x)`, which is an `Array` for many other array types (`SubArray`, `Transpose`, ...). |
There was a problem hiding this comment.
Where exactly in the DI code does this happen?
| prep.cache, | ||
| (f, prep.df), | ||
| (x, dx), | ||
| (x, input_tangent(x, dx, backend)), |
There was a problem hiding this comment.
I'm not sure we should do this: we don't do it for Enzyme either and let it error when the tangent doesn't have the correct type. Either we do it for all backends (but normalizing tangent types is a nightmare, since the ChainRules approach is yet another one), or we don't do it for any of them.
| """ | ||
| input_tangent(x, dx, ::AnyAutoMooncake) = dx | ||
|
|
||
| function input_tangent(x::AbstractArray, dx::AbstractArray, backend::AnyAutoMooncake) |
There was a problem hiding this comment.
Only having this fix for AbstractArrays and not for other types is kind of a band aid on the underlying problem
DI builds tangents of an array
xwithsimilar(x), which is anArrayfor inputs such as aSubArrayor aTranspose.AutoMooncakeForwardpassed that tangent to Mooncake unchanged, sopushforwardandjacobianthrew for these inputs: without friendly tangents,Tangent types do not match primal types; with them,AssertionError: typeof(tangent) <: tangent_type(P).This converts the input tangent in the one-argument and two-argument forward pushforwards. When it isn't already
tangent_type(typeof(x))(or, with friendly tangents, of typetypeof(x)), it is copied into a zeroed copy ofxand, without friendly tangents, converted withMooncake.primal_to_tangent!!. Inputs whose tangent type matches, such asArray, take the same path as before.The new test in
test/Back/Mooncake/test.jlcovers a vector view, a matrix column view and aTranspose, with and without friendly tangents, forpushforward,jacobianand in-placepushforward!. Locally, theMooncaketest group passes (31480 tests, Julia 1.13, Mooncake 0.5.60).Reverse mode has a related gap that this PR doesn't touch: for these inputs
gradientreturns a raw MooncakeTangent, andpushforward/jacobianviaAutoMooncakefail.Found through ImplicitDifferentiation.jl#225, whose tests skip a
SubArrayinput underAutoMooncakeForwardbecause of this.