Skip to content

[PyTorch] Keep DistributedWeight objects on ctx across saved-tensor hooks - #3545

Open
xrennvidia wants to merge 7 commits into
NVIDIA:mainfrom
xrennvidia:xren/dist-weight-saved-tensor-hooks
Open

xrennvidia wants to merge 7 commits into
NVIDIA:mainfrom
xrennvidia:xren/dist-weight-saved-tensor-hooks

Conversation

@xrennvidia

@xrennvidia xrennvidia commented Sep 18, 2026 •

Copy link
Copy Markdown
Collaborator

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

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Change A
  • Change B

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 18, 2026
@greptile-apps

greptile-apps Bot commented Sep 18, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 4/5

[Medium risk] Fixes distributed weight handling under PyTorch saved-tensor hooks.

The PR is not yet safe to merge because a second backward on a retained graph can lose the distributed weights.

Findings

  1. P1 Second backward loses distributed weights ▶

Summary

The PR keeps DistributedWeight objects available to Linear and GroupedLinear backward when saved-tensor hooks unpack saved weights as plain tensors.

  • Adds regression tests for both modules and their fused weight-gradient paths.
  • Adds the tests to the PyTorch unit-test run.

Reviews (9) · Last reviewed commit: "Merge branch 'main' into xren/dist-weigh..."

Comment thread transformer_engine/pytorch/module/grouped_linear.py
xrennvidia and others added 2 commits September 18, 2026 02:16
…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>
@xrennvidia
xrennvidia force-pushed the xren/dist-weight-saved-tensor-hooks branch from 45d5cb7 to eea3da3 Compare September 18, 2026 09:16
@fanshiqing

Copy link
Copy Markdown
Member

LGTM, thanks for the fix!

@fanshiqing fanshiqing linked an issue Sep 23, 2026 that may be closed by this pull request
@ksivaman

Copy link
Copy Markdown
Member

@xrennvidia Can you rebase this with main? #3517 has been merged

…d-tensor-hooks

# Conflicts:
#	transformer_engine/pytorch/module/grouped_linear.py
@xrennvidia

Copy link
Copy Markdown
Collaborator Author

@xrennvidia Can you rebase this with main? #3517 has been merged

Done.

@greptile-apps

This comment has been minimized.

Comment thread tests/pytorch/test_module_distributed_weight_saved_tensor_hooks.py
@greptile-apps

greptile-apps Bot commented Sep 28, 2026

Copy link
Copy Markdown
Contributor

Want your agent to iterate on Greptile's feedback? Try greploops.

xrennvidia and others added 2 commits September 28, 2026 06:32
…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>
@xrennvidia
xrennvidia force-pushed the xren/dist-weight-saved-tensor-hooks branch from 9d8b650 to dddf418 Compare September 28, 2026 13:32
dist_weights = getattr(ctx, "dist_weights", None)
if dist_weights is not None:
weight_tensors = dist_weights
ctx.dist_weights = None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

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

@ksivaman

Copy link
Copy Markdown
Member

/te-ci pytorch L0

This branch has not been deployed

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

Labels

community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

GTP+TE integration

3 participants