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
98 changes: 98 additions & 0 deletions tests/pytorch/test_ptq_calibration_metadata_buffering.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,98 @@
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.

from types import SimpleNamespace

import pytest
import torch

from transformer_engine.pytorch.module import _common
from transformer_engine.pytorch.module import grouped_linear


@pytest.mark.parametrize(
("recipe", "metadata_name", "expected_value"),
(
("fp8_current_scaling", "scale_inv", 0.25),
("fp8_delayed_scaling", "amax", 448.0),
("nvfp4", "amax", 2688.0),
("nvfp4_rowwise", "amax_rowwise", 1344.0),
),
)
def test_scale_buffer_info_selects_recipe_metadata(
monkeypatch, recipe, metadata_name, expected_value
):
monkeypatch.setattr(_common, "get_quantization_recipe_name", lambda _: recipe)
tensor = SimpleNamespace(
_scale_inv=torch.tensor([0.25], dtype=torch.float32),
_amax_rowwise=torch.tensor([2688.0 if recipe == "nvfp4" else 1344.0], dtype=torch.float32),
)
quantizer = SimpleNamespace(amax=torch.tensor([448.0], dtype=torch.float32))

buffer_name, value = _common._get_scale_buffer_info("input", tensor, quantizer)

assert buffer_name == f"input_tensor_{metadata_name}_{recipe}_te_ptq_calibrated"
torch.testing.assert_close(value, torch.tensor([expected_value]))


@pytest.mark.parametrize("recipe", ("mxfp8", "fp8_block_scaling"))
def test_scale_buffer_info_skips_non_global_scaling_recipes(monkeypatch, recipe):
monkeypatch.setattr(_common, "get_quantization_recipe_name", lambda _: recipe)
tensor = SimpleNamespace(_rowwise_scale_inv=torch.ones(2, 2))

assert _common._get_scale_buffer_info("input", tensor, object()) is None


def test_grouped_scale_buffers_are_per_gemm(monkeypatch):
monkeypatch.setattr(_common, "get_quantization_recipe_name", lambda _: "fp8_current_scaling")
inputs = [
SimpleNamespace(_scale_inv=torch.tensor([0.25])),
SimpleNamespace(_scale_inv=torch.tensor([0.5])),
]
weights = [
SimpleNamespace(_scale_inv=torch.tensor([0.75])),
SimpleNamespace(_scale_inv=torch.tensor([1.0])),
]
scale_buffers = {}

grouped_linear._update_grouped_scale_buffers(
scale_buffers,
inputs,
weights,
object(),
object(),
activation_scale_decay=0.0,
)

assert set(scale_buffers) == {
"input_gemm0_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated",
"input_gemm1_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated",
"weight_gemm0_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated",
"weight_gemm1_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated",
}
torch.testing.assert_close(
scale_buffers["input_gemm1_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated"],
torch.tensor([0.5]),
)


@pytest.mark.parametrize(
("observed_scale", "expected_scale"),
(
# Decayed max is greater than the observed.
(1.0, 2.0),
# Decayed max is less than the observed.
(3.0, 3.0),
),
)
def test_activation_scale_buffer_uses_decaying_maximum(observed_scale, expected_scale):
name = "fc1_input_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated"
scale_buffers = {name: torch.tensor([4.0])}

_common._update_scale_buffers(
scale_buffers,
{name: torch.tensor([observed_scale])},
activation_scale_decay=0.5,
)
torch.testing.assert_close(scale_buffers[name], torch.tensor([expected_scale]))
29 changes: 4 additions & 25 deletions transformer_engine/debug/features/log_fp8_tensor_stats.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,35 +23,12 @@
)
from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer
from transformer_engine.pytorch.tensor.float8_blockwise_tensor import Float8BlockQuantizer

try:
from transformer_engine.pytorch.tensor.nvfp4_tensor import NVFP4Quantizer

_nvfp4_available = True
except ImportError:
_nvfp4_available = False
NVFP4Quantizer = None
from transformer_engine.pytorch.tensor.utils import get_quantization_recipe_name


ALL_RECIPE_NAMES = ["fp8_delayed_scaling", "fp8_current_scaling", "mxfp8", "fp8_block_scaling"]


def _get_recipe_name(quantizer: Optional[Quantizer]):
if quantizer is None:
return ""
if isinstance(quantizer, Float8Quantizer):
return "fp8_delayed_scaling"
if isinstance(quantizer, Float8CurrentScalingQuantizer):
return "fp8_current_scaling"
if isinstance(quantizer, MXFP8Quantizer):
return "mxfp8"
if isinstance(quantizer, Float8BlockQuantizer):
return "fp8_block_scaling"
if _nvfp4_available and isinstance(quantizer, NVFP4Quantizer):
return "nvfp4"
raise ValueError(f"Unsupported quantizer type: {type(quantizer)}")


