From 5a831a313b0215093099cf7aefffa01d58ef7852 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Tue, 14 Oct 2025 14:57:45 +0000 Subject: [PATCH 01/10] rename experimental -> custom_recipes Signed-off-by: Evgeny --- tests/pytorch/distributed/run_numerics_exact.py | 4 ++-- tests/pytorch/nvfp4/test_nvfp4_gemm_exact.py | 4 ++-- tests/pytorch/nvfp4/test_nvfp4_module_exact.py | 4 ++-- tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py | 4 ++-- tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py | 4 ++-- transformer_engine/pytorch/cpp_extensions/gemm.py | 2 +- .../pytorch/{experimental => custom_recipes}/__init__.py | 0 .../pytorch/{experimental => custom_recipes}/gemm.py | 4 ++-- .../{experimental => custom_recipes}/quantization.py | 0 .../quantization_nvfp4.py | 4 ++-- .../pytorch/{experimental => custom_recipes}/utils.py | 0 transformer_engine/pytorch/tensor/utils.py | 9 ++------- 12 files changed, 17 insertions(+), 22 deletions(-) rename transformer_engine/pytorch/{experimental => custom_recipes}/__init__.py (100%) rename transformer_engine/pytorch/{experimental => custom_recipes}/gemm.py (96%) rename transformer_engine/pytorch/{experimental => custom_recipes}/quantization.py (100%) rename transformer_engine/pytorch/{experimental => custom_recipes}/quantization_nvfp4.py (99%) rename transformer_engine/pytorch/{experimental => custom_recipes}/utils.py (100%) diff --git a/tests/pytorch/distributed/run_numerics_exact.py b/tests/pytorch/distributed/run_numerics_exact.py index 40be8e1f06..cea6c28264 100644 --- a/tests/pytorch/distributed/run_numerics_exact.py +++ b/tests/pytorch/distributed/run_numerics_exact.py @@ -25,8 +25,8 @@ ) from transformer_engine.pytorch.tensor.nvfp4_tensor import NVFP4Quantizer from transformer_engine.pytorch.constants import NVFP4_BLOCK_SCALING_SIZE -from transformer_engine.pytorch.experimental import quantization_nvfp4 -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes import quantization_nvfp4 +from transformer_engine.pytorch.custom_recipes import utils from run_layer_with_overlap import _compare_tensors diff --git a/tests/pytorch/nvfp4/test_nvfp4_gemm_exact.py b/tests/pytorch/nvfp4/test_nvfp4_gemm_exact.py index 42837fb40f..3ff1a98d5f 100644 --- a/tests/pytorch/nvfp4/test_nvfp4_gemm_exact.py +++ b/tests/pytorch/nvfp4/test_nvfp4_gemm_exact.py @@ -9,8 +9,8 @@ from transformer_engine.pytorch.fp8 import FP8GlobalStateManager from transformer_engine.pytorch.constants import TE_DType from transformer_engine.pytorch.tensor.nvfp4_tensor import NVFP4Quantizer -from transformer_engine.pytorch.experimental.quantization_nvfp4 import NVFP4QuantizerRef -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import NVFP4QuantizerRef +from transformer_engine.pytorch.custom_recipes import utils recipe_available, reason_for_no_recipe = FP8GlobalStateManager.is_nvfp4_available() diff --git a/tests/pytorch/nvfp4/test_nvfp4_module_exact.py b/tests/pytorch/nvfp4/test_nvfp4_module_exact.py index 1d1467640c..26b5b26a71 100644 --- a/tests/pytorch/nvfp4/test_nvfp4_module_exact.py +++ b/tests/pytorch/nvfp4/test_nvfp4_module_exact.py @@ -8,8 +8,8 @@ from transformer_engine.pytorch.fp8 import FP8GlobalStateManager from transformer_engine.pytorch.distributed import fp8_autocast from transformer_engine.common import recipe -from transformer_engine.pytorch.experimental import quantization_nvfp4 -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes import quantization_nvfp4 +from transformer_engine.pytorch.custom_recipes import utils recipe_available, reason_for_no_recipe = FP8GlobalStateManager.is_nvfp4_available() diff --git a/tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py b/tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py index cdcb2df1dc..0879e60d93 100644 --- a/tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py +++ b/tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py @@ -12,8 +12,8 @@ from transformer_engine.pytorch.tensor.nvfp4_tensor import ( NVFP4Quantizer, ) -from transformer_engine.pytorch.experimental.quantization_nvfp4 import NVFP4QuantizerRef -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import NVFP4QuantizerRef +from transformer_engine.pytorch.custom_recipes import utils from transformer_engine.pytorch.fp8 import fp8_autocast, get_fp4_te_dtype diff --git a/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py b/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py index 494fa63c09..c72e0d850b 100644 --- a/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py +++ b/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py @@ -20,8 +20,8 @@ from transformer_engine.pytorch.tensor.nvfp4_tensor import ( NVFP4Quantizer, ) -from transformer_engine.pytorch.experimental.quantization_nvfp4 import NVFP4QuantizerRef -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import NVFP4QuantizerRef +from transformer_engine.pytorch.custom_recipes import utils from transformer_engine.pytorch.fp8 import fp8_autocast, get_fp4_te_dtype import pytest diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index a45fafb68a..6282f26fc6 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -14,7 +14,7 @@ from ..tensor.quantized_tensor import Quantizer from ..tensor.storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage from ..tensor.utils import is_experimental -from ..experimental.gemm import experimental_gemm +from ..custom_recipes.gemm import experimental_gemm from ...debug.pytorch.debug_quantization import DebugQuantizer __all__ = [ diff --git a/transformer_engine/pytorch/experimental/__init__.py b/transformer_engine/pytorch/custom_recipes/__init__.py similarity index 100% rename from transformer_engine/pytorch/experimental/__init__.py rename to transformer_engine/pytorch/custom_recipes/__init__.py diff --git a/transformer_engine/pytorch/experimental/gemm.py b/transformer_engine/pytorch/custom_recipes/gemm.py similarity index 96% rename from transformer_engine/pytorch/experimental/gemm.py rename to transformer_engine/pytorch/custom_recipes/gemm.py index 0bd740d85d..75f8777033 100644 --- a/transformer_engine/pytorch/experimental/gemm.py +++ b/transformer_engine/pytorch/custom_recipes/gemm.py @@ -2,13 +2,13 @@ # # See LICENSE for license information. -"""GEMM API for experimental middleware between Transformer Engine and Kitchen.""" +"""GEMM API that enables custom GEMM logic for custom quantization recipes.""" from typing import Iterable, Optional import torch -from transformer_engine.pytorch.experimental.quantization import ( +from transformer_engine.pytorch.custom_recipes.quantization import ( MMParams, GEMMType, ) diff --git a/transformer_engine/pytorch/experimental/quantization.py b/transformer_engine/pytorch/custom_recipes/quantization.py similarity index 100% rename from transformer_engine/pytorch/experimental/quantization.py rename to transformer_engine/pytorch/custom_recipes/quantization.py diff --git a/transformer_engine/pytorch/experimental/quantization_nvfp4.py b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py similarity index 99% rename from transformer_engine/pytorch/experimental/quantization_nvfp4.py rename to transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py index fc50d07424..10615df1aa 100644 --- a/transformer_engine/pytorch/experimental/quantization_nvfp4.py +++ b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py @@ -9,8 +9,8 @@ import torch -from transformer_engine.pytorch.experimental import quantization -from transformer_engine.pytorch.experimental import utils +from transformer_engine.pytorch.custom_recipes import quantization +from transformer_engine.pytorch.custom_recipes import utils from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensorStorage, Quantizer diff --git a/transformer_engine/pytorch/experimental/utils.py b/transformer_engine/pytorch/custom_recipes/utils.py similarity index 100% rename from transformer_engine/pytorch/experimental/utils.py rename to transformer_engine/pytorch/custom_recipes/utils.py diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index cc02494013..e9a299e01f 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -455,16 +455,11 @@ def _cast_master_weights_to_fp8_blockwise_scaling( def is_experimental(x: Optional[Union[Quantizer, QuantizedTensorStorage]] = None) -> bool: - """Check if an environment or object is using experimental Kitchen middleware. + """Check if an object is experimental. Returns False if x is a torch.Tensor. """ - # Detect if the environment is experimental - if x is None: - return int(os.getenv("QAT_PARAMS", "0")) > 0 - - # Detect if the object is experimental - if isinstance(x, torch.Tensor): + if x is None or isinstance(x, torch.Tensor): return False if not isinstance(x, (Quantizer, QuantizedTensorStorage)): raise AssertionError("Object must be a Quantizer or QuantizedTensorStorage instance") From dce8556bbccceffa5f9f47d6f9ca0359c02d6e74 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Wed, 15 Oct 2025 14:45:23 +0000 Subject: [PATCH 02/10] Decouple python base classes (api) Signed-off-by: Evgeny --- tests/pytorch/attention/test_attention.py | 2 +- .../test_fusible_ops_with_userbuffers.py | 2 +- .../debug/pytorch/debug_quantization.py | 2 +- .../dot_product_attention/backends.py | 2 +- .../dot_product_attention/context_parallel.py | 4 +- .../pytorch/cpp_extensions/fused_attn.py | 2 +- .../pytorch/cpp_extensions/gemm.py | 2 +- transformer_engine/pytorch/cpu_offload.py | 2 +- .../pytorch/custom_recipes/gemm.py | 2 +- .../custom_recipes/quantization_nvfp4.py | 2 +- transformer_engine/pytorch/distributed.py | 2 +- transformer_engine/pytorch/module/base.py | 2 +- .../pytorch/module/grouped_linear.py | 2 +- .../pytorch/module/layernorm_linear.py | 2 +- .../pytorch/module/layernorm_mlp.py | 2 +- transformer_engine/pytorch/module/linear.py | 2 +- transformer_engine/pytorch/ops/_common.py | 2 +- .../ops/fused/userbuffers_backward_linear.py | 2 +- .../ops/fused/userbuffers_forward_linear.py | 2 +- transformer_engine/pytorch/ops/fuser.py | 2 +- transformer_engine/pytorch/permutation.py | 2 +- .../quantized_tensor.py => quantization.py} | 76 ++--------------- transformer_engine/pytorch/tensor/__init__.py | 2 +- .../pytorch/tensor/_quantization_helpers.py | 84 +++++++++++++++++++ .../pytorch/tensor/float8_blockwise_tensor.py | 7 +- .../pytorch/tensor/float8_tensor.py | 7 +- .../pytorch/tensor/mxfp8_tensor.py | 7 +- .../pytorch/tensor/nvfp4_tensor.py | 3 +- .../float8_blockwise_tensor_storage.py | 4 +- .../tensor/storage/float8_tensor_storage.py | 4 +- .../tensor/storage/mxfp8_tensor_storage.py | 4 +- .../tensor/storage/nvfp4_tensor_storage.py | 4 +- transformer_engine/pytorch/tensor/utils.py | 2 +- transformer_engine/pytorch/utils.py | 2 +- 34 files changed, 131 insertions(+), 119 deletions(-) rename transformer_engine/pytorch/{tensor/quantized_tensor.py => quantization.py} (89%) create mode 100644 transformer_engine/pytorch/tensor/_quantization_helpers.py diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index e3a4de73b0..cfbfa55d17 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -39,7 +39,7 @@ ) from transformer_engine.pytorch.utils import get_cudnn_version import transformer_engine_torch as tex -from transformer_engine.pytorch.tensor.quantized_tensor import ( +from transformer_engine.pytorch.quantization import ( Quantizer, prepare_for_saving, restore_from_saved, diff --git a/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py b/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py index d6ddfe27c9..f6db2df1e8 100644 --- a/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py +++ b/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py @@ -30,7 +30,7 @@ Float8CurrentScalingQuantizer, ) from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer -from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensor +from transformer_engine.pytorch.quantization import QuantizedTensor from transformer_engine.pytorch.tensor.float8_tensor import Float8Tensor from transformer_engine.pytorch.utils import is_bf16_compatible diff --git a/transformer_engine/debug/pytorch/debug_quantization.py b/transformer_engine/debug/pytorch/debug_quantization.py index 185bf15d05..aacf042977 100644 --- a/transformer_engine/debug/pytorch/debug_quantization.py +++ b/transformer_engine/debug/pytorch/debug_quantization.py @@ -15,7 +15,7 @@ import transformer_engine_torch as tex from transformer_engine.common.recipe import Recipe -from transformer_engine.pytorch.tensor.quantized_tensor import ( +from transformer_engine.pytorch.quantization import ( QuantizedTensor, Quantizer, QuantizedTensorStorage, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index d75481ad9a..1c359dae2e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -24,7 +24,7 @@ Float8Quantizer, Float8CurrentScalingQuantizer, ) -from transformer_engine.pytorch.tensor.quantized_tensor import ( +from transformer_engine.pytorch.quantization import ( QuantizedTensorStorage, prepare_for_saving, restore_from_saved, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index d1374e949e..73a2acc260 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -21,7 +21,7 @@ ) from transformer_engine.pytorch.fp8 import FP8GlobalStateManager from transformer_engine.pytorch.tensor.float8_tensor import Float8Tensor -from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensorStorage +from transformer_engine.pytorch.quantization import QuantizedTensorStorage from transformer_engine.pytorch.jit import jit_fuser from transformer_engine.pytorch.constants import ( dist_group_type, @@ -33,7 +33,7 @@ gather_along_first_dim, reduce_scatter_along_first_dim, ) -from transformer_engine.pytorch.tensor.quantized_tensor import ( +from transformer_engine.pytorch.quantization import ( prepare_for_saving, restore_from_saved, ) diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 94a12c4a09..446fbe18e5 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -15,7 +15,7 @@ NVTE_Softmax_Type, NVTE_Fused_Attn_Backend, ) -from ..tensor.quantized_tensor import Quantizer +from ..quantization import Quantizer __all__ = [ diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index 6282f26fc6..b4f905a7ad 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -11,7 +11,7 @@ from ..constants import TE_DType from ..utils import get_sm_count, _empty_tensor -from ..tensor.quantized_tensor import Quantizer +from ..quantization import Quantizer from ..tensor.storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage from ..tensor.utils import is_experimental from ..custom_recipes.gemm import experimental_gemm diff --git a/transformer_engine/pytorch/cpu_offload.py b/transformer_engine/pytorch/cpu_offload.py index 648b21eb4d..670e172306 100644 --- a/transformer_engine/pytorch/cpu_offload.py +++ b/transformer_engine/pytorch/cpu_offload.py @@ -10,7 +10,7 @@ import torch from transformer_engine.debug.pytorch.debug_state import TEDebugState -from .tensor.quantized_tensor import QuantizedTensorStorage +from .quantization import QuantizedTensorStorage from .tensor.float8_tensor import Float8Tensor __all__ = ["get_cpu_offload_context"] diff --git a/transformer_engine/pytorch/custom_recipes/gemm.py b/transformer_engine/pytorch/custom_recipes/gemm.py index 75f8777033..71d58a565a 100644 --- a/transformer_engine/pytorch/custom_recipes/gemm.py +++ b/transformer_engine/pytorch/custom_recipes/gemm.py @@ -12,7 +12,7 @@ MMParams, GEMMType, ) -from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensorStorage, Quantizer +from transformer_engine.pytorch.quantization import QuantizedTensorStorage, Quantizer from transformer_engine.pytorch.tensor.utils import is_experimental diff --git a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py index 10615df1aa..eaab4bfae7 100644 --- a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py +++ b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py @@ -11,7 +11,7 @@ from transformer_engine.pytorch.custom_recipes import quantization from transformer_engine.pytorch.custom_recipes import utils -from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensorStorage, Quantizer +from transformer_engine.pytorch.quantization import QuantizedTensorStorage, Quantizer def nvfp4_ref_rht_2d_quantizer_factory(role): diff --git a/transformer_engine/pytorch/distributed.py b/transformer_engine/pytorch/distributed.py index 51fbb50c4c..43e5fd7d43 100644 --- a/transformer_engine/pytorch/distributed.py +++ b/transformer_engine/pytorch/distributed.py @@ -41,7 +41,7 @@ from .tensor.mxfp8_tensor import MXFP8Quantizer from .tensor.nvfp4_tensor import NVFP4Quantizer from .tensor.float8_blockwise_tensor import Float8BlockQuantizer -from .tensor.quantized_tensor import QuantizedTensorStorage, QuantizedTensor, Quantizer +from .quantization import QuantizedTensorStorage, QuantizedTensor, Quantizer from .tensor.storage.float8_tensor_storage import Float8TensorStorage from .tensor.storage.mxfp8_tensor_storage import MXFP8TensorStorage from .tensor.storage.nvfp4_tensor_storage import NVFP4TensorStorage diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index 838ac5281c..b4f3626082 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -38,7 +38,7 @@ _fsdp_gather_tensors, ) from ..constants import dist_group_type -from ..tensor.quantized_tensor import QuantizedTensor, QuantizedTensorStorage, Quantizer +from ..quantization import QuantizedTensor, QuantizedTensorStorage, Quantizer from ..tensor.float8_tensor import Float8Quantizer, Float8CurrentScalingQuantizer from ..tensor.mxfp8_tensor import MXFP8Quantizer from ..tensor.float8_blockwise_tensor import Float8BlockQuantizer diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index ec05f684b8..5472c1d836 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -43,7 +43,7 @@ from ..cpu_offload import is_cpu_offload_enabled from ..tensor.float8_tensor import Float8CurrentScalingQuantizer, Float8Quantizer -from ..tensor.quantized_tensor import ( +from ..quantization import ( QuantizedTensorStorage, Quantizer, prepare_for_saving, diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 824fcc0a7d..890825bb51 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -56,7 +56,7 @@ from ..jit import no_torch_dynamo from ..graph import is_graph_capturing from ._common import apply_normalization, noop_cat, WeightGradStore -from ..tensor.quantized_tensor import ( +from ..quantization import ( QuantizedTensor, QuantizedTensorStorage, Quantizer, diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 8ef19d0520..627e37a259 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -70,7 +70,7 @@ from ..tensor.float8_blockwise_tensor import Float8BlockQuantizer from ._common import apply_normalization, WeightGradStore from ..cpu_offload import is_cpu_offload_enabled, mark_activation_offload -from ..tensor.quantized_tensor import ( +from ..quantization import ( QuantizedTensorStorage, Quantizer, prepare_for_saving, diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 12b7bac011..6101cddfef 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -57,7 +57,7 @@ from ..constants import GemmParallelModes, dist_group_type from ..jit import no_torch_dynamo from ..graph import is_graph_capturing -from ..tensor.quantized_tensor import ( +from ..quantization import ( QuantizedTensor, QuantizedTensorStorage, Quantizer, diff --git a/transformer_engine/pytorch/ops/_common.py b/transformer_engine/pytorch/ops/_common.py index 13db35fc78..71cb70aeb8 100644 --- a/transformer_engine/pytorch/ops/_common.py +++ b/transformer_engine/pytorch/ops/_common.py @@ -13,7 +13,7 @@ from .. import torch_version from ..fp8 import FP8GlobalStateManager from ..tensor.float8_tensor import Float8Tensor -from ..tensor.quantized_tensor import QuantizedTensorStorage +from ..quantization import QuantizedTensorStorage from ..utils import canonicalize_dtype diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py index d95b2298fe..0de4733b93 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py @@ -21,7 +21,7 @@ get_ub, get_workspace, ) -from ...tensor.quantized_tensor import Quantizer +from ...quantization import Quantizer from ...tensor.mxfp8_tensor import MXFP8Quantizer from ...utils import canonicalize_device, canonicalize_dtype, clear_tensor_data from ..basic import BasicLinear, Bias, ReduceScatter diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index cbbe529d6a..24c654c06d 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -21,7 +21,7 @@ get_workspace, _2X_ACC_FPROP, ) -from ...tensor.quantized_tensor import Quantizer +from ...quantization import Quantizer from ...tensor.float8_tensor import Float8Quantizer, Float8CurrentScalingQuantizer from ...tensor.storage.float8_tensor_storage import Float8TensorStorage from .._common import maybe_dequantize, is_quantized_tensor diff --git a/transformer_engine/pytorch/ops/fuser.py b/transformer_engine/pytorch/ops/fuser.py index 6f80a7a1f3..6d7a153a36 100644 --- a/transformer_engine/pytorch/ops/fuser.py +++ b/transformer_engine/pytorch/ops/fuser.py @@ -28,7 +28,7 @@ fuse_userbuffers_backward_linear, fuse_userbuffers_forward_linear, ) -from transformer_engine.pytorch.tensor.quantized_tensor import ( +from transformer_engine.pytorch.quantization import ( prepare_for_saving, restore_from_saved, ) diff --git a/transformer_engine/pytorch/permutation.py b/transformer_engine/pytorch/permutation.py index ea3e67a57c..f796f774cb 100644 --- a/transformer_engine/pytorch/permutation.py +++ b/transformer_engine/pytorch/permutation.py @@ -10,7 +10,7 @@ import transformer_engine_torch as tex import transformer_engine.pytorch.triton.permutation as triton_permutation from transformer_engine.pytorch.constants import TE_DType -from transformer_engine.pytorch.tensor.quantized_tensor import QuantizedTensor +from transformer_engine.pytorch.quantization import QuantizedTensor from transformer_engine.pytorch.tensor.float8_tensor import Float8Tensor from transformer_engine.pytorch.tensor.float8_blockwise_tensor import Float8BlockwiseQTensor from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Tensor diff --git a/transformer_engine/pytorch/tensor/quantized_tensor.py b/transformer_engine/pytorch/quantization.py similarity index 89% rename from transformer_engine/pytorch/tensor/quantized_tensor.py rename to transformer_engine/pytorch/quantization.py index a524d5c8de..15f5b6bd5e 100644 --- a/transformer_engine/pytorch/tensor/quantized_tensor.py +++ b/transformer_engine/pytorch/quantization.py @@ -2,10 +2,10 @@ # # See LICENSE for license information. -"""Tensor with quantized data""" +"""Pure Python base classes for quantization.""" from __future__ import annotations -from typing import Callable, Optional, Tuple, Iterable, Any, Dict, Union +from typing import Optional, Tuple, Iterable, Any, Dict, Union import abc import copy import warnings @@ -14,6 +14,11 @@ from torch.utils._pytree import tree_map from transformer_engine.common.recipe import Recipe +from transformer_engine.pytorch.tensor._quantization_helpers import ( + _QuantizeFunc, + _IdentityFunc, + _stride_from_shape, +) class QuantizedTensorStorage: @@ -310,73 +315,6 @@ def is_quantizable(self, inp: torch.Tensor) -> bool: # pylint: disable=unused-a return True -class _QuantizeFunc(torch.autograd.Function): - """Quantize tensor""" - - @staticmethod - def forward( - _ctx: Optional[torch.autograd.function.FunctionCtx], # unused - tensor: torch.Tensor, - quantize_impl: Callable, - ) -> QuantizedTensor: - # pylint: disable=missing-function-docstring - return quantize_impl(tensor) - - @staticmethod - def backward( - _ctx: torch.autograd.function.FunctionCtx, # unused - grad: torch.Tensor, - ) -> Tuple[Optional[torch.Tensor], ...]: - # pylint: disable=missing-function-docstring - # Assume that we want gradients in full precision - return grad, None - - -class _IdentityFunc(torch.autograd.Function): - """Identity function - - If constructor keyword-arguments are provided, then construct a - new Float8Tensor using the provided tensor's attributes. - - """ - - @staticmethod - def forward( - ctx, tensor: QuantizedTensor, init_kwargs: Optional[Dict[str, Any]] = None - ) -> QuantizedTensor: - # pylint: disable=missing-function-docstring - - # Return input tensor if constructor kwargs are not provided - if init_kwargs is None: - return tensor.detach() - - # Construct new tensor if constructor kwargs are provided - ctx.input_dtype = tensor.dtype - kwargs = tensor.get_metadata() - for key, val in init_kwargs.items(): - kwargs[key] = val - return type(tensor)(tensor.shape, tensor.dtype, **kwargs) - - @staticmethod - def backward(ctx, grad_output): - # pylint: disable=missing-function-docstring - grad_input = grad_output - if grad_input.dtype == ctx.input_dtype: - grad_input = grad_input.detach() - else: - grad_input = grad_input.to(ctx.input_dtype) - return grad_input, None - - -def _stride_from_shape(shape: list[int]): - if len(shape) == 0: - return [] - rstride = [1] - for d in reversed(shape[1:]): - rstride.append(rstride[-1] * d) - return list(reversed(rstride)) - - class QuantizedTensor(torch.Tensor): """Abstract base class for tensor with quantized data diff --git a/transformer_engine/pytorch/tensor/__init__.py b/transformer_engine/pytorch/tensor/__init__.py index 7689e20194..9f52841aec 100644 --- a/transformer_engine/pytorch/tensor/__init__.py +++ b/transformer_engine/pytorch/tensor/__init__.py @@ -6,7 +6,7 @@ import torch -from .quantized_tensor import ( +from ..quantization import ( QuantizedTensorStorage, QuantizedTensor, Quantizer, diff --git a/transformer_engine/pytorch/tensor/_quantization_helpers.py b/transformer_engine/pytorch/tensor/_quantization_helpers.py new file mode 100644 index 0000000000..25003d443f --- /dev/null +++ b/transformer_engine/pytorch/tensor/_quantization_helpers.py @@ -0,0 +1,84 @@ +# Copyright (c) 2022-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Private helper functions and classes for quantized tensor implementations. + +This module contains internal autograd functions and utilities that support +the quantization machinery. +""" + +from __future__ import annotations +from typing import Callable, Optional, Tuple, Any, Dict, TYPE_CHECKING +import torch + +if TYPE_CHECKING: + from transformer_engine.pytorch.quantization import QuantizedTensor + + +class _QuantizeFunc(torch.autograd.Function): + """Quantize tensor""" + + @staticmethod + def forward( + _ctx: Optional[torch.autograd.function.FunctionCtx], # unused + tensor: torch.Tensor, + quantize_impl: Callable, + ) -> QuantizedTensor: + # pylint: disable=missing-function-docstring + return quantize_impl(tensor) + + @staticmethod + def backward( + _ctx: torch.autograd.function.FunctionCtx, # unused + grad: torch.Tensor, + ) -> Tuple[Optional[torch.Tensor], ...]: + # pylint: disable=missing-function-docstring + # Assume that we want gradients in full precision + return grad, None + + +class _IdentityFunc(torch.autograd.Function): + """Identity function + + If constructor keyword-arguments are provided, then construct a + new Float8Tensor using the provided tensor's attributes. + + """ + + @staticmethod + def forward( + ctx, tensor: QuantizedTensor, init_kwargs: Optional[Dict[str, Any]] = None + ) -> QuantizedTensor: + # pylint: disable=missing-function-docstring + + # Return input tensor if constructor kwargs are not provided + if init_kwargs is None: + return tensor.detach() + + # Construct new tensor if constructor kwargs are provided + ctx.input_dtype = tensor.dtype + kwargs = tensor.get_metadata() + for key, val in init_kwargs.items(): + kwargs[key] = val + return type(tensor)(tensor.shape, tensor.dtype, **kwargs) + + @staticmethod + def backward(ctx, grad_output): + # pylint: disable=missing-function-docstring + grad_input = grad_output + if grad_input.dtype == ctx.input_dtype: + grad_input = grad_input.detach() + else: + grad_input = grad_input.to(ctx.input_dtype) + return grad_input, None + + +def _stride_from_shape(shape: list[int]): + """Calculate stride from shape for contiguous tensors""" + if len(shape) == 0: + return [] + rstride = [1] + for d in reversed(shape[1:]): + rstride.append(rstride[-1] * d) + return list(reversed(rstride)) diff --git a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py index 16631a3d0d..3887075f5a 100644 --- a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py @@ -14,11 +14,8 @@ from transformer_engine.common.recipe import Float8BlockScaling, Recipe from .storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage -from .quantized_tensor import ( - QuantizedTensor, - Quantizer, - _IdentityFunc, -) +from ..quantization import QuantizedTensor, Quantizer +from ._quantization_helpers import _IdentityFunc from ..utils import devices_match, round_up_to_nearest_multiple aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/float8_tensor.py b/transformer_engine/pytorch/tensor/float8_tensor.py index a4e68e53b0..1b3d937883 100644 --- a/transformer_engine/pytorch/tensor/float8_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_tensor.py @@ -14,11 +14,8 @@ from transformer_engine.common.recipe import DelayedScaling, Float8CurrentScaling, Recipe from ..utils import canonicalize_process_group, devices_match from .storage.float8_tensor_storage import Float8TensorStorage, _FromFloat8Func -from .quantized_tensor import ( - QuantizedTensor, - Quantizer, - _IdentityFunc, -) +from ..quantization import QuantizedTensor, Quantizer +from ._quantization_helpers import _IdentityFunc from ..constants import dist_group_type aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/mxfp8_tensor.py b/transformer_engine/pytorch/tensor/mxfp8_tensor.py index 700de24c4e..857c480296 100644 --- a/transformer_engine/pytorch/tensor/mxfp8_tensor.py +++ b/transformer_engine/pytorch/tensor/mxfp8_tensor.py @@ -17,11 +17,8 @@ from ..utils import devices_match, round_up_to_nearest_multiple from .storage.mxfp8_tensor_storage import MXFP8TensorStorage, _FromMXFP8Func -from .quantized_tensor import ( - QuantizedTensor, - Quantizer, - _IdentityFunc, -) +from ..quantization import QuantizedTensor, Quantizer +from ._quantization_helpers import _IdentityFunc aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/nvfp4_tensor.py b/transformer_engine/pytorch/tensor/nvfp4_tensor.py index ca2154f554..8e3689852d 100644 --- a/transformer_engine/pytorch/tensor/nvfp4_tensor.py +++ b/transformer_engine/pytorch/tensor/nvfp4_tensor.py @@ -22,7 +22,8 @@ ) from .storage.nvfp4_tensor_storage import NVFP4TensorStorage, _FromNVFP4Func -from .quantized_tensor import QuantizedTensor, Quantizer, _IdentityFunc +from ..quantization import QuantizedTensor, Quantizer +from ._quantization_helpers import _IdentityFunc aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py index 9040ea3a43..8c018169ae 100644 --- a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py @@ -13,11 +13,11 @@ from transformer_engine_torch import DType as TE_DType from transformer_engine_torch import Float8BlockScaleTensorFormat -from ..quantized_tensor import QuantizedTensorStorage +from ...quantization import QuantizedTensorStorage from ...constants import TE_DType_To_Torch -from ..quantized_tensor import Quantizer +from ...quantization import Quantizer from ...utils import _empty_tensor diff --git a/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py index b9533edb6e..7e683f5b7f 100644 --- a/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py @@ -12,12 +12,10 @@ import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ..quantized_tensor import QuantizedTensorStorage +from ...quantization import QuantizedTensorStorage, Quantizer from ...constants import TE_DType as torch_to_transformer_engine_dtype -from ..quantized_tensor import Quantizer - from ...utils import is_non_tn_fp8_gemm_supported, _empty_tensor diff --git a/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py index c1f30146c9..a67e3b5d36 100644 --- a/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py @@ -13,11 +13,11 @@ import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ..quantized_tensor import QuantizedTensorStorage +from ...quantization import QuantizedTensorStorage from ...constants import TE_DType as torch_to_transformer_engine_dtype -from ..quantized_tensor import Quantizer +from ...quantization import Quantizer from ...utils import _empty_tensor diff --git a/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py index 350103f7ca..a5cfedb464 100644 --- a/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py @@ -16,10 +16,10 @@ # import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ..quantized_tensor import QuantizedTensorStorage +from ...quantization import QuantizedTensorStorage # from ...constants import TE_DType as torch_to_transformer_engine_dtype -from ..quantized_tensor import Quantizer +from ...quantization import Quantizer from ...utils import _empty_tensor diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index e9a299e01f..c9753c4743 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -10,7 +10,7 @@ import transformer_engine_torch as tex from transformer_engine_torch import multi_tensor_scale, multi_tensor_compute_scale_and_scale_inv -from .quantized_tensor import QuantizedTensor, Quantizer, QuantizedTensorStorage +from ..quantization import QuantizedTensor, Quantizer, QuantizedTensorStorage from .float8_tensor import Float8Tensor, Float8Quantizer, Float8CurrentScalingQuantizer from .mxfp8_tensor import MXFP8Tensor, MXFP8Quantizer from .float8_blockwise_tensor import Float8BlockwiseQTensor, Float8BlockQuantizer diff --git a/transformer_engine/pytorch/utils.py b/transformer_engine/pytorch/utils.py index b1a7e3731d..aa29f3fa03 100644 --- a/transformer_engine/pytorch/utils.py +++ b/transformer_engine/pytorch/utils.py @@ -12,7 +12,7 @@ import torch from . import torch_version -from .tensor.quantized_tensor import Quantizer +from .quantization import Quantizer from ..debug.pytorch.debug_quantization import DebugQuantizedTensor From 4378ceee97f92cf06b6e7c56ceb4308a8782fbe9 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Wed, 15 Oct 2025 15:42:25 +0000 Subject: [PATCH 03/10] update test_custom_recipe Signed-off-by: Evgeny --- tests/pytorch/test_custom_recipe.py | 36 +++++++++++++++++++++++++++++ 1 file changed, 36 insertions(+) diff --git a/tests/pytorch/test_custom_recipe.py b/tests/pytorch/test_custom_recipe.py index cb840f1971..d625c7694d 100644 --- a/tests/pytorch/test_custom_recipe.py +++ b/tests/pytorch/test_custom_recipe.py @@ -17,6 +17,42 @@ Float8CurrentScalingQuantizer, ) from transformer_engine.pytorch.module.grouped_linear import GroupedLinear +from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import nvfp4_ref_rht_2d_quantizer_factory + + +@pytest.mark.parametrize("module_type", ["Linear", "LayerNormLinear", "OpsLinear"]) +def test_custom_recipe_sanity_modules_nvfp4(module_type): + """Test modules with NVFP4 custom recipe support""" + available, reason = check_fp8_support() + if not torch.cuda.is_available() or not available: + pytest.skip(f"FP8 unsupported on this device: {reason}") + + torch.manual_seed(0) + + # Simple linear layer with dims divisible by 16 + in_features = 64 + out_features = 64 + batch = 32 + + if module_type == "Linear": + model = Linear(in_features, out_features, params_dtype=torch.bfloat16, bias=False).cuda() + elif module_type == "LayerNormLinear": + model = LayerNormLinear(in_features, out_features, params_dtype=torch.bfloat16, bias=False).cuda() + else: # OpsLinear + model = te_ops.Linear(in_features, out_features, device="cuda", dtype=torch.bfloat16, bias=False) + inp = torch.randn(batch, in_features, device="cuda", dtype=torch.bfloat16, requires_grad=True) + + # Use NVFP4 quantizer factory + custom_recipe = recipe.CustomRecipe(qfactory=nvfp4_ref_rht_2d_quantizer_factory) + + # Execute with custom recipe + with fp8_autocast(enabled=True, fp8_recipe=custom_recipe): + out = model(inp) + loss = out.float().sum() + loss.backward() + + # Basic sanity: gradients exist + assert inp.grad is not None @pytest.mark.parametrize("module_type", ["Linear", "LayerNormLinear", "OpsLinear", "LayerNormMLP"]) From 35757711d424643851ebc3ca38f6f06a0a691e3d Mon Sep 17 00:00:00 2001 From: Evgeny Date: Wed, 15 Oct 2025 15:59:34 +0000 Subject: [PATCH 04/10] Rename experimental -> custom Signed-off-by: Evgeny --- tests/pytorch/distributed/run_numerics_exact.py | 2 +- tests/pytorch/distributed/test_numerics_exact.py | 2 +- transformer_engine/pytorch/cpp_extensions/gemm.py | 10 +++++----- transformer_engine/pytorch/custom_recipes/gemm.py | 6 +++--- .../pytorch/custom_recipes/quantization_nvfp4.py | 8 ++++---- .../pytorch/module/layernorm_linear.py | 10 +++++----- transformer_engine/pytorch/module/layernorm_mlp.py | 10 +++++----- transformer_engine/pytorch/module/linear.py | 12 ++++++------ transformer_engine/pytorch/tensor/utils.py | 6 +++--- 9 files changed, 33 insertions(+), 33 deletions(-) diff --git a/tests/pytorch/distributed/run_numerics_exact.py b/tests/pytorch/distributed/run_numerics_exact.py index cea6c28264..d8577bad4c 100644 --- a/tests/pytorch/distributed/run_numerics_exact.py +++ b/tests/pytorch/distributed/run_numerics_exact.py @@ -489,7 +489,7 @@ def _test_linear(parallel_mode=None, sequence_parallel=False, **kwargs): sequence_parallel (bool): Enable sequence parallelism if True. kwargs (dict): Additional arguments for the linear layer. - QUANTIZATION options: nvfp4 <=> experimental nvfp4 as a reference + QUANTIZATION options: nvfp4 <=> custom nvfp4 as a reference """ params_dtype = torch.bfloat16 use_bias = kwargs.get("bias", True) diff --git a/tests/pytorch/distributed/test_numerics_exact.py b/tests/pytorch/distributed/test_numerics_exact.py index 890a248044..593e0e901f 100644 --- a/tests/pytorch/distributed/test_numerics_exact.py +++ b/tests/pytorch/distributed/test_numerics_exact.py @@ -14,7 +14,7 @@ Distributed numerics tests This numerical test aims for zero tolerance test for absolute confidence in numerics. - In the case of NVFP4, with the experimental NVFP4 quantization, we matched bitwise + In the case of NVFP4, with the custom NVFP4 quantization, we matched bitwise result with the native silicon. For distrbuted test cases, we can do the same by thing by comparing BF16 AG results with the low precision AG results at layer level. """ diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index b4f905a7ad..0da88f85a9 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -13,8 +13,8 @@ from ..quantization import Quantizer from ..tensor.storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage -from ..tensor.utils import is_experimental -from ..custom_recipes.gemm import experimental_gemm +from ..tensor.utils import is_custom +from ..custom_recipes.gemm import custom_gemm from ...debug.pytorch.debug_quantization import DebugQuantizer __all__ = [ @@ -79,9 +79,9 @@ def general_gemm( if not out.is_contiguous(): raise ValueError("Output tensor is not contiguous.") - # If A or B are experimental tensors -> dispatch to quantizers's qgemm implementation - if is_experimental(A) or is_experimental(B): - return experimental_gemm( + # If A or B are custom tensors -> dispatch to quantizers's qgemm implementation + if is_custom(A) or is_custom(B): + return custom_gemm( A, B, workspace, diff --git a/transformer_engine/pytorch/custom_recipes/gemm.py b/transformer_engine/pytorch/custom_recipes/gemm.py index 71d58a565a..41faf57bb5 100644 --- a/transformer_engine/pytorch/custom_recipes/gemm.py +++ b/transformer_engine/pytorch/custom_recipes/gemm.py @@ -13,10 +13,10 @@ GEMMType, ) from transformer_engine.pytorch.quantization import QuantizedTensorStorage, Quantizer -from transformer_engine.pytorch.tensor.utils import is_experimental +from transformer_engine.pytorch.tensor.utils import is_custom -def experimental_gemm( +def custom_gemm( A: QuantizedTensorStorage, B: QuantizedTensorStorage, workspace: torch.Tensor, # pylint: disable=unused-argument @@ -32,7 +32,7 @@ def experimental_gemm( grad: bool = False, ) -> Iterable[Optional[torch.Tensor]]: """Dispatch GEMM to quantizer's qgemm method.""" - assert is_experimental(A) and is_experimental(B), "A and B must be experimental tensors" + assert is_custom(A) and is_custom(B), "A and B must be custom tensors" A, B = B, A diff --git a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py index eaab4bfae7..2e9f4de053 100644 --- a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py +++ b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py @@ -229,8 +229,8 @@ class NVFP4TensorRef(QuantizedTensorStorage): _quantizer: Optional[Quantizer] = None @property - def experimental(self) -> bool: - """Flag to indicate this quantizer is using experimental Kitchen middleware.""" + def custom(self) -> bool: + """Flag to indicate this quantized tensor is custom.""" return True def prepare_for_saving( @@ -362,8 +362,8 @@ def __init__( self.with_random_sign_mask = with_random_sign_mask @property - def experimental(self) -> bool: - """Flag to indicate this quantizer is using experimental Kitchen middleware""" + def custom(self) -> bool: + """Flag to indicate this quantizer is custom.""" return True @staticmethod diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 890825bb51..1e2509c18f 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -16,7 +16,7 @@ from transformer_engine.common.recipe import Recipe from transformer_engine.pytorch import torch_version -from transformer_engine.pytorch.tensor.utils import is_experimental +from transformer_engine.pytorch.tensor.utils import is_custom from .base import ( fill_userbuffers_buffer_for_all_gather, get_workspace, @@ -194,13 +194,13 @@ def forward( # Avoid quantized norm kernel if norm output will be returned # or if a gather of ln_out must be in high precision. - experimental = is_experimental(input_quantizer) + custom = is_custom(input_quantizer) with_quantized_norm = ( fp8 and not debug and not return_layernorm_output and not return_layernorm_output_gathered - and not experimental # TODO(negvet): and not FP8GlobalStateManager.get_fp8_recipe().custom() + and not custom # TODO(negvet): and not FP8GlobalStateManager.get_fp8_recipe().custom() ) # Apply normalization @@ -246,8 +246,8 @@ def forward( quantizer = None if fp8 or debug: quantizer = input_quantizer - # experimental recipe doesn't need to support quantized AG - if not with_quantized_norm and not experimental: + # custom recipe doesn't need to support quantized AG + if not with_quantized_norm and not custom: ln_out = quantizer(ln_out) quantizer.set_usage(rowwise=True, columnwise=False) if ub_overlap_ag_fprop: # Initialize Userbuffers all-gather diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 627e37a259..0872174c05 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -17,7 +17,7 @@ from transformer_engine.common.recipe import Recipe from transformer_engine.pytorch import torch_version -from transformer_engine.pytorch.tensor.utils import is_experimental +from transformer_engine.pytorch.tensor.utils import is_custom from .base import ( fill_userbuffers_buffer_for_all_gather, get_workspace, @@ -268,13 +268,13 @@ def forward( # high precision layernorm output and output of the linear are returned # for debug: : layernorm output = High precision to enable processing of this norm - experimental = is_experimental(fc1_input_quantizer) + custom = is_custom(fc1_input_quantizer) with_quantized_norm = ( fp8 and not debug and not return_layernorm_output and not return_layernorm_output_gathered - and not experimental + and not custom ) # Apply normalization @@ -314,8 +314,8 @@ def forward( quantizer = None if fp8 or debug: quantizer = fc1_input_quantizer - # experimental recipe doesn't need to support quantized AG - if not with_quantized_norm and not experimental: + # custom recipe doesn't need to support quantized AG + if not with_quantized_norm and not custom: ln_out = fc1_input_quantizer(ln_out) fc1_input_quantizer.set_usage(rowwise=True, columnwise=False) if ub_overlap_ag: diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 6101cddfef..d7160f2180 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -66,7 +66,7 @@ ) from ..tensor.float8_tensor import Float8CurrentScalingQuantizer, Float8Quantizer from ..tensor.mxfp8_tensor import MXFP8Quantizer -from ..tensor.utils import is_experimental +from ..tensor.utils import is_custom from ..export import is_in_onnx_export_mode, assert_warmed_up from ..cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...debug.pytorch.debug_state import TEDebugState @@ -153,8 +153,8 @@ def forward( ub_obj = get_ub(ub_name + "_fprop", fp8) ub_type = tex.CommOverlapType.AG - # experimental recipe check - experimental = is_experimental(input_quantizer) or is_experimental(weight_quantizer) + # custom recipe check + custom = is_custom(input_quantizer) or is_custom(weight_quantizer) # ------------------------------------------------------ # Prepare input tensor @@ -178,7 +178,7 @@ def forward( if fp8 or debug: if input_quantizer is None: raise ValueError("Missing quantizer for input tensor") - if not isinstance(inputmat, QuantizedTensorStorage) and not experimental: + if not isinstance(inputmat, QuantizedTensorStorage) and not custom: own_quantized_input = True input_quantizer.set_usage(rowwise=True, columnwise=backward_needs_input) if isinstance( @@ -448,7 +448,7 @@ def forward( ctx.main_grad_func = lambda: weight.main_grad ctx.debug = debug - ctx.experimental = experimental + ctx.custom = custom ctx.cpu_offloading = cpu_offloading ctx.is_first_microbatch = is_first_microbatch ctx.use_bias = bias is not None @@ -616,7 +616,7 @@ def backward(ctx, grad_output: torch.Tensor) -> Tuple[Union[torch.Tensor, None], if isinstance(inputmat, QuantizedTensorStorage): # Input tensor is already quantized pass - elif ctx.debug or ctx.experimental: + elif ctx.debug or ctx.custom: # Debug quantizer will be applied immediately before wgrad GEMM pass else: diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index c9753c4743..4525f265d7 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -454,8 +454,8 @@ def _cast_master_weights_to_fp8_blockwise_scaling( ) -def is_experimental(x: Optional[Union[Quantizer, QuantizedTensorStorage]] = None) -> bool: - """Check if an object is experimental. +def is_custom(x: Optional[Union[Quantizer, QuantizedTensorStorage]] = None) -> bool: + """Check if an object is custom. Returns False if x is a torch.Tensor. """ @@ -463,4 +463,4 @@ def is_experimental(x: Optional[Union[Quantizer, QuantizedTensorStorage]] = None return False if not isinstance(x, (Quantizer, QuantizedTensorStorage)): raise AssertionError("Object must be a Quantizer or QuantizedTensorStorage instance") - return hasattr(x, "experimental") and x.experimental + return hasattr(x, "custom") and x.custom From f686e1dc72435c85174cb19e4ad62e5ea3cfd328 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Mon, 20 Oct 2025 12:59:04 +0000 Subject: [PATCH 05/10] Minor Signed-off-by: Evgeny --- .../pytorch/tensor/storage/float8_blockwise_tensor_storage.py | 2 +- transformer_engine/pytorch/tensor/utils.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py index b2e7581a0f..e75d05ce7e 100644 --- a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py @@ -14,7 +14,7 @@ from transformer_engine_torch import Float8BlockScaleTensorFormat from ...quantization_base import QuantizedTensorStorage, Quantizer -git + from ...constants import TE_DType_To_Torch from ...utils import _empty_tensor diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index 5d7b337b1c..1b079ea766 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -4,7 +4,6 @@ """Helper functions for using fp8 tensors as weights""" -import os from typing import Optional, Union import torch import transformer_engine_torch as tex From 578ed57f034d786c0de0ca6ff9230e7a03546e04 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 20 Oct 2025 13:00:53 +0000 Subject: [PATCH 06/10] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/pytorch/test_custom_recipe.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/test_custom_recipe.py b/tests/pytorch/test_custom_recipe.py index d35f65c389..6725c5fc49 100644 --- a/tests/pytorch/test_custom_recipe.py +++ b/tests/pytorch/test_custom_recipe.py @@ -18,7 +18,9 @@ ) import transformer_engine.pytorch.ops as te_ops from transformer_engine.pytorch.module.grouped_linear import GroupedLinear -from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import nvfp4_ref_rht_2d_quantizer_factory +from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import ( + nvfp4_ref_rht_2d_quantizer_factory, +) @pytest.mark.parametrize("module_type", ["Linear", "LayerNormLinear", "OpsLinear"]) @@ -38,9 +40,13 @@ def test_custom_recipe_sanity_modules_nvfp4(module_type): if module_type == "Linear": model = Linear(in_features, out_features, params_dtype=torch.bfloat16, bias=False).cuda() elif module_type == "LayerNormLinear": - model = LayerNormLinear(in_features, out_features, params_dtype=torch.bfloat16, bias=False).cuda() + model = LayerNormLinear( + in_features, out_features, params_dtype=torch.bfloat16, bias=False + ).cuda() else: # OpsLinear - model = te_ops.Linear(in_features, out_features, device="cuda", dtype=torch.bfloat16, bias=False) + model = te_ops.Linear( + in_features, out_features, device="cuda", dtype=torch.bfloat16, bias=False + ) inp = torch.randn(batch, in_features, device="cuda", dtype=torch.bfloat16, requires_grad=True) # Use NVFP4 quantizer factory From 25b3f28e1018a845b94b70e3091f5bfa7b8202ae Mon Sep 17 00:00:00 2001 From: Evgeny Date: Mon, 20 Oct 2025 14:53:27 +0000 Subject: [PATCH 07/10] Fix import Signed-off-by: Evgeny --- tests/pytorch/attention/test_attention.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 3892ab00dc..a4591f0d7e 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -45,7 +45,7 @@ ) from transformer_engine.pytorch.utils import get_cudnn_version import transformer_engine_torch as tex -from transformer_engine.pytorch.quantization import ( +from transformer_engine.pytorch.quantization_base import ( Quantizer, prepare_for_saving, restore_from_saved, From 74899524662a681b34b58ddd63d445161d92cf6a Mon Sep 17 00:00:00 2001 From: Evgeny Tsykunov Date: Tue, 21 Oct 2025 17:23:23 +0200 Subject: [PATCH 08/10] Update tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py Co-authored-by: Kirthi Shankar Sivamani Signed-off-by: Evgeny Tsykunov --- tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py b/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py index 7125b04198..904dfc2eab 100644 --- a/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py +++ b/tests/pytorch/nvfp4/test_nvfp4_rht_quantize_exact.py @@ -15,7 +15,6 @@ from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import NVFP4QuantizerRef from transformer_engine.pytorch.custom_recipes import utils from transformer_engine.pytorch.constants import TE_DType -from transformer_engine.pytorch.fp8 import fp8_autocast, get_fp4_te_dtype from transformer_engine.common.recipe import NVFP4BlockScaling import pytest From ce99997e65623807ad13bd7415f35996e79255f5 Mon Sep 17 00:00:00 2001 From: Evgeny Tsykunov Date: Tue, 21 Oct 2025 17:23:34 +0200 Subject: [PATCH 09/10] Update tests/pytorch/test_custom_recipe.py Co-authored-by: Kirthi Shankar Sivamani Signed-off-by: Evgeny Tsykunov --- tests/pytorch/test_custom_recipe.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/pytorch/test_custom_recipe.py b/tests/pytorch/test_custom_recipe.py index 6725c5fc49..64f1c3d159 100644 --- a/tests/pytorch/test_custom_recipe.py +++ b/tests/pytorch/test_custom_recipe.py @@ -17,7 +17,6 @@ Float8CurrentScalingQuantizer, ) import transformer_engine.pytorch.ops as te_ops -from transformer_engine.pytorch.module.grouped_linear import GroupedLinear from transformer_engine.pytorch.custom_recipes.quantization_nvfp4 import ( nvfp4_ref_rht_2d_quantizer_factory, ) From 77a7d5cfe0f71418572fc04753b093b14c4137b7 Mon Sep 17 00:00:00 2001 From: Evgeny Date: Wed, 22 Oct 2025 09:15:15 +0000 Subject: [PATCH 10/10] quantization_base -> quantized_tensor rename Signed-off-by: Evgeny --- tests/pytorch/attention/test_attention.py | 2 +- transformer_engine/debug/pytorch/debug_quantization.py | 2 +- transformer_engine/pytorch/__init__.py | 10 +++++----- .../attention/dot_product_attention/backends.py | 2 +- .../dot_product_attention/context_parallel.py | 4 ++-- .../pytorch/cpp_extensions/fused_attn.py | 2 +- transformer_engine/pytorch/cpp_extensions/gemm.py | 2 +- transformer_engine/pytorch/cpu_offload.py | 2 +- transformer_engine/pytorch/custom_recipes/gemm.py | 2 +- .../pytorch/custom_recipes/quantization_nvfp4.py | 2 +- transformer_engine/pytorch/distributed.py | 2 +- transformer_engine/pytorch/module/base.py | 2 +- transformer_engine/pytorch/module/grouped_linear.py | 2 +- transformer_engine/pytorch/module/layernorm_linear.py | 2 +- transformer_engine/pytorch/module/layernorm_mlp.py | 2 +- transformer_engine/pytorch/module/linear.py | 2 +- transformer_engine/pytorch/ops/_common.py | 2 +- .../pytorch/ops/fused/userbuffers_backward_linear.py | 2 +- .../pytorch/ops/fused/userbuffers_forward_linear.py | 2 +- transformer_engine/pytorch/ops/fuser.py | 2 +- transformer_engine/pytorch/permutation.py | 2 +- .../{quantization_base.py => quantized_tensor.py} | 0 transformer_engine/pytorch/tensor/__init__.py | 2 +- .../pytorch/tensor/_quantization_helpers.py | 2 +- .../pytorch/tensor/float8_blockwise_tensor.py | 2 +- transformer_engine/pytorch/tensor/float8_tensor.py | 2 +- transformer_engine/pytorch/tensor/mxfp8_tensor.py | 2 +- transformer_engine/pytorch/tensor/nvfp4_tensor.py | 2 +- .../tensor/storage/float8_blockwise_tensor_storage.py | 2 +- .../pytorch/tensor/storage/float8_tensor_storage.py | 2 +- .../pytorch/tensor/storage/mxfp8_tensor_storage.py | 2 +- .../pytorch/tensor/storage/nvfp4_tensor_storage.py | 2 +- transformer_engine/pytorch/tensor/utils.py | 2 +- transformer_engine/pytorch/utils.py | 2 +- 34 files changed, 38 insertions(+), 38 deletions(-) rename transformer_engine/pytorch/{quantization_base.py => quantized_tensor.py} (100%) diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index a4591f0d7e..3150c06abb 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -45,7 +45,7 @@ ) from transformer_engine.pytorch.utils import get_cudnn_version import transformer_engine_torch as tex -from transformer_engine.pytorch.quantization_base import ( +from transformer_engine.pytorch.quantized_tensor import ( Quantizer, prepare_for_saving, restore_from_saved, diff --git a/transformer_engine/debug/pytorch/debug_quantization.py b/transformer_engine/debug/pytorch/debug_quantization.py index df35ad0fba..7f45a24e20 100644 --- a/transformer_engine/debug/pytorch/debug_quantization.py +++ b/transformer_engine/debug/pytorch/debug_quantization.py @@ -15,7 +15,7 @@ import transformer_engine_torch as tex from transformer_engine.common.recipe import Recipe -from transformer_engine.pytorch.quantization_base import ( +from transformer_engine.pytorch.quantized_tensor import ( QuantizedTensor, Quantizer, QuantizedTensorStorage, diff --git a/transformer_engine/pytorch/__init__.py b/transformer_engine/pytorch/__init__.py index 3293e6bfd9..9d894a389b 100644 --- a/transformer_engine/pytorch/__init__.py +++ b/transformer_engine/pytorch/__init__.py @@ -66,11 +66,11 @@ def torch_version() -> tuple[int, ...]: from transformer_engine.pytorch import optimizers from transformer_engine.pytorch.export import onnx_export from transformer_engine.pytorch.cross_entropy import parallel_cross_entropy -from transformer_engine.pytorch.quantization_base import QuantizedTensorStorage -from transformer_engine.pytorch.quantization_base import QuantizedTensor -from transformer_engine.pytorch.quantization_base import Quantizer -from transformer_engine.pytorch.quantization_base import prepare_for_saving -from transformer_engine.pytorch.quantization_base import restore_from_saved +from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage +from transformer_engine.pytorch.quantized_tensor import QuantizedTensor +from transformer_engine.pytorch.quantized_tensor import Quantizer +from transformer_engine.pytorch.quantized_tensor import prepare_for_saving +from transformer_engine.pytorch.quantized_tensor import restore_from_saved from transformer_engine.pytorch.tensor import Float8Quantizer from transformer_engine.pytorch.tensor import Float8CurrentScalingQuantizer from transformer_engine.pytorch.tensor import MXFP8Quantizer diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 7e2c7c6318..6c19d868a1 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -24,7 +24,7 @@ Float8Quantizer, Float8CurrentScalingQuantizer, ) -from transformer_engine.pytorch.quantization_base import ( +from transformer_engine.pytorch.quantized_tensor import ( QuantizedTensorStorage, prepare_for_saving, restore_from_saved, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index de6942efc1..e5ee8cc7db 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -21,7 +21,7 @@ ) from transformer_engine.pytorch.quantization import FP8GlobalStateManager from transformer_engine.pytorch.tensor.float8_tensor import Float8Tensor -from transformer_engine.pytorch.quantization_base import QuantizedTensorStorage +from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage from transformer_engine.pytorch.jit import jit_fuser from transformer_engine.pytorch.constants import ( dist_group_type, @@ -33,7 +33,7 @@ gather_along_first_dim, reduce_scatter_along_first_dim, ) -from transformer_engine.pytorch.quantization_base import ( +from transformer_engine.pytorch.quantized_tensor import ( prepare_for_saving, restore_from_saved, ) diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 0d9c66a099..f80c001a1d 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -15,7 +15,7 @@ NVTE_Softmax_Type, NVTE_Fused_Attn_Backend, ) -from ..quantization_base import Quantizer +from ..quantized_tensor import Quantizer __all__ = [ diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index b0cd165dba..dd04112982 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -11,7 +11,7 @@ from ..constants import TE_DType from ..utils import get_sm_count, _empty_tensor -from ..quantization_base import Quantizer +from ..quantized_tensor import Quantizer from ..tensor.storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage from ..tensor.utils import is_custom from ..custom_recipes.gemm import custom_gemm diff --git a/transformer_engine/pytorch/cpu_offload.py b/transformer_engine/pytorch/cpu_offload.py index 869067b622..6edc126200 100644 --- a/transformer_engine/pytorch/cpu_offload.py +++ b/transformer_engine/pytorch/cpu_offload.py @@ -10,7 +10,7 @@ import torch from transformer_engine.debug.pytorch.debug_state import TEDebugState -from .quantization_base import QuantizedTensorStorage +from .quantized_tensor import QuantizedTensorStorage from .tensor.float8_tensor import Float8Tensor __all__ = ["get_cpu_offload_context"] diff --git a/transformer_engine/pytorch/custom_recipes/gemm.py b/transformer_engine/pytorch/custom_recipes/gemm.py index 12fd7c162f..cc98a8a57a 100644 --- a/transformer_engine/pytorch/custom_recipes/gemm.py +++ b/transformer_engine/pytorch/custom_recipes/gemm.py @@ -12,7 +12,7 @@ MMParams, GEMMType, ) -from transformer_engine.pytorch.quantization_base import QuantizedTensorStorage, Quantizer +from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage, Quantizer from transformer_engine.pytorch.tensor.utils import is_custom diff --git a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py index fdbfd32084..1ce9079eb1 100644 --- a/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py +++ b/transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py @@ -11,7 +11,7 @@ from transformer_engine.pytorch.custom_recipes import quantization from transformer_engine.pytorch.custom_recipes import utils -from transformer_engine.pytorch.quantization_base import QuantizedTensorStorage, Quantizer +from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage, Quantizer def nvfp4_ref_rht_2d_quantizer_factory(role): diff --git a/transformer_engine/pytorch/distributed.py b/transformer_engine/pytorch/distributed.py index cd3bdb04bc..8c14d5ab7f 100644 --- a/transformer_engine/pytorch/distributed.py +++ b/transformer_engine/pytorch/distributed.py @@ -41,7 +41,7 @@ from .tensor.mxfp8_tensor import MXFP8Quantizer from .tensor.nvfp4_tensor import NVFP4Quantizer from .tensor.float8_blockwise_tensor import Float8BlockQuantizer -from .quantization_base import QuantizedTensorStorage, QuantizedTensor, Quantizer +from .quantized_tensor import QuantizedTensorStorage, QuantizedTensor, Quantizer from .tensor.storage.float8_tensor_storage import Float8TensorStorage from .tensor.storage.mxfp8_tensor_storage import MXFP8TensorStorage from .tensor.storage.nvfp4_tensor_storage import NVFP4TensorStorage diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index e22dcbc2b7..7f571ce011 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -38,7 +38,7 @@ _fsdp_gather_tensors, ) from ..constants import dist_group_type -from ..quantization_base import QuantizedTensor, QuantizedTensorStorage, Quantizer +from ..quantized_tensor import QuantizedTensor, QuantizedTensorStorage, Quantizer from ..tensor.float8_tensor import Float8Quantizer, Float8CurrentScalingQuantizer from ..tensor.mxfp8_tensor import MXFP8Quantizer from ..tensor.float8_blockwise_tensor import Float8BlockQuantizer diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index cfacae2dd3..aae85e2cab 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -43,7 +43,7 @@ from ..cpu_offload import is_cpu_offload_enabled from ..tensor.float8_tensor import Float8CurrentScalingQuantizer, Float8Quantizer -from ..quantization_base import ( +from ..quantized_tensor import ( QuantizedTensorStorage, Quantizer, prepare_for_saving, diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 1023606288..933c7cde53 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -56,7 +56,7 @@ from ..jit import no_torch_dynamo from ..graph import is_graph_capturing from ._common import apply_normalization, noop_cat, WeightGradStore -from ..quantization_base import ( +from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, Quantizer, diff --git a/transformer_engine/pytorch/module/layernorm_mlp.py b/transformer_engine/pytorch/module/layernorm_mlp.py index 200d09bb6f..bae0f28251 100644 --- a/transformer_engine/pytorch/module/layernorm_mlp.py +++ b/transformer_engine/pytorch/module/layernorm_mlp.py @@ -70,7 +70,7 @@ from ..tensor.float8_blockwise_tensor import Float8BlockQuantizer from ._common import apply_normalization, WeightGradStore from ..cpu_offload import is_cpu_offload_enabled, mark_activation_offload -from ..quantization_base import ( +from ..quantized_tensor import ( QuantizedTensorStorage, Quantizer, prepare_for_saving, diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index bf2803e33c..ccb84e6642 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -57,7 +57,7 @@ from ..constants import GemmParallelModes, dist_group_type from ..jit import no_torch_dynamo from ..graph import is_graph_capturing -from ..quantization_base import ( +from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, Quantizer, diff --git a/transformer_engine/pytorch/ops/_common.py b/transformer_engine/pytorch/ops/_common.py index f07ab3d3ab..a07ffea43f 100644 --- a/transformer_engine/pytorch/ops/_common.py +++ b/transformer_engine/pytorch/ops/_common.py @@ -13,7 +13,7 @@ from .. import torch_version from ..quantization import FP8GlobalStateManager from ..tensor.float8_tensor import Float8Tensor -from ..quantization_base import QuantizedTensorStorage +from ..quantized_tensor import QuantizedTensorStorage from ..utils import canonicalize_dtype diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py index 18ac80d6c1..fd1820d15d 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_backward_linear.py @@ -21,7 +21,7 @@ get_ub, get_workspace, ) -from ...quantization_base import Quantizer +from ...quantized_tensor import Quantizer from ...tensor.mxfp8_tensor import MXFP8Quantizer from ...utils import canonicalize_device, canonicalize_dtype, clear_tensor_data from ..basic import BasicLinear, Bias, ReduceScatter diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index 08b89efabb..057eb576d7 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -21,7 +21,7 @@ get_workspace, _2X_ACC_FPROP, ) -from ...quantization_base import Quantizer +from ...quantized_tensor import Quantizer from ...tensor.float8_tensor import Float8Quantizer, Float8CurrentScalingQuantizer from ...tensor.storage.float8_tensor_storage import Float8TensorStorage from .._common import maybe_dequantize, is_quantized_tensor diff --git a/transformer_engine/pytorch/ops/fuser.py b/transformer_engine/pytorch/ops/fuser.py index f5dbdc476e..6026a40b65 100644 --- a/transformer_engine/pytorch/ops/fuser.py +++ b/transformer_engine/pytorch/ops/fuser.py @@ -28,7 +28,7 @@ fuse_userbuffers_backward_linear, fuse_userbuffers_forward_linear, ) -from transformer_engine.pytorch.quantization_base import ( +from transformer_engine.pytorch.quantized_tensor import ( prepare_for_saving, restore_from_saved, ) diff --git a/transformer_engine/pytorch/permutation.py b/transformer_engine/pytorch/permutation.py index 77c9c77dae..f73bc9a966 100644 --- a/transformer_engine/pytorch/permutation.py +++ b/transformer_engine/pytorch/permutation.py @@ -10,7 +10,7 @@ import transformer_engine_torch as tex import transformer_engine.pytorch.triton.permutation as triton_permutation from transformer_engine.pytorch.constants import TE_DType -from transformer_engine.pytorch.quantization_base import QuantizedTensor +from transformer_engine.pytorch.quantized_tensor import QuantizedTensor from transformer_engine.pytorch.tensor.float8_tensor import Float8Tensor from transformer_engine.pytorch.tensor.float8_blockwise_tensor import Float8BlockwiseQTensor from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Tensor diff --git a/transformer_engine/pytorch/quantization_base.py b/transformer_engine/pytorch/quantized_tensor.py similarity index 100% rename from transformer_engine/pytorch/quantization_base.py rename to transformer_engine/pytorch/quantized_tensor.py diff --git a/transformer_engine/pytorch/tensor/__init__.py b/transformer_engine/pytorch/tensor/__init__.py index e8e93aaeab..ada624a902 100644 --- a/transformer_engine/pytorch/tensor/__init__.py +++ b/transformer_engine/pytorch/tensor/__init__.py @@ -6,7 +6,7 @@ import torch -from ..quantization_base import ( +from ..quantized_tensor import ( QuantizedTensorStorage, QuantizedTensor, Quantizer, diff --git a/transformer_engine/pytorch/tensor/_quantization_helpers.py b/transformer_engine/pytorch/tensor/_quantization_helpers.py index 7bf90d49e1..2214edbff2 100644 --- a/transformer_engine/pytorch/tensor/_quantization_helpers.py +++ b/transformer_engine/pytorch/tensor/_quantization_helpers.py @@ -13,7 +13,7 @@ import torch if TYPE_CHECKING: - from transformer_engine.pytorch.quantization_base import QuantizedTensor + from transformer_engine.pytorch.quantized_tensor import QuantizedTensor class _QuantizeFunc(torch.autograd.Function): diff --git a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py index 5e4d6feeed..8054374c81 100644 --- a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py @@ -14,7 +14,7 @@ from transformer_engine.common.recipe import Float8BlockScaling, Recipe from .storage.float8_blockwise_tensor_storage import Float8BlockwiseQTensorStorage -from ..quantization_base import QuantizedTensor, Quantizer +from ..quantized_tensor import QuantizedTensor, Quantizer from ._quantization_helpers import _IdentityFunc from ..utils import devices_match, round_up_to_nearest_multiple diff --git a/transformer_engine/pytorch/tensor/float8_tensor.py b/transformer_engine/pytorch/tensor/float8_tensor.py index 53c92f707e..de112bb3fd 100644 --- a/transformer_engine/pytorch/tensor/float8_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_tensor.py @@ -14,7 +14,7 @@ from transformer_engine.common.recipe import DelayedScaling, Float8CurrentScaling, Recipe from ..utils import canonicalize_process_group, devices_match from .storage.float8_tensor_storage import Float8TensorStorage, _FromFloat8Func -from ..quantization_base import QuantizedTensor, Quantizer +from ..quantized_tensor import QuantizedTensor, Quantizer from ._quantization_helpers import _IdentityFunc from ..constants import dist_group_type diff --git a/transformer_engine/pytorch/tensor/mxfp8_tensor.py b/transformer_engine/pytorch/tensor/mxfp8_tensor.py index 239d7a2d8e..5ef5708fdb 100644 --- a/transformer_engine/pytorch/tensor/mxfp8_tensor.py +++ b/transformer_engine/pytorch/tensor/mxfp8_tensor.py @@ -17,7 +17,7 @@ from ..utils import devices_match, round_up_to_nearest_multiple from .storage.mxfp8_tensor_storage import MXFP8TensorStorage, _FromMXFP8Func -from ..quantization_base import QuantizedTensor, Quantizer +from ..quantized_tensor import QuantizedTensor, Quantizer from ._quantization_helpers import _IdentityFunc aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/nvfp4_tensor.py b/transformer_engine/pytorch/tensor/nvfp4_tensor.py index a1b13b8988..7a5f8858f2 100644 --- a/transformer_engine/pytorch/tensor/nvfp4_tensor.py +++ b/transformer_engine/pytorch/tensor/nvfp4_tensor.py @@ -22,7 +22,7 @@ ) from .storage.nvfp4_tensor_storage import NVFP4TensorStorage, _FromNVFP4Func -from ..quantization_base import QuantizedTensor, Quantizer +from ..quantized_tensor import QuantizedTensor, Quantizer from ._quantization_helpers import _IdentityFunc aten = torch.ops.aten diff --git a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py index e75d05ce7e..c2d5e8b3fa 100644 --- a/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/float8_blockwise_tensor_storage.py @@ -13,7 +13,7 @@ from transformer_engine_torch import DType as TE_DType from transformer_engine_torch import Float8BlockScaleTensorFormat -from ...quantization_base import QuantizedTensorStorage, Quantizer +from ...quantized_tensor import QuantizedTensorStorage, Quantizer from ...constants import TE_DType_To_Torch diff --git a/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py index 4b801c064a..a31f6a3799 100644 --- a/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/float8_tensor_storage.py @@ -12,7 +12,7 @@ import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ...quantization_base import QuantizedTensorStorage, Quantizer +from ...quantized_tensor import QuantizedTensorStorage, Quantizer from ...constants import TE_DType as torch_to_transformer_engine_dtype diff --git a/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py index d730aba2d9..2cca0829db 100644 --- a/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/mxfp8_tensor_storage.py @@ -13,7 +13,7 @@ import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ...quantization_base import QuantizedTensorStorage, Quantizer +from ...quantized_tensor import QuantizedTensorStorage, Quantizer from ...constants import TE_DType as torch_to_transformer_engine_dtype diff --git a/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py b/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py index 06782c2369..67543a8e2a 100644 --- a/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py +++ b/transformer_engine/pytorch/tensor/storage/nvfp4_tensor_storage.py @@ -16,7 +16,7 @@ # import transformer_engine_torch as tex from transformer_engine_torch import DType as TE_DType -from ...quantization_base import QuantizedTensorStorage, Quantizer +from ...quantized_tensor import QuantizedTensorStorage, Quantizer # from ...constants import TE_DType as torch_to_transformer_engine_dtype from ...utils import _empty_tensor diff --git a/transformer_engine/pytorch/tensor/utils.py b/transformer_engine/pytorch/tensor/utils.py index 6ca2b6a57b..8354823b32 100644 --- a/transformer_engine/pytorch/tensor/utils.py +++ b/transformer_engine/pytorch/tensor/utils.py @@ -10,7 +10,7 @@ import transformer_engine_torch as tex from transformer_engine_torch import multi_tensor_scale, multi_tensor_compute_scale_and_scale_inv -from ..quantization_base import QuantizedTensor, Quantizer, QuantizedTensorStorage +from ..quantized_tensor import QuantizedTensor, Quantizer, QuantizedTensorStorage from .float8_tensor import Float8Tensor, Float8Quantizer, Float8CurrentScalingQuantizer from .mxfp8_tensor import MXFP8Tensor, MXFP8Quantizer from .float8_blockwise_tensor import Float8BlockwiseQTensor, Float8BlockQuantizer diff --git a/transformer_engine/pytorch/utils.py b/transformer_engine/pytorch/utils.py index 98a0a85bc3..90c6289963 100644 --- a/transformer_engine/pytorch/utils.py +++ b/transformer_engine/pytorch/utils.py @@ -12,7 +12,7 @@ import torch from . import torch_version -from .quantization_base import Quantizer +from .quantized_tensor import Quantizer from ..debug.pytorch.debug_quantization import DebugQuantizedTensor