[Pytorch] Enable TE Op to consume extra_outputs from a previously run Op in TE Sequential - #3320
[Pytorch] Enable TE Op to consume extra_outputs from a previously run Op in TE Sequential#3320vthumbe1503 wants to merge 44 commits into
Conversation
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
…h error handling tests Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
for more information, see https://pre-commit.ci
Greptile SummaryThe PR adds named extra-tensor channels for routing outputs between operations in one PyTorch OperationFuser, including public-output controls and backward gradient fan-out.
Confidence Score: 4/5The PR does not yet appear safe to merge because transient and discarded fusers still permanently prevent valid later channel configuration on the affected operations. Standalone operation calls construct temporary OperationFusers that permanently lock their basic operations, and Sequential mutation discards cached fusers without releasing locks held by surviving operations; both previously reported lifecycle failures remain reachable even though stale routing on retained fusers is now prevented. Files Needing Attention: transformer_engine/pytorch/ops/fuser.py, transformer_engine/pytorch/ops/op.py Important Files Changed
Sequence DiagramsequenceDiagram
participant Caller
participant Seq as Sequential
participant Fuser as OperationFuser
participant Producer
participant Consumer
Caller->>Seq: forward(input, public extra inputs)
Seq->>Fuser: execute fused group
Fuser->>Producer: fuser_forward(input)
Producer-->>Fuser: main output + named extra output
Fuser->>Consumer: fuser_forward(main output, channel tensor)
Consumer-->>Fuser: final output
Fuser-->>Caller: final output + public extra outputs
Caller->>Fuser: backward(output gradients)
Fuser->>Consumer: backward
Consumer-->>Fuser: channel gradient
Fuser->>Producer: backward(accumulated channel gradient)
Reviews (26): Last reviewed commit: "a bit of doc" | Re-trigger Greptile |
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
There was a problem hiding this comment.
Additional review comments from Codex:
1. [High] Internal channel outputs lose the signal that their gradient is required.
transformer_engine/pytorch/ops/fuser.py:170 only calls requires_grad_ for public outputs at line 194. A fresh tensor created by an internal producer inside
torch.autograd.Function.forward therefore arrives at its consumer with requires_grad=False. Existing operations such as transformer_engine/pytorch/ops/basic/
swiglu.py:468 and ScaledSReLU use that flag to decide whether to compute the extra-input gradient. They consequently return None, silently dropping the
gradient to a differentiable producer such as the documented router-probability dispatch. The new tests do not expose this because MakeExtraOutput returns the
original input tensor.
2. [Medium] Channel routing violates the declared iterable output contract.
transformer_engine/pytorch/ops/fuser.py:140 applies len() and indexing to a producer’s extra outputs, but transformer_engine/pytorch/ops/op.py:92 permits any
Iterable[Iterable[Tensor]]. A custom Dispatch returning a generator works under the old flattening logic but now fails when a later channel consumer executes.
3. [Medium] A standalone operation call permanently prevents later channel configuration.
transformer_engine/pytorch/ops/op.py:598 constructs a temporary OperationFuser, while transformer_engine/pytorch/ops/fuser.py:509 permanently locks every
attached operation. Calling an operation once through its normal forward, then placing it into a channel-connected Sequential, makes either setter raise even
though the temporary fuser no longer exists.
4. [Medium] The tests do not exercise two major routing branches.
tests/pytorch/test_fusible_ops.py:460 registers only a forward fusion, so backward remains unfused and never exercises the same-fusion skip at
transformer_engine/pytorch/ops/fuser.py:323. The multi-output test at tests/pytorch/test_fusible_ops.py:624 only checks duplicate-name rejection; its custom
operation never runs. Thus mixed bound/unbound slot ordering, filtered autograd returns, and the modified two-input GroupedLinear(scale_bias=True) behavior
remain unproved.
## Suggested repairs
- For finding 1: Preserve the gradient-requirement flag on every extra output before classifying it as public or internal, or carry equivalent explicit per-slot
metadata. Add a producer that creates a fresh tensor and verify gradient propagation through ScaledSwiGLU or ScaledSReLU.
- For finding 2: Materialize and validate each operation’s extra outputs as a tuple immediately after fuser_forward; store that tuple for later consumers and
lifetime tracking.
- For finding 3: Make transient fusers created by BasicOperation.forward non-locking, while persistent Sequential fusers retain immutable routing. Add a call-
then-bind regression test.
- For finding 4: Add a joint/backward fused residual operation and assert fusion selection plus input/parameter/channel gradients. Add a successful multi-input/
multi-output routing test and a GroupedLinear(scale_bias=True) channel case.
I would like you to also take a look at the tests - multiple of them duplicate each other (e.g. test_channel_fan_out_accumulates_grad is a stronger duplicate of test_internal_extra_tensor_channel_fanout).
| and cycles are not supported. | ||
| - A channel has exactly one producer, but its output may fan out to | ||
| multiple consumers. | ||
| - Every named output channel must have at least one consumer, and the |
There was a problem hiding this comment.
This limitation that there has to be at least one consumer in the named channel seems
arbitrary to me. If we do not strictly need this behavior then we shouldn't have that
as it would introduce friction when somebody needs to refactor the code using those
named channels by splitting the sequential - now they also need to remove the channel
names. In fact, I would expect people to generally want to name their extra outputs and
inputs even if they would not be reused inside the sequential. That could also enable
us to accept and return the dictionary rather than a list (which would make it less
fragile).
There was a problem hiding this comment.
Originally I wanted to restrict the channel naming as a way to just do internal routing of tensors to reduce possibility of errors and the friction was kind of intentional. But I see your point of making it more seamless for user in future to construct a big a sequential op. If they want to refactor a code from 1 to 2 below
- single sequential having internal routing
- Two sequentials with one sequential passing extra output as extra input to another sequential
This can indeed be a problem since there might be some use-cases just supporting 2 but not 1.
And so I have removed that restriction. However, supporting extra_input and extra_output as dictionaries would be a problem from backwards compatibility perspective. Also, I want to restrict the scope of this PR. And allowing for dict based extra_input and extra output can be a seperate PR.
I have one extra requirement from named extra input channel added currently. If two different extra inputs share the same channel name, and is not internally connected to extra output of a previous op. Caller/User should still provide the extra_input two times.
This is done so that user's code doesnt have to change while naming an input channel vs not naming it. Also as you can see introducing extra_input dict is also going to make this tricky from backward compatibility perspective.
There was a problem hiding this comment.
Update: Tim has a valid point below of extra_input's complexity. For the named extra_input which is not connected to any producer op, we are simply raising an error. Its users responsibility to unname the channel if it is already named(by setting it to None) and is not connected to any internal producer.
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
for more information, see https://pre-commit.ci
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
…/TransformerEngine into enable_extra_out_consumption
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
|
/te-ci L1 pytorch |
| Channels cannot connect operations in different ``OperationFuser`` | ||
| instances. In particular, an ordinary PyTorch module inside a | ||
| ``Sequential`` splits the fusible operations on either side into | ||
| separate fusers. The following channel connection is therefore not |
There was a problem hiding this comment.
It would be nice if Sequential could handle channels across OperationFusers, but the implementation would be quite hairy and not worth it for the current effort.
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
Co-authored-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
…tted Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
|
/te-ci pytorch |
| if self.has_stale_op_channels(): | ||
| raise RuntimeError( | ||
| "Extra tensor channels changed after this OperationFuser captured " | ||
| "its routing. Construct a new OperationFuser." | ||
| ) |
There was a problem hiding this comment.
Do we need to do this on every call? We could just make it impossible to change them with
the API itself (e.g. have the fuser mark the ops as finalized when it takes them).
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
|
/te-ci pytorch |
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: