Repository navigation
GroupedTensor Machinery for MXFP8 Dispatch input passed from Megatron - #3660
vthumbe1503 wants to merge 2 commits into
Conversation
Signed-off-by: vthumbe1503 <vthumbe@nvidia.com>
|
| def register_alias(self, tensor: "GroupedTensor") -> None: | ||
| """Register a wrapper that shares this storage's payload allocations.""" | ||
| self._aliases.append(weakref.ref(tensor)) |
There was a problem hiding this comment.
Dead aliases keep accumulating
register_alias adds a weak reference for every new wrapper, but dead references are removed only during resize_(0). Repeated detach() calls on a long-lived grouped parameter keep growing this list after the detached wrappers are gone. This wastes host memory and makes later resizing walk a larger list. Remove dead entries when aliases expire or when registering new aliases.
| result = Identity.apply(grouped_tensor) | ||
|
|
||
| assert isinstance(result, GroupedTensor) | ||
| assert result is not grouped_tensor | ||
| assert result._base is grouped_tensor | ||
| assert result.grad_fn is not None | ||
| assert result.rowwise_data is grouped_tensor.rowwise_data | ||
| assert result.first_dims is grouped_tensor.first_dims | ||
| assert result.tensor_offsets is grouped_tensor.tensor_offsets |
There was a problem hiding this comment.
Backward behavior is not tested
test_autograd_identity_supports_varying_shapes checks that grad_fn exists but never runs backward. This PR also changes how alias wrappers set requires_grad, so the test would miss a failure when gradients pass through the new view. Run backward with an explicit gradient and check that the input receives the expected gradient.
Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: