[PyTorch] Keep DistributedWeight objects on ctx across saved-tensor hooks - #3545
xrennvidia wants to merge 7 commits into
Conversation
|
…ooks _Linear and _GroupedLinear save the DistributedWeight parameter itself for backward and gate their distributed-weight path on getting that object back (is_distributed_weight(saved_weight)). That holds only while no torch.autograd.graph.saved_tensors_hooks are installed: with hooks active, SavedVariable::unpack rebuilds a fresh plain tensor from the unpack hook's result, dropping the Python subclass and its is_distributed_weight marker. The backward then takes the plain-parameter branch, whose weakrefs point at the transient all-gathered weights, and fails with "weight was removed while fuse_wgrad_accumulation=True". Seen with Megatron-Core fine-grained activation offloading (which installs such hooks around expert fc1/act) on bf16 GTP-sharded GroupedLinear experts. Quantized weights were unaffected because their storage objects already stay on ctx via prepare_for_saving/restore_from_saved, and the op-fuser path reads the module's own weight attribute instead of the saved tensor. Keep the DistributedWeight objects as ordinary ctx / LinearBwdArgs references and prefer them over the saved tensors in backward. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Signed-off-by: Xiaowei Ren <xren@nvidia.com>
Regression coverage for keeping DistributedWeight objects on the autograd
context: with torch.autograd.graph.saved_tensors_hooks installed, autograd
unpacks a saved leaf as a fresh plain tensor, so gating backward on
is_distributed_weight(saved_weight) silently takes the plain-parameter branch.
Tests, using an in-repo fake DistributedWeight with distinct forward/backward
materialize scales and hook call counters:
* premise: a saved DistributedWeight comes back as torch.Tensor under hooks;
* Linear and GroupedLinear (split-quantize path), with and without fused
weight-gradient accumulation: under identity hooks every DistributedWeight
hook still fires and outputs / dgrad / wgrad match the unhooked run bitwise.
Without the fix the hooked cases fail in the dgrad GEMM (wrong weight
representation) or with "weight was removed while fuse_wgrad_accumulation=True".
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: Xiaowei Ren <xren@nvidia.com>
45d5cb7 to
eea3da3
Compare
|
LGTM, thanks for the fix! |
|
@xrennvidia Can you rebase this with main? #3517 has been merged |
…d-tensor-hooks # Conflicts: # transformer_engine/pytorch/module/grouped_linear.py
Done. |
This comment has been minimized.
This comment has been minimized.
|
Want your agent to iterate on Greptile's feedback? Try greploops. |
…upedLinear path The use_grouped_tensor=True backward rebuilt origin_weights from the unpacked saved tensors, so under saved-tensor hooks the DistributedWeight subclass was lost and materialize/finalize dispatch was skipped. Apply the same ctx.dist_weights fallback used by the split-quantize path and cover both GEMM paths in the test. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Signed-off-by: Xiaowei Ren <xren@nvidia.com>
…h is unsupported GroupedLinear silently falls back to split-quantize when the device or cuBLASLt cannot run the native grouped-tensor path, which would make the grouped-tensor test case a duplicate rather than coverage of the native backward dispatch. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Signed-off-by: Xiaowei Ren <xren@nvidia.com>
9d8b650 to
dddf418
Compare
| dist_weights = getattr(ctx, "dist_weights", None) | ||
| if dist_weights is not None: | ||
| weight_tensors = dist_weights | ||
| ctx.dist_weights = None |
There was a problem hiding this comment.
Second backward loses distributed weights
When GroupedLinear runs under saved-tensor hooks and its output is backpropagated twice with retain_graph=True, the first backward clears ctx.dist_weights here; the split backward path clears it too. On the second backward, the saved weights unpack as plain tensors, so the distributed-weight path is lost again. Backward can then fail when it tries to use the unsaved gathered weights, or skip distributed gradient finalization. Keep the original weights available for each backward while the graph is retained.
|
/te-ci pytorch L0 |
Description
_Linear and _GroupedLinear save the DistributedWeight parameter itself for
backward and gate their distributed-weight path on getting that object back
(is_distributed_weight(saved_weight)). That holds only while no
torch.autograd.graph.saved_tensors_hooks are installed: with hooks active,
SavedVariable::unpack rebuilds a fresh plain tensor from the unpack hook's
result, dropping the Python subclass and its is_distributed_weight marker.
The backward then takes the plain-parameter branch, whose weakrefs point at
the transient all-gathered weights, and fails with
"weight was removed while fuse_wgrad_accumulation=True".
Seen with Megatron-Core fine-grained activation offloading (which installs
such hooks around expert fc1/act) on bf16 GTP-sharded GroupedLinear experts.
Quantized weights were unaffected because their storage objects already stay
on ctx via prepare_for_saving/restore_from_saved, and the op-fuser path reads
the module's own weight attribute instead of the saved tensor.
Keep the DistributedWeight objects as ordinary ctx / LinearBwdArgs references
and prefer them over the saved tensors in backward.
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: