Skip to content
Open
Show file tree
Hide file tree
Changes from 12 commits
Commits
Show all changes
32 commits
Select commit Hold shift + click to select a range
b1574f0
produce/consume extra output
vthumbe1503 Jul 28, 2026
63192ab
allow for fusions with producer/consumer being part of same fuser wit…
vthumbe1503 Aug 4, 2026
3b4b523
cleanup
vthumbe1503 Aug 4, 2026
de38ed8
minor cleanup
vthumbe1503 Aug 4, 2026
385b0d5
dispatch combine impl
vthumbe1503 Aug 4, 2026
ad3b044
fusible ops test
vthumbe1503 Aug 5, 2026
5fb0d3a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 5, 2026
2ba4f6a
Merge remote-tracking branch 'nvidia_origin/main' into enable_extra_o…
vthumbe1503 Aug 5, 2026
3af2ecc
keep just ops infra changes
vthumbe1503 Aug 5, 2026
d7d6380
cleanup with residual tests
vthumbe1503 Aug 5, 2026
74f563a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 5, 2026
29d23f2
Merge branch 'main' into enable_extra_out_consumption
vthumbe1503 Aug 6, 2026
87e2b36
address review comment
vthumbe1503 Aug 6, 2026
80601dc
update to cleaner documentation
vthumbe1503 Aug 7, 2026
5070e34
address review comments
vthumbe1503 Aug 7, 2026
ae41ad3
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 7, 2026
f82cbed
some cleanup
vthumbe1503 Aug 9, 2026
5a4e1ec
update docs
vthumbe1503 Aug 9, 2026
0a479c7
pin channels through channel version
vthumbe1503 Aug 9, 2026
d679998
unecessary handling removal
vthumbe1503 Aug 9, 2026
8f7ba95
simplify
vthumbe1503 Aug 9, 2026
c62bb15
doc update + extra_grad = None case
vthumbe1503 Aug 9, 2026
a93b820
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 9, 2026
35b73b1
test cleanup
vthumbe1503 Aug 9, 2026
6801a6d
no need to check staleness in every forward call
vthumbe1503 Aug 9, 2026
6688e8a
remove redundant tests
vthumbe1503 Aug 9, 2026
12430c2
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 9, 2026
b189550
revert from bad names
vthumbe1503 Aug 9, 2026
a4cc112
keep simple
vthumbe1503 Aug 9, 2026
7edaf89
Merge branch 'enable_extra_out_consumption' of github.com:vthumbe1503…
vthumbe1503 Aug 9, 2026
5ba6055
unecessary checks
vthumbe1503 Aug 9, 2026
76826dc
minor doc
vthumbe1503 Aug 9, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
112 changes: 92 additions & 20 deletions docs/examples/op_fuser/op_fuser.rst
Original file line number Diff line number Diff line change
Expand Up @@ -113,43 +113,115 @@ quantized compute.
Branching operations
^^^^^^^^^^^^^^^^^^^^

The operation fuser supports very limited branching behavior. While
the operations must be in sequential order, some operations can accept
extra inputs or produce extra outputs. For example, ``AddExtraInput``
will add an extra input tensor to the intermediate tensor and
``MakeExtraOutput`` will return the intermediate tensor as an extra
output. When calling a ``Sequential`` that contains any of these
branching operations, the extra inputs should be passed in as
arguments and the extra outputs will be returned.
The operation fuser supports limited branching behavior. While the
operations must be in sequential order, basic operations may declare
extra tensor inputs and outputs. By default, an extra tensor slot has
no channel assigned and is part of the public ``Sequential`` interface:
Comment thread
vthumbe1503 marked this conversation as resolved.
Outdated
the caller provides extra inputs as arguments, and extra outputs are
returned after the main output. Assigning the same channel name to an
output slot and a later input slot connects them internally instead.

.. code-block:: python

import torch
import transformer_engine.pytorch as te

