Skip to content

fix: convert input tangents for non-Array inputs in forward-mode Mooncake - #1074

Draft
ParadaCarleton wants to merge 3 commits into
JuliaDiff:mainfrom
ParadaCarleton:mooncake-forward-noarray
Draft

ParadaCarleton wants to merge 3 commits into
JuliaDiff:mainfrom
ParadaCarleton:mooncake-forward-noarray

Conversation

@ParadaCarleton

Copy link
Copy Markdown

DI builds tangents of an array x with similar(x), which is an Array for inputs such as a SubArray or a Transpose. AutoMooncakeForward passed that tangent to Mooncake unchanged, so pushforward and jacobian threw 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 type typeof(x)), it is copied into a zeroed copy of x and, without friendly tangents, converted with Mooncake.primal_to_tangent!!. Inputs whose tangent type matches, such as Array, take the same path as before.

The new test in test/Back/Mooncake/test.jl covers a vector view, a matrix column view and a Transpose, with and without friendly tangents, for pushforward, jacobian and in-place pushforward!. Locally, the Mooncake test 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 gradient returns a raw Mooncake Tangent, and pushforward/jacobian via AutoMooncake fail.

Found through ImplicitDifferentiation.jl#225, whose tests skip a SubArray input under AutoMooncakeForward because of this.

@codecov

codecov Bot commented Sep 29, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 97.41%. Comparing base (bef2881) to head (55e05bd).

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           
Flag Coverage Δ
DI 97.92% <100.00%> (+<0.01%) ⬆️
DIT 96.04% <ø> (ø)

Flags with carried forward coverage won't be shown. Click here to find out more.

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@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.

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`, ...).

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.

Where exactly in the DI code does this happen?

prep.cache,
(f, prep.df),
(x, dx),
(x, input_tangent(x, dx, backend)),

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 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)

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.

Only having this fix for AbstractArrays and not for other types is kind of a band aid on the underlying problem

@gdalle
gdalle marked this pull request as draft September 30, 2026 16:07

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.

2 participants