[bugfix] _DimZeroAllToAll silently fails to send grads under torch.compile - #8491
Conversation
Signed-off-by: Stas Bekman <stas@stason.org>
…mpile Signed-off-by: Stas Bekman <stas@stason.org>
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: 39d1e65e46
ℹ️ About Codex in GitHub
Codex has been enabled to automatically review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
When you sign up for Codex through ChatGPT, Codex can also answer questions or update the PR, like "@codex address that feedback".
| def forward(ctx: Any, group: dist.ProcessGroup, input: Tensor) -> Tensor: | ||
| world_size = dist.get_world_size(group) | ||
| assert input.shape[0] == world_size, f"Dim 0 {input.shape[0]} is not world size" | ||
| def _dim_zero_all_to_all(group: dist.ProcessGroup, input: Tensor) -> Tensor: |
There was a problem hiding this comment.
In general I'd recommend dropping all the original collective autograd functions in Deepspeed which were written when pytorch didn't have them and replace with torch.functional version of each, it'd be simpler and better for future proofing.
I already proposed doing that when porting Ulysses to the newer version, a year ago, but at the end decided to leave things as they are being unaware those functions silently failed torch.compile and destroyed grads.
Signed-off-by: Stas Bekman <stas@stason.org>
A Ulysses job whose decoder layers are
torch.compiled delivers no gradient at all to the attention query, key and value projections. The forward pass is numerically correct, the loss is finite, no warning is raised, and attention simply stops learning: one optimizer step leaves those tensors bit-identical with their optimizer moments at exactly zero, while the output projection, the MLP and everything on the residual stream update normally.The sequence/head exchange was a
torch.autograd.Functionwhose forward allocatedtorch.empty_like(input)and filled it withall_to_all_single, a collective that writes into an output argument instead of returning a value. Dynamo does not treat an autograd function as an opaque call: it inlines the forward body and differentiates what it traced. The traced body allocates an empty tensor and mutates it through an operation the tracer does not model as producing data, so in the graph the exchange output carries no dependence on the exchange input, the derivative is exactly zero, and the hand-written backward -- the transposed exchange -- never runs. At execution time the collective still moves the bytes, which is why nothing about the forward looks wrong._dim_zero_all_to_allindeepspeed/sequence/layer.pyreplaces that function and callstorch.distributed.nn.functional.all_to_all_single, which returns the exchanged tensor. The dependence is then in the graph and the generated backward is the transposed exchange. Both call sites indeepspeed/runtime/sequence_parallel/ulysses_sp.pyare updated; there is no other caller.The commented-out line the previous implementation carried -- the same differentiable collective, called from inside the autograd function -- does not fix this. Holding everything else constant and varying only the implementation, on 2 and 4 ranks with the exchange inside
torch.compile, against the eager gradient as reference:torch.autograd.Function(previous)torch.distributed.nn.functional.all_to_all_singleinside atorch.autograd.Functiontorch.distributed.nn.functional.all_to_all_singlecalled directly (this PR)torch.distributed._functional_collectives.all_to_all_singleinside atorch.autograd.Functiontorch.distributed._functional_collectives.all_to_all_singlecalled directlyThe two failing rows share the enclosing autograd function, which is what the tracer inlines. Removing it is the load-bearing part of the change; the public functional collective is preferred over
_functional_collectivesbecause it is not a private API.Whether a given release inlines that body is a property of the tracer, not of this code: the same exchange on torch 2.9.1 keeps its gradient through a compiled region and loses it on torch 2.11.0. Code that reaches the wrong answer only on newer torch is worth fixing at the source rather than pinning around.
single_all_to_allin the same module has the same shape --torch.empty_likefollowed by a mutating collective -- on theDistributedAttentionpath. It is not changed here: itsasync_op=Truebranch returns a work handle that a functional collective does not provide, so it needs a different treatment and its own measurement.Testing
tests/unit/ulysses_alst/test_dim_zero_all_to_all.pyis new. Two ranks, both arms eager and compiled:sum(exchange(x) ** 2)is2 * xon every rank whichever slice a rank holds. The compiled arm of this test returns a zero gradient against the previous implementation and the exact value against this one.iof a rank's output holds the slot that rank owns in ranki's input, so a substitution that is differentiable but exchanges the wrong data cannot pass.The failing behaviour, measured on the previous implementation with a two-line driver -- one
torch.compiled call of the exchange, losssum(y ** 2), H200, torch 2.11.0+cu130:The loss agrees to all printed digits between the two arms (
56.828213on 2 ranks), which is the reason this is invisible in a training run.Run against the previous exchange behind the new name, the test file reports
1 failed, 3 passed: the compiled gradient arm fails with all 64 elements mismatched and a greatest absolute difference of 5.99, while the eager gradient arm and both forward arms pass. Against this change it reports4 passed. Two H200s, torch 2.11.0+cu130.tests/unit/ulysses_alst/test_ulysses_sp_hf.pycannot be collected in the environment used here for a reason that predates this change: importing the module fails on a missing optionaltransformersdependency, before any DeepSpeed code runs.End to end on a Hugging Face causal model with 24 attention heads, 4 key/value heads, head dimension 256, 1024 tokens, ZeRO-2,
sdpa, decoder layers compiled: before the changek_projandv_projreceive 0 of 5,242,880 gradient elements at sequence parallelism 2 and at 8; after it they receive 5,239,733 and 5,242,536. On a gated-attention modelq_projreads exactly half nonzero before the change, because the gate half of that projection is applied downstream of the exchange and keeps its gradient -- which is why a whole-tensor gradient norm onq_projdoes not reveal the defect.Also on this branch, unrelated to the above:
TiledMLP.backwardreshapes rather than views when flattening the activation and the incoming gradient, so a caller passing a non-contiguous tensor is not rejected, with a case intests/unit/ulysses_alst/test_tiled_compute.py.Not run: the
DistributedAttentionpath described above, and any measurement of throughput change from the substituted collective.