# Construct MLP with residual connection
fc1 = te.ops.Sequential(
# Keep a residual connection inside one Sequential.
Comment thread
vthumbe1503 marked this conversation as resolved.
Outdated
make_residual = te.ops.MakeExtraOutput()
add_residual = te.ops.AddExtraInput()
make_residual.set_extra_output_channel(0, "residual")
add_residual.set_extra_input_channel(0, "residual")

block = te.ops.Sequential(
te.ops.LayerNorm(4096),
te.ops.MakeExtraOutput(), # Output residual
make_residual,
te.ops.Linear(4096, 28672),
te.ops.SwiGLU(),
)
fc2 = te.ops.Sequential(
te.ops.Linear(14336, 4096),
te.ops.AddExtraInput(), # Add residual
add_residual,
)

# Forward pass
x = torch.randn(16384, 4096, device="cuda")
y, residual = fc1(x)
y = fc2(y, residual)
y = block(x)

.. figure:: ./residual_layernorm_mlp.png
:align: center

Operations for an MLP block with a residual connection. Note that
the block has been split into two sections, each with one branching
operation.
Operations for an MLP block with a residual connection.

Extra tensor channels
"""""""""""""""""""""

An extra output and one or more later extra inputs can be assigned the
same channel name. This routes the tensor inside the
``OperationFuser`` and removes the bound slots from the public
``Sequential`` interface. In the residual example above, the caller
therefore receives only ``y`` and does not need to pass the residual
back into the block.

Channels are also useful for mixture-of-experts blocks. The following
example assumes custom ``Dispatch`` and ``Combine`` basic operations.
``Dispatch`` has one public extra input containing router probabilities
and three extra outputs: split sizes, token probabilities, and a
routing map. ``Combine`` consumes the routing map.

.. code-block:: python

import transformer_engine.pytorch as te
from my_ops import Dispatch, Combine

num_experts = 8
hidden_size = 4096
ffn_size = 14336

dispatch = Dispatch(num_experts)
fc1 = te.ops.GroupedLinear(
num_experts, hidden_size, 2 * ffn_size, bias=False
)
activation = te.ops.ScaledSwiGLU()
fc2 = te.ops.GroupedLinear(
num_experts, ffn_size, hidden_size, bias=False
)
combine = Combine(num_experts)

# Dispatch extra outputs:
# 0: split sizes, 1: token probabilities, 2: routing map
dispatch.set_extra_output_channel(0, "m_splits")
dispatch.set_extra_output_channel(1, "probs")
dispatch.set_extra_output_channel(2, "routing_map")

fc1.set_extra_input_channel(0, "m_splits")
activation.set_extra_input_channel(0, "probs")
fc2.set_extra_input_channel(0, "m_splits")
combine.set_extra_input_channel(0, "routing_map")

moe = te.ops.Sequential(dispatch, fc1, activation, fc2, combine)

# Dispatch's extra input has no channel, so the caller passes router_probs.
# Channels supply all later extra inputs internally.
y = moe(x, router_probs)

The following conditions apply to extra tensor channels:

- A producer must appear before all of its consumers. Backward edges
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

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.

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

@vthumbe1503 vthumbe1503 Aug 10, 2026

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

  1. single sequential having internal routing
  2. 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.

op1.set_extra_input_channel(0, "common_name") # has 1 extra input
op2.set_extra_input_channel(0, "common_name") # has 1 extra input
model = te.Sequential(op1,op2)
y = model(input, extra_input, extra_input) 
# y = model(input, extra_input) --> wrong

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.

channel names on the producer and consumers must match.
- A channel is scoped to one ``OperationFuser``. In a ``Sequential``,
ordinary PyTorch modules split adjacent fusible operations into
separate fusers, and channels cannot cross that boundary.
- The caller passes extra inputs that have no channel assigned and
receives extra outputs that have no channel assigned. Slots assigned
to channels are internal and do not appear in the ``Sequential``
arguments or return value.

Channel-connected basic operations may still be replaced by registered
``FusedOperation`` implementations. If a fused operation contains both
the producer and consumer of a channel, its ``fuser_forward`` and
``fuser_backward`` implementations are responsible for routing the
tensor and its gradient between those basic operations.

Developer guide
---------------
Expand Down
Loading
Loading