-
Notifications
You must be signed in to change notification settings - Fork 799
[Pytorch] Enable TE Op to consume extra_outputs from a previously run Op in TE Sequential #3320
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
vthumbe1503
wants to merge
34
commits into
NVIDIA:main
Choose a base branch
from
vthumbe1503:enable_extra_out_consumption
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
+816
−65
Open
Changes from all commits
Commits
Show all changes
34 commits
Select commit
Hold shift + click to select a range
b1574f0
produce/consume extra output
vthumbe1503 63192ab
allow for fusions with producer/consumer being part of same fuser wit…
vthumbe1503 3b4b523
cleanup
vthumbe1503 de38ed8
minor cleanup
vthumbe1503 385b0d5
dispatch combine impl
vthumbe1503 ad3b044
fusible ops test
vthumbe1503 5fb0d3a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 2ba4f6a
Merge remote-tracking branch 'nvidia_origin/main' into enable_extra_o…
vthumbe1503 3af2ecc
keep just ops infra changes
vthumbe1503 d7d6380
cleanup with residual tests
vthumbe1503 74f563a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 29d23f2
Merge branch 'main' into enable_extra_out_consumption
vthumbe1503 87e2b36
address review comment
vthumbe1503 80601dc
update to cleaner documentation
vthumbe1503 5070e34
address review comments
vthumbe1503 ae41ad3
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] f82cbed
some cleanup
vthumbe1503 5a4e1ec
update docs
vthumbe1503 0a479c7
pin channels through channel version
vthumbe1503 d679998
unecessary handling removal
vthumbe1503 8f7ba95
simplify
vthumbe1503 c62bb15
doc update + extra_grad = None case
vthumbe1503 a93b820
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 35b73b1
test cleanup
vthumbe1503 6801a6d
no need to check staleness in every forward call
vthumbe1503 6688e8a
remove redundant tests
vthumbe1503 12430c2
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] b189550
revert from bad names
vthumbe1503 a4cc112
keep simple
vthumbe1503 7edaf89
Merge branch 'enable_extra_out_consumption' of github.com:vthumbe1503…
vthumbe1503 5ba6055
unecessary checks
vthumbe1503 76826dc
minor doc
vthumbe1503 827f8e9
Merge branch 'main' into enable_extra_out_consumption
vthumbe1503 63a4ea3
fix lint
vthumbe1503 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -151,6 +151,139 @@ arguments and the extra outputs will be returned. | |
| the block has been split into two sections, each with one branching | ||
| operation. | ||
|
|
||
| Extra tensor channels | ||
| """"""""""""""""""""" | ||
|
|
||
| Extra inputs and Extra outputs may optionally specify a channel. Assigning | ||
| the same channel name to an extra output and one or more later extra | ||
| inputs routes the tensor internally within the same | ||
| ``OperationFuser``. An extra input connected to an earlier producer is | ||
| removed from the public ``Sequential`` arguments because the channel | ||
| supplies it. | ||
| Extra outputs remain in the public ``Sequential`` return value, | ||
| including outputs that are also consumed through a channel. | ||
|
|
||
| With a channel, the residual block above can be expressed using one | ||
| ``Sequential``: | ||
|
|
||
| .. code-block:: python | ||
|
|
||
| import torch | ||
| import transformer_engine.pytorch as te | ||
|
|
||
| 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), | ||
| make_residual, | ||
| te.ops.Linear(4096, 28672), | ||
| te.ops.SwiGLU(), | ||
| te.ops.Linear(14336, 4096), | ||
| add_residual, | ||
| ) | ||
|
|
||
| # The residual is routed internally and is also returned to the caller. | ||
| x = torch.randn(16384, 4096, device="cuda") | ||
| y, residual = block(x) | ||
|
|
||
| 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, while Dispatch's | ||
| # extra outputs are still returned in their original order. | ||
| y, m_splits, probs, routing_map = moe(x, router_probs) | ||
|
|
||
| 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 | ||
|
Comment on lines
+235
to
+238
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. It would be nice if |
||
| supported: | ||
|
|
||
| .. code-block:: python | ||
|
|
||
| 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( | ||
| make_residual, | ||
| torch.nn.Identity(), # Splits the operations into separate fusers. | ||
| add_residual, | ||
| ) | ||
|
|
||
| Use the public extra output and extra input interfaces, as in the | ||
| two-``Sequential`` example above, when the producer and consumer cannot | ||
| be placed in the same ``OperationFuser``. | ||
|
|
||
| The following conditions apply to extra tensor channels: | ||
|
|
||
| - A producer must appear before all of its consumers. Backward edges | ||
| and cycles are not supported. | ||
| - An output channel name has at most one producer, but its output may | ||
| fan out to multiple consumers. | ||
| - A named output does not require a consumer. It is still returned as | ||
| a public extra output. | ||
| - 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 are not connected to an earlier | ||
| producer in the same fuser. Channel-connected extra input slots do | ||
| not appear in the ``Sequential`` arguments. | ||
| - The caller receives every extra output in the original basic-operation | ||
| and slot order. This includes channel-bound outputs that are also | ||
| consumed internally. Gradients supplied for a returned output are | ||
| combined with gradients from its internal channel consumers. | ||
| - Channel bindings are captured when an ``OperationFuser`` (or the | ||
| fusers inside a ``Sequential``) is first constructed. Changing | ||
| ``set_extra_input_channel`` / ``set_extra_output_channel`` afterward | ||
| requires constructing a new ``OperationFuser`` or ``Sequential``. | ||
|
|
||
| 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 | ||
| --------------- | ||
|
|
||
|
|
||
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
If you read pendantically, then this technically describes our current behavior where we ignore unmatched input channels. However, this is subtle and non-obvious. Better to error out if the input channel is invalid.