Skip to content

fix: reduce-scatter gradients through the CP all-gather - #221

Open
shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/cp-gather-grad
Open

shiaho777 wants to merge 1 commit into
modelscope:mainfrom
shiaho777:fix/cp-gather-grad

Conversation

@shiaho777

Copy link
Copy Markdown
Contributor

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. flake8 is clean on the changed file.

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

restore patched function

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Inconsistent with the semantics of reduce_scatter

@hjh0119

hjh0119 commented Oct 10, 2026

Copy link
Copy Markdown
Collaborator

I would prefer not to modify the reconstruct_tensor_cp method; the related PLE gradient calculation bug will be addressed in #230.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants