From 8869e4647828c772edc75c3e0271a82b6a66702b Mon Sep 17 00:00:00 2001 From: Cory Ye Date: Sun, 26 Jul 2026 21:06:13 -0700 Subject: [PATCH 1/5] Add optional tooling to accumulate quantization scaling factors for inference. Signed-off-by: Cory Ye --- .../debug/features/log_fp8_tensor_stats.py | 29 +------ transformer_engine/pytorch/module/_common.py | 67 +++++++++++++- .../pytorch/module/grouped_linear.py | 87 ++++++++++++++++++- .../pytorch/module/layernorm_linear.py | 61 ++++++++++++- .../pytorch/module/layernorm_mlp.py | 84 +++++++++++++++++- transformer_engine/pytorch/module/linear.py | 79 ++++++++++++++++- transformer_engine/pytorch/tensor/utils.py | 20 +++++ 7 files changed, 392 insertions(+), 35 deletions(-) diff --git a/transformer_engine/debug/features/log_fp8_tensor_stats.py b/transformer_engine/debug/features/log_fp8_tensor_stats.py index 96f1b644cf..e05a90f0d8 100644 --- a/transformer_engine/debug/features/log_fp8_tensor_stats.py +++ b/transformer_engine/debug/features/log_fp8_tensor_stats.py @@ -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) @@ -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( diff --git a/transformer_engine/pytorch/module/_common.py b/transformer_engine/pytorch/module/_common.py index 9b914e2d6c..0eddee7523 100644 --- a/transformer_engine/pytorch/module/_common.py +++ b/transformer_engine/pytorch/module/_common.py @@ -6,16 +6,81 @@ 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_current_scaling": + metadata_name = "scale_inv" + metadata = getattr(tensor, "_scale_inv", None) + else: + metadata_name = "amax_rowwise" + metadata = getattr( + tensor, + "_amax_rowwise", + getattr(quantizer, "amax", None), + ) + 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 buffer_name.startswith("input"): + observed_scale = scale.detach().float() + 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 activation_scale_decay == 0.0: + # If not using scale decay, just buffer the current scaling. + scale_buffers[buffer_name] = observed_scale + continue + 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: + # Keep a reference to the current weight 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. diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index 0b51ca6a1e..b5ae45e9b0 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -4,7 +4,7 @@ """GroupedLinear API""" -from typing import Union, Optional, Callable, Tuple, List +from typing import Union, Optional, Callable, Tuple, List, Dict from itertools import chain import os import warnings @@ -31,7 +31,7 @@ _clear_high_precision_init_val, _get_high_precision_init_val, ) -from ._common import WeightGradStore +from ._common import _get_scale_buffer_info, _update_scale_buffers, WeightGradStore from ..quantization import FP8GlobalStateManager, QuantizerRole from ..utils import ( divide, @@ -83,6 +83,29 @@ __all__ = ["GroupedLinear"] +def _update_grouped_scale_buffers( + scale_buffers: Dict[str, Optional[torch.Tensor]], + input_tensors: List[Union[torch.Tensor, QuantizedTensorStorage]], + weight_tensors: List[Union[torch.Tensor, QuantizedTensorStorage]], + input_quantizer: Optional[Quantizer], + weight_quantizer: Optional[Quantizer], + activation_scale_decay: float, +) -> None: + """Update GroupedLinear PTQ calibration buffers with per-GEMM metadata.""" + scale_updates = {} + for index, tensor in enumerate(input_tensors): + scale_buffer = _get_scale_buffer_info(f"input_gemm{index}", tensor, input_quantizer) + if scale_buffer is not None: + scale_updates[scale_buffer[0]] = scale_buffer[1] + for index, tensor in enumerate(weight_tensors): + scale_buffer = _get_scale_buffer_info( + f"weight_gemm{index}", tensor, weight_quantizer + ) + if scale_buffer is not None: + scale_updates[scale_buffer[0]] = scale_buffer[1] + _update_scale_buffers(scale_buffers, scale_updates, activation_scale_decay) + + class _GroupedLinear(torch.autograd.Function): """GroupedLinear semi-top level module Calls custom cuda extensions. @@ -326,6 +349,8 @@ def _forward_grouped_tensor( weight_workspaces: List[Optional[QuantizedTensorStorage]], cache_weight: bool, skip_fp8_weight_update: Optional[torch.Tensor], + scale_buffers: Optional[Dict[str, Optional[torch.Tensor]]], + quantized_scaling_factor_buffering_decay: float, weights: Tuple[torch.Tensor, ...], biases: Tuple[torch.Tensor, ...], out: Optional[torch.Tensor] = None, @@ -414,6 +439,19 @@ def _forward_grouped_tensor( use_split_accumulator=use_split_accumulator, ) + if scale_buffers is not None: + grouped_inputs = grouped_x.quantized_tensors + if grouped_inputs is None: + grouped_inputs = grouped_x.split_into_quantized_tensors() + _update_grouped_scale_buffers( + scale_buffers, + grouped_inputs, + weights_for_gemm, + input_quantizers[0], + weight_quantizers[0], + quantized_scaling_factor_buffering_decay, + ) + if is_grad_enabled: if weight_requires_grad: # (For FP8 per tensor current scaling on Hopper --> Free Rowwise Data @@ -517,6 +555,8 @@ def forward( skip_fp8_weight_update, save_original_input, debug, + scale_buffers, + quantized_scaling_factor_buffering_decay, ) = non_tensor_args if fp8: backward_override = FP8GlobalStateManager.get_fp8_recipe().backward_override @@ -623,6 +663,10 @@ def forward( weight_workspaces=weight_workspaces, cache_weight=cache_weight, skip_fp8_weight_update=skip_fp8_weight_update, + scale_buffers=scale_buffers, + quantized_scaling_factor_buffering_decay=( + quantized_scaling_factor_buffering_decay + ), weights=weights, biases=biases, out=out, @@ -675,6 +719,16 @@ def forward( else: weights_fp8 = [cast_if_needed(weight, activation_dtype) for weight in weights] + if scale_buffers is not None: + _update_grouped_scale_buffers( + scale_buffers, + inputmats, + weights_fp8, + input_quantizers[0], + weight_quantizers[0], + quantized_scaling_factor_buffering_decay, + ) + # Initialize biases bias_dtype = activation_dtype if fp8 and activation_dtype == torch.float32: @@ -1370,6 +1424,10 @@ class GroupedLinear(TransformerEngineBaseModule): cast tensor. In some scenarios, the input tensor is used by multiple modules, and saving the original input tensor may reduce the memory usage. Cannot work with FP8 DelayedScaling recipe. + buffer_quantized_scaling_factors : bool, default = False + If set to ``True``, maintain nonpersistent input and weight quantization + metadata buffers for inference checkpoint export. Buffers store metadata + per grouped GEMM, using inverse scales directly for FP8 current scaling. single_grouped_weight : bool, default = False If set to ``True``, grouped weights are stored as a single grouped parameter instead of one parameter per GEMM. @@ -1415,6 +1473,8 @@ def __init__( single_grouped_weight: bool = False, single_grouped_bias: bool = False, name: Optional[str] = None, + buffer_quantized_scaling_factors: bool = False, + quantized_scaling_factor_buffering_decay: float = 0.0, ) -> None: super().__init__(name) @@ -1430,6 +1490,10 @@ def __init__( self.ub_overlap_ag = ub_overlap_ag self.ub_name = ub_name self.save_original_input = save_original_input + self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors + self.quantized_scaling_factor_buffering_decay = ( + quantized_scaling_factor_buffering_decay + ) single_grouped_weight, single_grouped_bias = resolve_grouped_linear_single_param_flags( single_grouped_weight, single_grouped_bias ) @@ -1946,6 +2010,13 @@ def forward( if cache_weight else [None] * num_gemms ) + scale_buffers = None + if self.buffer_quantized_scaling_factors: + scale_buffers = { + name: value + for name, value in self._buffers.items() + if name.endswith("_te_ptq_calibrated") + } non_tensor_args = ( self.apply_bias, @@ -1969,6 +2040,8 @@ def forward( skip_fp8_weight_update, self.save_original_input, debug, + scale_buffers, + self.quantized_scaling_factor_buffering_decay, ) out, new_workspaces = linear_fn( *autograd_ctx, @@ -1981,6 +2054,16 @@ def forward( *bias_tensors, ) + if scale_buffers is not None: + # Assign scaling-factor calibration buffers to the model. + # Materializing a new buffer requires a CUDA graph warmup step. + for name, value in scale_buffers.items(): + if value is not None: + if name in self._buffers: + setattr(self, name, value) + else: + self.register_buffer(name, value, persistent=False) + if cache_weight: for i, ws in enumerate(new_workspaces): if ws is not None: diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index a588e21a0c..b317433d3e 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -17,7 +17,10 @@ from transformer_engine.common.recipe import Recipe from transformer_engine.pytorch.torch_version import torch_version -from transformer_engine.pytorch.tensor.utils import clear_columnwise_cache, is_custom +from transformer_engine.pytorch.tensor.utils import ( + clear_columnwise_cache, + is_custom, +) from .base import ( fill_userbuffers_buffer_for_all_gather, get_ub, @@ -66,6 +69,8 @@ from ..jit import no_torch_dynamo from ..graph import is_graph_capturing from ._common import ( + _get_scale_buffer_info, + _update_scale_buffers, apply_normalization, noop_cat, set_quantizer_amax_reduction_group, @@ -157,6 +162,8 @@ def forward( symmetric_ar_type, debug, is_fsdp2, + quantized_scaling_factor_buffering_decay, + scale_buffers, ) = non_tensor_args if fp8: backward_override = FP8GlobalStateManager.get_fp8_recipe().backward_override @@ -376,6 +383,24 @@ def forward( if weight_quantizer is not None: weight_quantizer.calibrate(weight) + if scale_buffers is not None: + scale_updates = {} + input_scale_buffer = _get_scale_buffer_info( + "input", ln_out_total, input_quantizer + ) + if input_scale_buffer is not None: + scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] + weight_scale_buffer = _get_scale_buffer_info( + "weight", weightmat, weight_quantizer + ) + if weight_scale_buffer is not None: + scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] + _update_scale_buffers( + scale_buffers, + scale_updates, + quantized_scaling_factor_buffering_decay, + ) + # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = _2X_ACC_FPROP if fp8: @@ -1299,6 +1324,15 @@ class LayerNormLinear(TransformerEngineBaseModule): This can help in latency bound communication situations. Requires PyTorch version 2.7.0 or higher. When set to ``None``, standard all-reduce is used. + buffer_quantized_scaling_factors : bool, default = False + If set to ``True``, maintain nonpersistent input and weight quantization + metadata buffers for inference checkpoint export. Per-tensor buffers + store raw global amaxes, except FP8 current scaling buffers, which store + inverse scales directly. + quantized_scaling_factor_buffering_decay : float, default = 0.0 + Decay applied to buffered activation scaling factors before incorporating + each new observation. Defaults to 0.0, in which case only the most recent + scaling factor is buffered. """ def __init__( @@ -1331,6 +1365,8 @@ def __init__( delay_wgrad_compute: bool = False, symmetric_ar_type: Optional[str] = None, name: Optional[str] = None, + buffer_quantized_scaling_factors: bool = False, + quantized_scaling_factor_buffering_decay: float = 0.0, ) -> None: super().__init__(name) @@ -1593,6 +1629,11 @@ def __init__( if name in self.weight_names or name in self.bias_names: param.skip_backward_post_hook = True + self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors + self.quantized_scaling_factor_buffering_decay = ( + quantized_scaling_factor_buffering_decay + ) + def set_meta_tensor(self, fwd: bool, recipe: Recipe) -> None: """Init scales and amaxes for fwd | bwd.""" super().set_meta_tensor(fwd, recipe) @@ -1763,6 +1804,14 @@ def forward( self._fp8_workspaces.get(cache_name) if cache_name is not None else None ) + scale_buffers = None + if self.buffer_quantized_scaling_factors: + scale_buffers = { + name: value + for name, value in self._buffers.items() + if name.endswith("_te_ptq_calibrated") + } + non_tensor_args = ( self.eps, is_first_microbatch, @@ -1803,6 +1852,8 @@ def forward( self.symmetric_ar_type, debug, self.is_fsdp2, + self.quantized_scaling_factor_buffering_decay, + scale_buffers, ) out, ln_out, new_weight_workspace = fwd_fn( *autograd_ctx, @@ -1815,6 +1866,14 @@ def forward( non_tensor_args, ) + if scale_buffers is not None: + for name, value in scale_buffers.items(): + if value is not None: + if name in self._buffers: + setattr(self, name, value) + else: + self.register_buffer(name, value, persistent=False) + if new_weight_workspace is not None and cache_name is not None: if isinstance(new_weight_workspace, torch.Tensor): new_weight_workspace = new_weight_workspace.detach() diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 19e20d775c..7b3de3ef66 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -18,7 +18,10 @@ from transformer_engine.common.recipe import Recipe from transformer_engine.pytorch.torch_version import torch_version -from transformer_engine.pytorch.tensor.utils import clear_columnwise_cache, is_custom +from transformer_engine.pytorch.tensor.utils import ( + clear_columnwise_cache, + is_custom, +) from .base import ( fill_userbuffers_buffer_for_all_gather, _ub_communicators, @@ -74,7 +77,13 @@ from ..tensor.mxfp8_tensor import MXFP8Quantizer from ..tensor.nvfp4_tensor import NVFP4Quantizer from ..tensor.float8_blockwise_tensor import Float8BlockQuantizer -from ._common import apply_normalization, set_quantizer_amax_reduction_group, WeightGradStore +from ._common import ( + _get_scale_buffer_info, + _update_scale_buffers, + apply_normalization, + set_quantizer_amax_reduction_group, + WeightGradStore, +) from ..cpu_offload import ( is_cpu_offload_enabled, start_offload, @@ -241,6 +250,8 @@ def _forward( checkpoint, debug, is_fsdp2, + quantized_scaling_factor_buffering_decay, + scale_buffers, recompute_for_bwd, ) = non_tensor_args if fp8: @@ -343,6 +354,10 @@ def _forward( "checkpoint": checkpoint, "debug": debug, "is_fsdp2": is_fsdp2, + "quantized_scaling_factor_buffering_decay": ( + quantized_scaling_factor_buffering_decay + ), + "scale_buffers": scale_buffers, "recompute_for_bwd": True, # set this to true for recomputation phase } # Make sure input dimensions are compatible @@ -557,6 +572,16 @@ def _forward( if fc1_weight_quantizer is not None: fc1_weight_quantizer.calibrate(fc1_weight) + fc1_input_scale_buffer = None + fc1_weight_scale_buffer = None + if scale_buffers is not None: + fc1_input_scale_buffer = _get_scale_buffer_info( + "fc1_input", ln_out_total, fc1_input_quantizer + ) + fc1_weight_scale_buffer = _get_scale_buffer_info( + "fc1_weight", fc1_weight_final, fc1_weight_quantizer + ) + # ------------------------------------------------------ # FC1 GEMM # ------------------------------------------------------ @@ -671,6 +696,28 @@ def _forward( if fc2_weight_quantizer is not None: fc2_weight_quantizer.calibrate(fc2_weight) + if scale_buffers is not None: + scale_updates = {} + if fc1_input_scale_buffer is not None: + scale_updates[fc1_input_scale_buffer[0]] = fc1_input_scale_buffer[1] + if fc1_weight_scale_buffer is not None: + scale_updates[fc1_weight_scale_buffer[0]] = fc1_weight_scale_buffer[1] + fc2_input_scale_buffer = _get_scale_buffer_info( + "fc2_input", act_out, fc2_input_quantizer + ) + if fc2_input_scale_buffer is not None: + scale_updates[fc2_input_scale_buffer[0]] = fc2_input_scale_buffer[1] + fc2_weight_scale_buffer = _get_scale_buffer_info( + "fc2_weight", fc2_weight_final, fc2_weight_quantizer + ) + if fc2_weight_scale_buffer is not None: + scale_updates[fc2_weight_scale_buffer[0]] = fc2_weight_scale_buffer[1] + _update_scale_buffers( + scale_buffers, + scale_updates, + quantized_scaling_factor_buffering_decay, + ) + # Configure Userbuffers reduce-scatter if needed ub_obj_fc2out = None reduce_scatter_out = None @@ -1939,6 +1986,15 @@ class LayerNormMLP(TransformerEngineBaseModule): whether to use selective activation checkpointing, where activations are not saved for bwd, and instead are recomputed (skipping fc2, as it is not needed for backward). Trades compute for memory. default is false, in which activations are saved in fwd. not supported for onnx forward + buffer_quantized_scaling_factors : bool, default = False + If set to ``True``, maintain nonpersistent activation and weight quantization + metadata buffers for both internal linear layers for inference checkpoint + export. Per-tensor buffers store raw global amaxes, except FP8 current + scaling buffers, which store inverse scales directly. + quantized_scaling_factor_buffering_decay : float, default = 0.0 + Decay applied to buffered activation scaling factors before incorporating + each new observation. Defaults to 0.0, in which case only the most recent + scaling factor is buffered. """ def __init__( @@ -1975,6 +2031,8 @@ def __init__( delay_wgrad_compute: bool = False, symmetric_ar_type: Optional[str] = None, checkpoint: bool = False, + buffer_quantized_scaling_factors: bool = False, + quantized_scaling_factor_buffering_decay: float = 0.0, ) -> None: super().__init__(name) @@ -1996,6 +2054,10 @@ def __init__( self.zero_centered_gamma = zero_centered_gamma self.symmetric_ar_type = symmetric_ar_type self.checkpoint = checkpoint + self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors + self.quantized_scaling_factor_buffering_decay = ( + quantized_scaling_factor_buffering_decay + ) # GEMM-GELU fusion is currently only supported with split GEMM-AG overlap self.gemm_gelu_fusion = ( @@ -2376,6 +2438,14 @@ def forward( self._fp8_workspaces.get(cache_name_fc2) if cache_name_fc2 is not None else None ) + scale_buffers = None + if self.buffer_quantized_scaling_factors: + scale_buffers = { + name: value + for name, value in self._buffers.items() + if name.endswith("_te_ptq_calibrated") + } + non_tensor_args = ( self.eps, is_first_microbatch, @@ -2426,6 +2496,8 @@ def forward( self.checkpoint, debug, self.is_fsdp2, + self.quantized_scaling_factor_buffering_decay, + scale_buffers, ) out, ln_out, new_fc1_ws, new_fc2_ws = fwd_fn( *autograd_ctx, @@ -2441,6 +2513,14 @@ def forward( non_tensor_args, ) + if scale_buffers is not None: + for name, value in scale_buffers.items(): + if value is not None: + if name in self._buffers: + setattr(self, name, value) + else: + self.register_buffer(name, value, persistent=False) + if new_fc1_ws is not None and cache_name_fc1 is not None: if isinstance(new_fc1_ws, torch.Tensor): new_fc1_ws = new_fc1_ws.detach() diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 6b0f941bc4..049ede7fa7 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -30,7 +30,13 @@ _2X_ACC_DGRAD, _2X_ACC_WGRAD, ) -from ._common import noop_cat, set_quantizer_amax_reduction_group, WeightGradStore +from ._common import ( + _get_scale_buffer_info, + _update_scale_buffers, + noop_cat, + set_quantizer_amax_reduction_group, + WeightGradStore, +) from ..quantization import FP8GlobalStateManager, QuantizerRole from ..utils import ( cast_if_needed, @@ -76,7 +82,10 @@ ) from ..tensor.float8_tensor import Float8CurrentScalingQuantizer, Float8Quantizer from ..tensor.mxfp8_tensor import MXFP8Quantizer -from ..tensor.utils import clear_columnwise_cache, is_custom +from ..tensor.utils import ( + clear_columnwise_cache, + is_custom, +) from ..export import is_in_onnx_export_mode, assert_warmed_up from ..cpu_offload import ( is_cpu_offload_enabled, @@ -160,6 +169,10 @@ class LinearFwdArgs: fuse_wgrad_accumulation: bool wgrad_store: Optional[Any] + # Inference Scaling Factor Calibration Buffering + scale_buffers: Optional[Dict[str, Optional[torch.Tensor]]] + quantized_scaling_factor_buffering_decay: float + # --- Misc --- cpu_offloading: bool is_grad_enabled: bool @@ -266,8 +279,8 @@ def _linear_forward_impl( Returns ``(out, new_weight_workspace, tensors_to_save_from_forward, None, ctx_attrs)``. ``new_weight_workspace`` is the freshly produced FP8 weight - workspace (returned alongside ``out`` so the caller can refresh its - cache). The last three are ``None`` when gradients are disabled. + workspace returned alongside ``out`` so the caller can refresh its cache. + Scaling-factor checkpoint buffers are updated through ``args.scale_buffers``. """ weight = args.weight @@ -480,6 +493,25 @@ def _linear_forward_impl( if weight_quantizer is not None: weight_quantizer.calibrate(weight) + # Capture scaling metadata while it is still available. + if args.scale_buffers is not None: + scale_updates = {} + input_scale_buffer = _get_scale_buffer_info( + "input", inputmat_total, input_quantizer + ) + if input_scale_buffer is not None: + scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] + weight_scale_buffer = _get_scale_buffer_info( + "weight", weightmat, weight_quantizer + ) + if weight_scale_buffer is not None: + scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] + _update_scale_buffers( + args.scale_buffers, + scale_updates, + args.quantized_scaling_factor_buffering_decay, + ) + # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = _2X_ACC_FPROP if fp8: @@ -1532,6 +1564,18 @@ class Linear(TransformerEngineBaseModule): cast tensor. In some scenarios, the input tensor is used by multiple modules, and saving the original input tensor may reduce the memory usage. Cannot work with FP8 DelayedScaling recipe. + buffer_quantized_scaling_factors : bool, default = False + If set to ``True``, maintain nonpersistent input and weight quantization + metadata buffers for inference checkpoint export. Per-tensor buffers + store raw global amaxes, except FP8 current scaling buffers, which store + inverse scales directly. + Each buffer is materialized only when its tensor uses a quantizer with a + per-tensor scaling factor. Used to propagate scaling factors from training + into inference. + quantized_scaling_factor_buffering_decay : float, default = 0.0 + Decay applied to buffered activation scaling factors before incorporating + each new observation. Defaults to 0.0, in which case only the most recent + scaling factor is buffered. """ def __init__( @@ -1561,6 +1605,8 @@ def __init__( symmetric_ar_type: Optional[str] = None, save_original_input: bool = False, name: Optional[str] = None, + buffer_quantized_scaling_factors: bool = False, + quantized_scaling_factor_buffering_decay: float = 0.0, ) -> None: super().__init__(name) @@ -1788,6 +1834,9 @@ def __init__( if name in self.weight_names or name in self.bias_names: param.skip_backward_post_hook = True + self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors + self.quantized_scaling_factor_buffering_decay = quantized_scaling_factor_buffering_decay + def get_quantizer_roles( self, *, @@ -1972,6 +2021,13 @@ def forward( bias_tensor if (self.apply_bias and not self.gemm_bias_unfused_add) else None ) wgrad_store = self.wgrad_store if self.wgrad_store.delay_wgrad_compute() else None + scale_buffers = None + if self.buffer_quantized_scaling_factors: + scale_buffers = { + name: value + for name, value in self._buffers.items() + if name.endswith("_te_ptq_calibrated") + } fwd_args = LinearFwdArgs( # tensors weight=weight_tensor, @@ -2028,6 +2084,11 @@ def forward( # weight-grad scheduling fuse_wgrad_accumulation=self.fuse_wgrad_accumulation, wgrad_store=wgrad_store, + # Inference Scaling Factor Calibration Buffering + scale_buffers=scale_buffers, + quantized_scaling_factor_buffering_decay=( + self.quantized_scaling_factor_buffering_decay + ), # misc cpu_offloading=is_cpu_offload_enabled(), is_grad_enabled=is_grad_enabled, @@ -2040,6 +2101,16 @@ def forward( fwd_args, ) + if scale_buffers is not None: + # Assign the scaling factor calibration buffers to model. + # Requires CUDA graph warmup step. + for name, value in scale_buffers.items(): + if value is not None: + if name in self._buffers: + setattr(self, name, value) + else: + self.register_buffer(name, value, persistent=False) + if new_weight_workspace is not None and cache_name is not None: if isinstance(new_weight_workspace, torch.Tensor): new_weight_workspace = new_weight_workspace.detach() diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index 1789a0c98e..e56d618e5d 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -24,6 +24,26 @@ from ..constants import NVFP4_BLOCK_SCALING_SIZE, DType +def get_quantization_recipe_name(quantizer: Optional[Quantizer]) -> str: + """Get a stable recipe name from a quantizer.""" + quantizer = getattr(quantizer, "parent_quantizer", 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 isinstance(quantizer, NVFP4Quantizer): + if quantizer.row_scaled_nvfp4: + return "nvfp4_rowwise" + return "nvfp4" + raise ValueError(f"Unsupported quantizer type: {type(quantizer)}") + + def replace_raw_data(tensor: QuantizedTensor, new_raw_data: torch.Tensor): r"""Change a quantized tensor's data buffer while preserving values From 3058e6511db09db641e60c8713475c2c25bda2d3 Mon Sep 17 00:00:00 2001 From: Cory Ye Date: Fri, 31 Jul 2026 11:56:03 -0700 Subject: [PATCH 2/5] Add tests. Signed-off-by: Cory Ye --- ...test_ptq_calibration_metadata_buffering.py | 104 ++++++++++++++++++ transformer_engine/pytorch/module/_common.py | 28 +++-- 2 files changed, 125 insertions(+), 7 deletions(-) create mode 100644 tests/pytorch/test_ptq_calibration_metadata_buffering.py diff --git a/tests/pytorch/test_ptq_calibration_metadata_buffering.py b/tests/pytorch/test_ptq_calibration_metadata_buffering.py new file mode 100644 index 0000000000..94f1ba7664 --- /dev/null +++ b/tests/pytorch/test_ptq_calibration_metadata_buffering.py @@ -0,0 +1,104 @@ +# Copyright (c) 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 = "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])) diff --git a/transformer_engine/pytorch/module/_common.py b/transformer_engine/pytorch/module/_common.py index 0eddee7523..85a5f45b52 100644 --- a/transformer_engine/pytorch/module/_common.py +++ b/transformer_engine/pytorch/module/_common.py @@ -26,16 +26,30 @@ def _get_scale_buffer_info( recipe = get_quantization_recipe_name(quantizer) if not recipe: return None - if recipe == "fp8_current_scaling": + + 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) - else: + 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", - getattr(quantizer, "amax", None), - ) + 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 From ad4df8014789d9be1cd76b7333dac038d78b1833 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 31 Jul 2026 19:08:25 +0000 Subject: [PATCH 3/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../test_ptq_calibration_metadata_buffering.py | 12 +++--------- transformer_engine/pytorch/module/grouped_linear.py | 12 +++--------- .../pytorch/module/layernorm_linear.py | 12 +++--------- transformer_engine/pytorch/module/layernorm_mlp.py | 4 +--- transformer_engine/pytorch/module/linear.py | 8 ++------ 5 files changed, 12 insertions(+), 36 deletions(-) diff --git a/tests/pytorch/test_ptq_calibration_metadata_buffering.py b/tests/pytorch/test_ptq_calibration_metadata_buffering.py index 94f1ba7664..0f0fa1d3f0 100644 --- a/tests/pytorch/test_ptq_calibration_metadata_buffering.py +++ b/tests/pytorch/test_ptq_calibration_metadata_buffering.py @@ -26,9 +26,7 @@ def test_scale_buffer_info_selects_recipe_metadata( 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 - ), + _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)) @@ -47,9 +45,7 @@ def test_scale_buffer_info_skips_non_global_scaling_recipes(monkeypatch, recipe) def test_grouped_scale_buffers_are_per_gemm(monkeypatch): - monkeypatch.setattr( - _common, "get_quantization_recipe_name", lambda _: "fp8_current_scaling" - ) + 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])), @@ -76,9 +72,7 @@ def test_grouped_scale_buffers_are_per_gemm(monkeypatch): "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" - ], + scale_buffers["input_gemm1_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated"], torch.tensor([0.5]), ) diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index b5ae45e9b0..37cb59df7a 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -98,9 +98,7 @@ def _update_grouped_scale_buffers( if scale_buffer is not None: scale_updates[scale_buffer[0]] = scale_buffer[1] for index, tensor in enumerate(weight_tensors): - scale_buffer = _get_scale_buffer_info( - f"weight_gemm{index}", tensor, weight_quantizer - ) + scale_buffer = _get_scale_buffer_info(f"weight_gemm{index}", tensor, weight_quantizer) if scale_buffer is not None: scale_updates[scale_buffer[0]] = scale_buffer[1] _update_scale_buffers(scale_buffers, scale_updates, activation_scale_decay) @@ -664,9 +662,7 @@ def forward( cache_weight=cache_weight, skip_fp8_weight_update=skip_fp8_weight_update, scale_buffers=scale_buffers, - quantized_scaling_factor_buffering_decay=( - quantized_scaling_factor_buffering_decay - ), + quantized_scaling_factor_buffering_decay=(quantized_scaling_factor_buffering_decay), weights=weights, biases=biases, out=out, @@ -1491,9 +1487,7 @@ def __init__( self.ub_name = ub_name self.save_original_input = save_original_input self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors - self.quantized_scaling_factor_buffering_decay = ( - quantized_scaling_factor_buffering_decay - ) + self.quantized_scaling_factor_buffering_decay = quantized_scaling_factor_buffering_decay single_grouped_weight, single_grouped_bias = resolve_grouped_linear_single_param_flags( single_grouped_weight, single_grouped_bias ) diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index b317433d3e..529aeb104d 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -385,14 +385,10 @@ def forward( if scale_buffers is not None: scale_updates = {} - input_scale_buffer = _get_scale_buffer_info( - "input", ln_out_total, input_quantizer - ) + input_scale_buffer = _get_scale_buffer_info("input", ln_out_total, input_quantizer) if input_scale_buffer is not None: scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] - weight_scale_buffer = _get_scale_buffer_info( - "weight", weightmat, weight_quantizer - ) + weight_scale_buffer = _get_scale_buffer_info("weight", weightmat, weight_quantizer) if weight_scale_buffer is not None: scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] _update_scale_buffers( @@ -1630,9 +1626,7 @@ def __init__( param.skip_backward_post_hook = True self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors - self.quantized_scaling_factor_buffering_decay = ( - quantized_scaling_factor_buffering_decay - ) + self.quantized_scaling_factor_buffering_decay = quantized_scaling_factor_buffering_decay def set_meta_tensor(self, fwd: bool, recipe: Recipe) -> None: """Init scales and amaxes for fwd | bwd.""" diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 7b3de3ef66..f59f98ca80 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -2055,9 +2055,7 @@ def __init__( self.symmetric_ar_type = symmetric_ar_type self.checkpoint = checkpoint self.buffer_quantized_scaling_factors = buffer_quantized_scaling_factors - self.quantized_scaling_factor_buffering_decay = ( - quantized_scaling_factor_buffering_decay - ) + self.quantized_scaling_factor_buffering_decay = quantized_scaling_factor_buffering_decay # GEMM-GELU fusion is currently only supported with split GEMM-AG overlap self.gemm_gelu_fusion = ( diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 049ede7fa7..377848fb07 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -496,14 +496,10 @@ def _linear_forward_impl( # Capture scaling metadata while it is still available. if args.scale_buffers is not None: scale_updates = {} - input_scale_buffer = _get_scale_buffer_info( - "input", inputmat_total, input_quantizer - ) + input_scale_buffer = _get_scale_buffer_info("input", inputmat_total, input_quantizer) if input_scale_buffer is not None: scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] - weight_scale_buffer = _get_scale_buffer_info( - "weight", weightmat, weight_quantizer - ) + weight_scale_buffer = _get_scale_buffer_info("weight", weightmat, weight_quantizer) if weight_scale_buffer is not None: scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] _update_scale_buffers( From 50491115dcd4523fccd4b607490dfc8cdeda932d Mon Sep 17 00:00:00 2001 From: Cory Ye Date: Fri, 31 Jul 2026 12:42:22 -0700 Subject: [PATCH 4/5] Fix minor bugs. Signed-off-by: Cory Ye --- ...test_ptq_calibration_metadata_buffering.py | 4 +-- transformer_engine/pytorch/module/_common.py | 10 +++---- .../pytorch/module/grouped_linear.py | 18 ++++++++++--- .../pytorch/module/layernorm_linear.py | 18 +++++++------ .../pytorch/module/layernorm_mlp.py | 26 ++++++++++++++----- transformer_engine/pytorch/module/linear.py | 18 +++++++------ transformer_engine/pytorch/tensor/utils.py | 4 ++- 7 files changed, 62 insertions(+), 36 deletions(-) diff --git a/tests/pytorch/test_ptq_calibration_metadata_buffering.py b/tests/pytorch/test_ptq_calibration_metadata_buffering.py index 0f0fa1d3f0..da9b2c11a9 100644 --- a/tests/pytorch/test_ptq_calibration_metadata_buffering.py +++ b/tests/pytorch/test_ptq_calibration_metadata_buffering.py @@ -1,4 +1,4 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # See LICENSE for license information. @@ -87,7 +87,7 @@ def test_grouped_scale_buffers_are_per_gemm(monkeypatch): ), ) def test_activation_scale_buffer_uses_decaying_maximum(observed_scale, expected_scale): - name = "input_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated" + name = "fc1_input_tensor_scale_inv_fp8_current_scaling_te_ptq_calibrated" scale_buffers = {name: torch.tensor([4.0])} _common._update_scale_buffers( diff --git a/transformer_engine/pytorch/module/_common.py b/transformer_engine/pytorch/module/_common.py index 85a5f45b52..7ec51d1962 100644 --- a/transformer_engine/pytorch/module/_common.py +++ b/transformer_engine/pytorch/module/_common.py @@ -63,7 +63,7 @@ def _update_scale_buffers( for buffer_name, scale in scale_updates.items(): if scale is None: continue - if buffer_name.startswith("input"): + if activation_scale_decay > 0.0: observed_scale = scale.detach().float() scale_buffer = scale_buffers.get(buffer_name) if scale_buffer is not None and scale_buffer.shape != observed_scale.shape: @@ -71,10 +71,6 @@ def _update_scale_buffers( "Quantized scaling-factor buffer shape changed from " f"{tuple(scale_buffer.shape)} to {tuple(observed_scale.shape)}" ) - if activation_scale_decay == 0.0: - # If not using scale decay, just buffer the current scaling. - scale_buffers[buffer_name] = observed_scale - continue if scale_buffer is None: # Initialize the rolling activation scaling factor. # Requires CUDA graph warmup step. @@ -89,8 +85,8 @@ def _update_scale_buffers( out=scale_buffer, ) else: - # Keep a reference to the current weight metadata without - # allocating or copying a separate buffer. + # 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() diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index 37cb59df7a..d37d15a865 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -92,16 +92,26 @@ def _update_grouped_scale_buffers( activation_scale_decay: float, ) -> None: """Update GroupedLinear PTQ calibration buffers with per-GEMM metadata.""" - scale_updates = {} + activation_scale_updates = {} for index, tensor in enumerate(input_tensors): scale_buffer = _get_scale_buffer_info(f"input_gemm{index}", tensor, input_quantizer) if scale_buffer is not None: - scale_updates[scale_buffer[0]] = scale_buffer[1] + activation_scale_updates[scale_buffer[0]] = scale_buffer[1] + weight_scale_updates = {} for index, tensor in enumerate(weight_tensors): scale_buffer = _get_scale_buffer_info(f"weight_gemm{index}", tensor, weight_quantizer) if scale_buffer is not None: - scale_updates[scale_buffer[0]] = scale_buffer[1] - _update_scale_buffers(scale_buffers, scale_updates, activation_scale_decay) + weight_scale_updates[scale_buffer[0]] = scale_buffer[1] + _update_scale_buffers( + scale_buffers, + activation_scale_updates, + activation_scale_decay, + ) + _update_scale_buffers( + scale_buffers, + weight_scale_updates, + activation_scale_decay=0.0, + ) class _GroupedLinear(torch.autograd.Function): diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 529aeb104d..e234ce4726 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -384,18 +384,20 @@ def forward( weight_quantizer.calibrate(weight) if scale_buffers is not None: - scale_updates = {} input_scale_buffer = _get_scale_buffer_info("input", ln_out_total, input_quantizer) if input_scale_buffer is not None: - scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] + _update_scale_buffers( + scale_buffers, + {input_scale_buffer[0]: input_scale_buffer[1]}, + quantized_scaling_factor_buffering_decay, + ) weight_scale_buffer = _get_scale_buffer_info("weight", weightmat, weight_quantizer) if weight_scale_buffer is not None: - scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] - _update_scale_buffers( - scale_buffers, - scale_updates, - quantized_scaling_factor_buffering_decay, - ) + _update_scale_buffers( + scale_buffers, + {weight_scale_buffer[0]: weight_scale_buffer[1]}, + activation_scale_decay=0.0, + ) # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = _2X_ACC_FPROP diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index f59f98ca80..76f7848f52 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -697,26 +697,40 @@ def _forward( fc2_weight_quantizer.calibrate(fc2_weight) if scale_buffers is not None: - scale_updates = {} + activation_scale_updates = {} + weight_scale_updates = {} if fc1_input_scale_buffer is not None: - scale_updates[fc1_input_scale_buffer[0]] = fc1_input_scale_buffer[1] + activation_scale_updates[fc1_input_scale_buffer[0]] = ( + fc1_input_scale_buffer[1] + ) if fc1_weight_scale_buffer is not None: - scale_updates[fc1_weight_scale_buffer[0]] = fc1_weight_scale_buffer[1] + weight_scale_updates[fc1_weight_scale_buffer[0]] = ( + fc1_weight_scale_buffer[1] + ) fc2_input_scale_buffer = _get_scale_buffer_info( "fc2_input", act_out, fc2_input_quantizer ) if fc2_input_scale_buffer is not None: - scale_updates[fc2_input_scale_buffer[0]] = fc2_input_scale_buffer[1] + activation_scale_updates[fc2_input_scale_buffer[0]] = ( + fc2_input_scale_buffer[1] + ) fc2_weight_scale_buffer = _get_scale_buffer_info( "fc2_weight", fc2_weight_final, fc2_weight_quantizer ) if fc2_weight_scale_buffer is not None: - scale_updates[fc2_weight_scale_buffer[0]] = fc2_weight_scale_buffer[1] + weight_scale_updates[fc2_weight_scale_buffer[0]] = ( + fc2_weight_scale_buffer[1] + ) _update_scale_buffers( scale_buffers, - scale_updates, + activation_scale_updates, quantized_scaling_factor_buffering_decay, ) + _update_scale_buffers( + scale_buffers, + weight_scale_updates, + activation_scale_decay=0.0, + ) # Configure Userbuffers reduce-scatter if needed ub_obj_fc2out = None diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 377848fb07..a8541eb085 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -495,18 +495,20 @@ def _linear_forward_impl( # Capture scaling metadata while it is still available. if args.scale_buffers is not None: - scale_updates = {} input_scale_buffer = _get_scale_buffer_info("input", inputmat_total, input_quantizer) if input_scale_buffer is not None: - scale_updates[input_scale_buffer[0]] = input_scale_buffer[1] + _update_scale_buffers( + args.scale_buffers, + {input_scale_buffer[0]: input_scale_buffer[1]}, + args.quantized_scaling_factor_buffering_decay, + ) weight_scale_buffer = _get_scale_buffer_info("weight", weightmat, weight_quantizer) if weight_scale_buffer is not None: - scale_updates[weight_scale_buffer[0]] = weight_scale_buffer[1] - _update_scale_buffers( - args.scale_buffers, - scale_updates, - args.quantized_scaling_factor_buffering_decay, - ) + _update_scale_buffers( + args.scale_buffers, + {weight_scale_buffer[0]: weight_scale_buffer[1]}, + activation_scale_decay=0.0, + ) # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = _2X_ACC_FPROP diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index e56d618e5d..f69926246a 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -41,7 +41,9 @@ def get_quantization_recipe_name(quantizer: Optional[Quantizer]) -> str: if quantizer.row_scaled_nvfp4: return "nvfp4_rowwise" return "nvfp4" - raise ValueError(f"Unsupported quantizer type: {type(quantizer)}") + # Custom recipes may provide arbitrary Quantizer implementations without a + # stable recipe name or globally checkpointable scaling metadata. + return "" def replace_raw_data(tensor: QuantizedTensor, new_raw_data: torch.Tensor): From b830ba69f7da7d92914b79f070f21f461271f7f8 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 31 Jul 2026 19:46:32 +0000 Subject: [PATCH 5/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../pytorch/module/layernorm_mlp.py | 16 ++++------------ 1 file changed, 4 insertions(+), 12 deletions(-) diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 76f7848f52..25b7c7f446 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -700,27 +700,19 @@ def _forward( activation_scale_updates = {} weight_scale_updates = {} if fc1_input_scale_buffer is not None: - activation_scale_updates[fc1_input_scale_buffer[0]] = ( - fc1_input_scale_buffer[1] - ) + activation_scale_updates[fc1_input_scale_buffer[0]] = fc1_input_scale_buffer[1] if fc1_weight_scale_buffer is not None: - weight_scale_updates[fc1_weight_scale_buffer[0]] = ( - fc1_weight_scale_buffer[1] - ) + weight_scale_updates[fc1_weight_scale_buffer[0]] = fc1_weight_scale_buffer[1] fc2_input_scale_buffer = _get_scale_buffer_info( "fc2_input", act_out, fc2_input_quantizer ) if fc2_input_scale_buffer is not None: - activation_scale_updates[fc2_input_scale_buffer[0]] = ( - fc2_input_scale_buffer[1] - ) + activation_scale_updates[fc2_input_scale_buffer[0]] = fc2_input_scale_buffer[1] fc2_weight_scale_buffer = _get_scale_buffer_info( "fc2_weight", fc2_weight_final, fc2_weight_quantizer ) if fc2_weight_scale_buffer is not None: - weight_scale_updates[fc2_weight_scale_buffer[0]] = ( - fc2_weight_scale_buffer[1] - ) + weight_scale_updates[fc2_weight_scale_buffer[0]] = fc2_weight_scale_buffer[1] _update_scale_buffers( scale_buffers, activation_scale_updates,