Repository navigation
Conversation
reconstruct_tensor_cp used torch.distributed.all_gather and then spliced the local shard back into the result. all_gather does not track autograd, so only the local shard kept a gradient. PLE gathers the sequence, runs a causal short conv, and slices the output back. Tokens at the start of a rank depend on the previous rank, and that gradient was dropped. When the tensor requires grad, the gather is an autograd function whose backward reduce-scatters the full-sequence gradient onto the owning shard. Integer inputs such as token ids still use the plain all-gather. Zigzag reordering stays outside the gather, so its backward permutes the gradient back to rank order before the reduce-scatter. Checked with python3 -m pytest tests/test_cp_all_gather_grad.py. A two-rank stand-in gathers [1, 2] with the other rank at +10, and backward of the sum writes [2, 2] into the local gradient.
hjh0119
reviewed
Oct 10, 2026
Comment on lines
+90
to
+93
| finally: | ||
| sys.modules.clear() | ||
| sys.modules.update(saved) | ||
| # Each rank's compute contributes 1 to every gathered row. reduce-scatter |
Comment on lines
+81
to
+83
| def reduce_scatter(output, chunks, group=None): | ||
| seen['chunks'] = [chunk.detach().clone() for chunk in chunks] | ||
| output.copy_(sum(chunks)) |
Collaborator
There was a problem hiding this comment.
Inconsistent with the semantics of reduce_scatter
Collaborator
|
I would prefer not to modify the reconstruct_tensor_cp method; the related PLE gradient calculation bug will be addressed in #230. |
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
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
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.
reconstruct_tensor_cpusedtorch.distributed.all_gatherand then spliced the local shard back into the result.all_gatherdoes not track autograd, so only the local shard kept a gradient. PLE gathers the sequence, runs a causal short conv, and slices the output back. Tokens at the start of a rank depend on the previous rank, and that gradient was dropped.When the tensor requires grad, the gather is an autograd function whose backward reduce-scatters the full-sequence gradient onto the owning shard. Integer inputs such as token ids still use the plain all-gather. Zigzag reordering stays outside the gather, so its backward permutes the gradient back to rank order before the reduce-scatter.
Checked with
python3 -m pytest tests/test_cp_all_gather_grad.py. A two-rank stand-in gathers[1, 2]with the other rank at +10, and backward of the sum writes[2, 2]into the local gradient.flake8is clean on the changed file.