Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 43 additions & 3 deletions tests/pytorch/test_grouped_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,10 @@

"""Tests for GroupedTensor class"""

from typing import List, Optional, Tuple
import os
from types import SimpleNamespace
from typing import List, Optional, Tuple

import pytest
import torch
import transformer_engine.pytorch as te
Expand All @@ -18,8 +20,7 @@
MXFP8Quantizer,
NVFP4Quantizer,
)
from transformer_engine.pytorch.constants import TE_DType_To_Torch
from transformer_engine.pytorch.utils import is_non_tn_fp8_gemm_supported
from transformer_engine.pytorch.utils import is_non_tn_fp8_gemm_supported, mark_grouped_tensor
import transformer_engine_torch as tex

# Import test utilities
Expand Down Expand Up @@ -48,6 +49,45 @@
)
)


def test_mark_grouped_tensor_supports_plain_tensor():
tensor = torch.empty(16)

mark_grouped_tensor(tensor)

assert tensor.grouped_tensor_scale_inv is False


def test_mark_grouped_tensor_supports_unquantized_rowwise_storage():
rowwise_data = torch.empty(16)
grouped_tensor = SimpleNamespace(
rowwise_data=rowwise_data,
columnwise_data=None,
columnwise_scale_inv=None,
)

mark_grouped_tensor(grouped_tensor)

assert rowwise_data.grouped_tensor_scale_inv is False


def test_mark_grouped_tensor_marks_quantized_columnwise_storage():
rowwise_data = torch.empty(16)
columnwise_data = torch.empty(16)
columnwise_scale_inv = torch.empty(4)
grouped_tensor = SimpleNamespace(
rowwise_data=rowwise_data,
columnwise_data=columnwise_data,
columnwise_scale_inv=columnwise_scale_inv,
)

mark_grouped_tensor(grouped_tensor)

assert not hasattr(rowwise_data, "grouped_tensor_scale_inv")
assert columnwise_data.grouped_tensor_scale_inv is False
assert columnwise_scale_inv.grouped_tensor_scale_inv is True


_quantization_params = [
pytest.param(
"fp8_delayed_scaling",
Expand Down
5 changes: 5 additions & 0 deletions transformer_engine/pytorch/module/grouped_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
clear_tensor_data,
get_device_compute_capability,
init_method_constant,
mark_grouped_tensor,
requires_grad,
resolve_grouped_linear_single_param_flags,
get_nvtx_range_context,
Expand Down Expand Up @@ -553,6 +554,10 @@ def _forward_grouped_tensor(
if not inp.requires_grad:
weights_to_save = [None] * len(weights_to_save)

# Megatron-LM paged stashing uses this marker to identify the dynamic activation
# buffers among the tensors saved by the GroupedLinear autograd function. The
# operation-fuser grouped MLP applies the same marker to its saved activations.
mark_grouped_tensor(input_to_save)
tensors_to_save, tensor_objects = prepare_for_saving(
input_to_save,
*weights_to_save,
Expand Down
38 changes: 28 additions & 10 deletions transformer_engine/pytorch/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -311,21 +311,39 @@ def mark_grouped_tensor(*tensors: List[Any]):
Megatron-LM to detect which tensors are dynamic (varying shapes)
and remove the padding before doing the `save_for_backward` to
save memory.
Note: Only columnwise data is saved for backward."""

Plain tensors are saved directly. Grouped tensors are decomposed by
`prepare_for_saving`: unquantized BF16/FP16 activations save their
rowwise data, while quantized activations save their columnwise data
and scale metadata.
"""
for tensor in tensors:
if tensor is None:
continue
if hasattr(tensor, "columnwise_data"):
assert (
tensor.columnwise_data is not None
), "Columnwise data is not set for grouped tensor"

if not hasattr(tensor, "columnwise_data"):
# Plain tensor, e.g. a fused-MLP activation input.
setattr(tensor, "grouped_tensor_scale_inv", False)
continue

# Grouped tensor: mark the underlying tensors that `prepare_for_saving`
# will pass to `save_for_backward`, rather than the storage wrapper.
if tensor.columnwise_data is None:
# Unquantized BF16/FP16 grouped tensor.
saved_activation = tensor.rowwise_data
saved_scale_inv = None
else:
# Quantized grouped tensor saved in the representation used by wgrad.
saved_activation = tensor.columnwise_data
saved_scale_inv = tensor.columnwise_scale_inv
assert (
tensor.columnwise_scale_inv is not None
saved_scale_inv is not None
), "Columnwise scale inverse is not set for grouped tensor"
setattr(tensor.columnwise_data, "grouped_tensor_scale_inv", False)
setattr(tensor.columnwise_scale_inv, "grouped_tensor_scale_inv", True)
else:
setattr(tensor, "grouped_tensor_scale_inv", False)

assert saved_activation is not None, "Grouped tensor has no activation data"
setattr(saved_activation, "grouped_tensor_scale_inv", False)
if saved_scale_inv is not None:
setattr(saved_scale_inv, "grouped_tensor_scale_inv", True)


def split_tensor_along_dim(
Expand Down
Loading