Skip to content

[Fix][Relax] Accept a returned tuple bound to a Var in Gradient - #20588

Open
wwoosshh wants to merge 1 commit into
apache:mainfrom
wwoosshh:fix/20498-gradient-tuple-var
Open

wwoosshh wants to merge 1 commit into
apache:mainfrom
wwoosshh:fix/20498-gradient-tuple-var

Conversation

@wwoosshh

@wwoosshh wwoosshh commented Oct 7, 2026

Copy link
Copy Markdown

Fixes #20498.

relax.transform.Gradient picks the differentiation target from the function's return value. It handled a returned Tuple expression, but when the tuple was first bound to a Var, as in

res = (lv1, lv2, lv3)
R.output(res)
return res

the Var was treated as a single return value: target_index=1 and target_index=2 failed with "When the function has only one return value, target_index can only be 0", and target_index=0 failed because the tuple itself is not a scalar. Returning (lv1, lv2, lv3) directly works for target_index 1 and 2.

This PR looks through such a binding in GradientMutator, so both forms select the same target. The backward part of the adjoint function is the same as for the tuple literal, and the original return value stays res, so it returns (res, (x_adjoint_out, y_adjoint_out)).

A returned tuple Var that is not bound to a Tuple expression (for example the result of R.split) still cannot be used, because its fields are not Vars. It is now rejected with a message that says so, instead of the "only one return value" check.

One existing test case changes. TargetNotTensor in test_report_error returned gv = R.tuple(lv1, lv1) and expected an error, because the whole tuple was taken as the target. With this change, target_index=0 selects lv1, which is a valid target, so I changed the case to return a tuple whose selected field is itself a tuple. It still checks that a non-tensor target is rejected. If you would rather keep rejecting a tuple bound to a Var, I can change this PR to only improve the error message.

Testing: added test_target_index_tuple_bound_to_var, which is test_target_index with the tuple bound to a Var, and an R.split case to test_report_error. Both fail on main and pass with this change. All 23 tests in tests/python/relax/test_transform_gradient.py pass, and with the reproducer from the issue both return forms now behave the same for target_index 0, 1 and 2.

Generated-by: Claude Code (Claude Opus 5.5)

Gradient picked the differentiation target from the return value only
when the function returned a Tuple expression directly. When the tuple
was first bound to a Var, e.g. `res = (a, b, c); return res`, the Var
was treated as a single return value: every target_index other than 0
hit "When the function has only one return value", and target_index 0
failed because the tuple itself is not a scalar.

Look through such a binding so that both forms select the same target.
A returned tuple Var that is not bound to a Tuple expression (e.g. the
result of R.split) is now rejected with a message that says so, instead
of the "only one return value" check.

The TargetNotTensor case in test_report_error returned
`gv = R.tuple(lv1, lv1)` and expected an error. With this change,
target_index 0 selects lv1, which is a valid target, so the case now
returns a tuple whose selected field is itself a tuple.

Fixes apache#20498

Generated-by: Claude Code (Claude Opus 5.5)
@wwoosshh

wwoosshh commented Oct 7, 2026

Copy link
Copy Markdown
Author

cc @tlopex @tqchen

This fixes #20498: Gradient now handles a returned tuple that is bound to a Var the same way as a returned tuple literal. One existing error case in test_report_error had to change, and the PR description explains why. Could you take a look when you have time?

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.

[Bug] Gradient rejects a function that returns a tuple bound to a Var, regardless of target_index

1 participant