def _get_new_quantizer(recipe_name, fp8_dtype):
if recipe_name == "fp8_block_scaling":
return Float8BlockQuantizer(fp8_dtype=fp8_dtype, rowwise=True, columnwise=True)
Expand Down Expand Up @@ -336,7 +313,9 @@ def inspect_tensor(
)
return

recipe_name = _get_recipe_name(quantizer)
recipe_name = get_quantization_recipe_name(quantizer)
if recipe_name == "nvfp4_rowwise":
recipe_name = "nvfp4"

for stat in config["stats"]:
self.check_if_stat_is_supported(
Expand Down
77 changes: 76 additions & 1 deletion transformer_engine/pytorch/module/_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,16 +6,91 @@

import dataclasses
import queue
from typing import Any, Callable, List, Optional, Tuple, Union
from typing import Any, Callable, Dict, List, Optional, Tuple, Union

import torch

from .. import cpp_extensions as tex
from ..constants import TE_DType
from ..export import is_in_onnx_export_mode
from ..tensor.utils import get_quantization_recipe_name
from ..utils import get_default_init_method


def _get_scale_buffer_info(
tensor_name: str,
tensor: Any,
quantizer: Any,
) -> Optional[Tuple[str, Optional[torch.Tensor]]]:
"""Get the calibration buffer name and value for a quantized tensor."""
recipe = get_quantization_recipe_name(quantizer)
if not recipe:
return None

if recipe == "fp8_delayed_scaling":
metadata_name = "amax"
metadata = getattr(quantizer, "amax", None)
elif recipe == "fp8_current_scaling":
metadata_name = "scale_inv"
metadata = getattr(tensor, "_scale_inv", None)
elif recipe == "nvfp4":
metadata_name = "amax"
metadata = getattr(tensor, "_amax_rowwise", None)
elif recipe == "nvfp4_rowwise":
metadata_name = "amax_rowwise"
metadata = getattr(tensor, "_amax_rowwise", None)
elif recipe == "mxfp8":
# MXFP8 only exposes blockwise E8M0-encoded inverse scales, not a
# global FP32 scaling factor suitable for PTQ checkpoint export.
return None
elif recipe == "fp8_block_scaling":
# FP8 block scaling only exposes blockwise inverse scales, not a
# global FP32 scaling factor suitable for PTQ checkpoint export.
return None
else:
raise ValueError(f"Unsupported quantization recipe {recipe!r}")

buffer_name = f"{tensor_name}_tensor_{metadata_name}_{recipe}_te_ptq_calibrated"
return buffer_name, metadata


def _update_scale_buffers(
scale_buffers: Dict[str, Optional[torch.Tensor]],
scale_updates: Dict[str, Optional[torch.Tensor]],
activation_scale_decay: float = 0.0,
) -> None:
"""Merge observed scaling factors into checkpoint buffers."""
for buffer_name, scale in scale_updates.items():
if scale is None:
continue
if activation_scale_decay > 0.0:
observed_scale = scale.detach().float()
Comment thread
greptile-apps[bot] marked this conversation as resolved.
scale_buffer = scale_buffers.get(buffer_name)
if scale_buffer is not None and scale_buffer.shape != observed_scale.shape:
raise RuntimeError(
"Quantized scaling-factor buffer shape changed from "
f"{tuple(scale_buffer.shape)} to {tuple(observed_scale.shape)}"
)
if scale_buffer is None:
# Initialize the rolling activation scaling factor.
# Requires CUDA graph warmup step.
scale_buffer = torch.zeros_like(observed_scale)
scale_buffers[buffer_name] = scale_buffer
# Track a decaying maximum so early-training activation
# outliers do not permanently determine the inference scale.
scale_buffer.mul_(activation_scale_decay)
torch.maximum(
scale_buffer,
observed_scale,
out=scale_buffer,
)
else:
# Without scale history, keep a reference to the current metadata
# without allocating or copying a separate buffer.
# Requires CUDA graph warmup step.
scale_buffers[buffer_name] = scale.detach()


def set_quantizer_amax_reduction_group(quantizer, amax_reduction_group) -> None:
"""Set the amax reduction group on a quantizer; no-op if it doesn't support it.

Expand Down
Loading
Loading