From 9e74533d3644d4c247402e104aea48c690a36ca2 Mon Sep 17 00:00:00 2001 From: Dingqing Yang Date: Mon, 17 Aug 2026 16:12:06 -0700 Subject: [PATCH 1/4] Preserve quantized tensor subclasses on detach Signed-off-by: Dingqing Yang --- tests/pytorch/test_hybrid_quantization.py | 19 ++++++++ tests/pytorch/test_identity_quantizer.py | 18 ++++++++ tests/pytorch/test_quantized_tensor.py | 43 +++++++++++++++++++ .../pytorch/quantized_tensor.py | 2 +- .../pytorch/tensor/float8_blockwise_tensor.py | 2 +- .../pytorch/tensor/float8_tensor.py | 2 +- .../pytorch/tensor/hybrid_tensor.py | 2 +- .../pytorch/tensor/identity_tensor.py | 2 +- .../pytorch/tensor/mxfp8_tensor.py | 3 +- .../pytorch/tensor/nvfp4_tensor.py | 3 +- 10 files changed, 87 insertions(+), 9 deletions(-) diff --git a/tests/pytorch/test_hybrid_quantization.py b/tests/pytorch/test_hybrid_quantization.py index 74ec0a05ec..e7cd538c13 100644 --- a/tests/pytorch/test_hybrid_quantization.py +++ b/tests/pytorch/test_hybrid_quantization.py @@ -1738,6 +1738,25 @@ def test_detach(self, hybrid_tensor): assert isinstance(detached, HybridQuantizedTensor) assert not detached.requires_grad + def test_detach_preserves_subclass(self, hybrid_tensor): + """HybridQuantizedTensor detach preserves its runtime subclass.""" + + class DerivedHybridQuantizedTensor(HybridQuantizedTensor): + pass + + hybrid_tensor.__class__ = DerivedHybridQuantizedTensor + source_data = hybrid_tensor.get_data_tensors() + + detached = hybrid_tensor.detach() + + assert type(detached) is DerivedHybridQuantizedTensor + for detached_data, source_data_tensor in zip(detached.get_data_tensors(), source_data): + assert detached_data is source_data_tensor + assert not detached.requires_grad + + parameter = torch.nn.Parameter(hybrid_tensor) + assert type(parameter) is DerivedHybridQuantizedTensor + def test_repr(self, hybrid_tensor): r = repr(hybrid_tensor) assert "HybridQuantizedTensor" in r diff --git a/tests/pytorch/test_identity_quantizer.py b/tests/pytorch/test_identity_quantizer.py index cb0785e3e0..ed5327855b 100644 --- a/tests/pytorch/test_identity_quantizer.py +++ b/tests/pytorch/test_identity_quantizer.py @@ -237,6 +237,24 @@ def test_quantize_returns_identity_tensor(self): out = IdentityQuantizer()(x) assert isinstance(out, IdentityTensor) + def test_detach_preserves_subclass(self): + """IdentityTensor detach preserves its runtime subclass.""" + + class DerivedIdentityTensor(IdentityTensor): + pass + + tensor = IdentityQuantizer()(torch.randn(8, 16, device="cuda", dtype=torch.bfloat16)) + tensor.__class__ = DerivedIdentityTensor + + detached = tensor.detach() + + assert type(detached) is DerivedIdentityTensor + assert detached._hp_data is tensor._hp_data + assert not detached.requires_grad + + parameter = torch.nn.Parameter(tensor) + assert type(parameter) is DerivedIdentityTensor + def test_internal_returns_storage(self): x = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16) q = IdentityQuantizer() diff --git a/tests/pytorch/test_quantized_tensor.py b/tests/pytorch/test_quantized_tensor.py index 4bcaacac90..035b35dcf7 100644 --- a/tests/pytorch/test_quantized_tensor.py +++ b/tests/pytorch/test_quantized_tensor.py @@ -563,6 +563,49 @@ def setup_class(cls) -> None: torch.manual_seed(seed) torch.cuda.manual_seed(seed) + @pytest.mark.parametrize( + "quantization", + [ + pytest.param( + "fp8", + marks=pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8), + ), + pytest.param( + "fp8_blockwise", + marks=pytest.mark.skipif( + not fp8_block_scaling_available, + reason=reason_for_no_fp8_block_scaling, + ), + ), + pytest.param( + "mxfp8", + marks=pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8), + ), + pytest.param( + "nvfp4", + marks=pytest.mark.skipif(not nvfp4_available, reason=reason_for_no_nvfp4), + ), + ], + ) + def test_detach_preserves_subclass(self, quantization: str) -> None: + """Detaching a quantized tensor preserves its runtime subclass.""" + quantizer = make_quantizer(quantization) + tensor = quantizer(torch.randn(128, 128, dtype=torch.bfloat16, device="cuda")) + derived_type = type(f"Derived{type(tensor).__name__}", (type(tensor),), {}) + tensor.__class__ = derived_type + + detached = tensor.detach() + + assert type(detached) is derived_type + for detached_data, source_data in zip( + detached.get_data_tensors(), tensor.get_data_tensors() + ): + assert detached_data is source_data + assert not detached.requires_grad + + parameter = torch.nn.Parameter(tensor) + assert type(parameter) is derived_type + @pytest.mark.parametrize("op", ("clone", "view", "reshape", "contiguous")) @pytest.mark.parametrize("quantization", _quantization_list) def test_identity_op( diff --git a/transformer_engine/pytorch/quantized_tensor.py b/transformer_engine/pytorch/quantized_tensor.py index a2e57277d1..deb13a13ed 100644 --- a/transformer_engine/pytorch/quantized_tensor.py +++ b/transformer_engine/pytorch/quantized_tensor.py @@ -798,7 +798,7 @@ def detach(self) -> QuantizedTensor: """Create new quantized tensor with same data Output tensor must be detached from the current autograd - graph. + graph and have the same runtime type as ``self``. """ raise NotImplementedError( diff --git a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py index d1a2b488ff..af9de1209d 100644 --- a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py @@ -362,7 +362,7 @@ def dequantize(self, *, dtype: Optional[torch.dtype] = None) -> torch.Tensor: def detach(self) -> Float8BlockwiseQTensor: # pylint: disable=missing-function-docstring - return Float8BlockwiseQTensor.make_like(self) + return self.__class__.make_like(self) def clone(self) -> Float8BlockwiseQTensor: # pylint: disable=missing-function-docstring diff --git a/transformer_engine/pytorch/tensor/float8_tensor.py b/transformer_engine/pytorch/tensor/float8_tensor.py index 5c31022123..eaa546747c 100644 --- a/transformer_engine/pytorch/tensor/float8_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_tensor.py @@ -521,7 +521,7 @@ def quantize_( def detach(self) -> Float8Tensor: # pylint: disable=missing-function-docstring - return Float8Tensor.make_like(self) + return self.__class__.make_like(self) def clone(self) -> Float8Tensor: # pylint: disable=missing-function-docstring diff --git a/transformer_engine/pytorch/tensor/hybrid_tensor.py b/transformer_engine/pytorch/tensor/hybrid_tensor.py index dc65c9894b..8df2ec8b4b 100644 --- a/transformer_engine/pytorch/tensor/hybrid_tensor.py +++ b/transformer_engine/pytorch/tensor/hybrid_tensor.py @@ -469,7 +469,7 @@ def detach(self) -> HybridQuantizedTensor: "HybridQuantizedTensor.detach() does not support storage-only " f"columnwise sub-storage {col_cls.__name__}" ) - return HybridQuantizedTensor( + return self.__class__( shape=self.shape, dtype=self.dtype, rowwise_storage=row, diff --git a/transformer_engine/pytorch/tensor/identity_tensor.py b/transformer_engine/pytorch/tensor/identity_tensor.py index 9fb980a755..8310afc653 100644 --- a/transformer_engine/pytorch/tensor/identity_tensor.py +++ b/transformer_engine/pytorch/tensor/identity_tensor.py @@ -266,7 +266,7 @@ def _wrap_data_view( self, data: torch.Tensor, *, requires_grad: Optional[bool] = None ) -> "IdentityTensor": requires_grad = self.requires_grad if requires_grad is None else requires_grad - return IdentityTensor( + return self.__class__( shape=data.shape, dtype=self.dtype, hp_data=data, diff --git a/transformer_engine/pytorch/tensor/mxfp8_tensor.py b/transformer_engine/pytorch/tensor/mxfp8_tensor.py index 54cb281bd6..6cf5fba15c 100644 --- a/transformer_engine/pytorch/tensor/mxfp8_tensor.py +++ b/transformer_engine/pytorch/tensor/mxfp8_tensor.py @@ -323,8 +323,7 @@ def quantize_( def detach(self) -> MXFP8Tensor: # pylint: disable=missing-function-docstring - # TODO(ksivamani): Fix the detach bug - return MXFP8Tensor.make_like(self) + return self.__class__.make_like(self) def clone(self) -> MXFP8Tensor: # pylint: disable=missing-function-docstring diff --git a/transformer_engine/pytorch/tensor/nvfp4_tensor.py b/transformer_engine/pytorch/tensor/nvfp4_tensor.py index 5589e200ea..0572f53adb 100644 --- a/transformer_engine/pytorch/tensor/nvfp4_tensor.py +++ b/transformer_engine/pytorch/tensor/nvfp4_tensor.py @@ -530,8 +530,7 @@ def quantize_( def detach(self) -> NVFP4Tensor: # pylint: disable=missing-function-docstring - # TODO(ksivamani): Fix the detach bug - return NVFP4Tensor.make_like(self) + return self.__class__.make_like(self) def clone(self) -> NVFP4Tensor: # pylint: disable=missing-function-docstring From 865fda0d5a123f69c87156fb8232e974e223f3f7 Mon Sep 17 00:00:00 2001 From: Dingqing Yang Date: Mon, 17 Aug 2026 16:45:39 -0700 Subject: [PATCH 2/4] Fix IdentityTensor detach alias assertion Signed-off-by: Dingqing Yang --- tests/pytorch/test_identity_quantizer.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/test_identity_quantizer.py b/tests/pytorch/test_identity_quantizer.py index ed5327855b..e74675a4ef 100644 --- a/tests/pytorch/test_identity_quantizer.py +++ b/tests/pytorch/test_identity_quantizer.py @@ -249,7 +249,9 @@ class DerivedIdentityTensor(IdentityTensor): detached = tensor.detach() assert type(detached) is DerivedIdentityTensor - assert detached._hp_data is tensor._hp_data + assert detached._hp_data.data_ptr() == tensor._hp_data.data_ptr() + assert detached._hp_data.stride() == tensor._hp_data.stride() + assert detached._hp_data.storage_offset() == tensor._hp_data.storage_offset() assert not detached.requires_grad parameter = torch.nn.Parameter(tensor) From bf64b4e8b2985ce7ff394b7f3cb240e764b24a3a Mon Sep 17 00:00:00 2001 From: Dingqing Yang Date: Mon, 17 Aug 2026 19:55:30 -0700 Subject: [PATCH 3/4] Support quantized tensor subclasses in C++ bindings Signed-off-by: Dingqing Yang --- tests/pytorch/test_quantized_tensor.py | 37 ++++++++++++++++++++++++ transformer_engine/pytorch/csrc/pybind.h | 13 +++++---- 2 files changed, 45 insertions(+), 5 deletions(-) diff --git a/tests/pytorch/test_quantized_tensor.py b/tests/pytorch/test_quantized_tensor.py index 035b35dcf7..4e64215cdf 100644 --- a/tests/pytorch/test_quantized_tensor.py +++ b/tests/pytorch/test_quantized_tensor.py @@ -606,6 +606,43 @@ def test_detach_preserves_subclass(self, quantization: str) -> None: parameter = torch.nn.Parameter(tensor) assert type(parameter) is derived_type + @pytest.mark.parametrize( + "quantization", + [ + pytest.param( + "fp8", + marks=pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8), + ), + pytest.param( + "fp8_blockwise", + marks=pytest.mark.skipif( + not fp8_block_scaling_available, + reason=reason_for_no_fp8_block_scaling, + ), + ), + pytest.param( + "mxfp8", + marks=pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8), + ), + pytest.param( + "nvfp4", + marks=pytest.mark.skipif(not nvfp4_available, reason=reason_for_no_nvfp4), + ), + ], + ) + def test_update_quantized_supports_subclass(self, quantization: str) -> None: + """Quantizers update derived quantized tensor wrappers in-place.""" + quantizer = make_quantizer(quantization) + source = torch.randn(128, 128, dtype=torch.bfloat16, device="cuda") + tensor = quantizer(source) + derived_type = type(f"Derived{type(tensor).__name__}", (type(tensor),), {}) + tensor.__class__ = derived_type + + quantizer.update_quantized(torch.zeros_like(source), tensor) + + assert type(tensor) is derived_type + torch.testing.assert_close(tensor.dequantize(), torch.zeros_like(source)) + @pytest.mark.parametrize("op", ("clone", "view", "reshape", "contiguous")) @pytest.mark.parametrize("quantization", _quantization_list) def test_identity_op( diff --git a/transformer_engine/pytorch/csrc/pybind.h b/transformer_engine/pytorch/csrc/pybind.h index 9e640537f9..3dcb7005ad 100644 --- a/transformer_engine/pytorch/csrc/pybind.h +++ b/transformer_engine/pytorch/csrc/pybind.h @@ -57,13 +57,15 @@ inline bool IsFloat8CurrentScalingQuantizers(PyObject *obj) { } inline bool IsFloat8Tensor(PyObject *obj) { - return Py_TYPE(obj) == Float8TensorPythonClass || Py_TYPE(obj) == Float8TensorStoragePythonClass; + return PyObject_TypeCheck(obj, Float8TensorPythonClass) || + PyObject_TypeCheck(obj, Float8TensorStoragePythonClass); } inline bool IsMXFP8Quantizers(PyObject *obj) { return Py_TYPE(obj) == MXFP8QuantizerClass; } inline bool IsMXFP8Tensor(PyObject *obj) { - return Py_TYPE(obj) == MXFP8TensorPythonClass || Py_TYPE(obj) == MXFP8TensorStoragePythonClass; + return PyObject_TypeCheck(obj, MXFP8TensorPythonClass) || + PyObject_TypeCheck(obj, MXFP8TensorStoragePythonClass); } inline bool IsFloat8BlockwiseQuantizers(PyObject *obj) { @@ -73,12 +75,13 @@ inline bool IsFloat8BlockwiseQuantizers(PyObject *obj) { inline bool IsNVFP4Quantizers(PyObject *obj) { return Py_TYPE(obj) == NVFP4QuantizerClass; } inline bool IsFloat8BlockwiseQTensor(PyObject *obj) { - return Py_TYPE(obj) == Float8BlockwiseQTensorPythonClass || - Py_TYPE(obj) == Float8BlockwiseQTensorStoragePythonClass; + return PyObject_TypeCheck(obj, Float8BlockwiseQTensorPythonClass) || + PyObject_TypeCheck(obj, Float8BlockwiseQTensorStoragePythonClass); } inline bool IsNVFP4Tensor(PyObject *obj) { - return Py_TYPE(obj) == NVFP4TensorPythonClass || Py_TYPE(obj) == NVFP4TensorStoragePythonClass; + return PyObject_TypeCheck(obj, NVFP4TensorPythonClass) || + PyObject_TypeCheck(obj, NVFP4TensorStoragePythonClass); } TensorWrapper NVTETensorFromFloat8Tensor(py::handle tensor, Quantizer *quantizer); From ceb1c8bb23c726ce2b003e3d81c7ae1cf2b98b6b Mon Sep 17 00:00:00 2001 From: Dingqing Yang Date: Tue, 18 Aug 2026 13:49:29 -0700 Subject: [PATCH 4/4] Centralize quantized tensor detach Signed-off-by: Dingqing Yang --- transformer_engine/pytorch/quantized_tensor.py | 7 ++----- .../pytorch/tensor/float8_blockwise_tensor.py | 4 ---- transformer_engine/pytorch/tensor/float8_tensor.py | 4 ---- transformer_engine/pytorch/tensor/mxfp8_tensor.py | 4 ---- transformer_engine/pytorch/tensor/nvfp4_tensor.py | 4 ---- 5 files changed, 2 insertions(+), 21 deletions(-) diff --git a/transformer_engine/pytorch/quantized_tensor.py b/transformer_engine/pytorch/quantized_tensor.py index deb13a13ed..7149a5a163 100644 --- a/transformer_engine/pytorch/quantized_tensor.py +++ b/transformer_engine/pytorch/quantized_tensor.py @@ -797,13 +797,10 @@ def quantize_(self, tensor: torch.Tensor) -> QuantizedTensor: def detach(self) -> QuantizedTensor: """Create new quantized tensor with same data - Output tensor must be detached from the current autograd - graph and have the same runtime type as ``self``. + Output tensor must be detached from the current autograd graph. """ - raise NotImplementedError( - f"{self.__class__.__name__} class does not implement detach function" - ) + return type(self).make_like(self) def clear(self): """Deallocate this tensor's memory. Typically not needed and must be used carefully""" diff --git a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py index af9de1209d..6105b18a73 100644 --- a/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_blockwise_tensor.py @@ -360,10 +360,6 @@ def dequantize(self, *, dtype: Optional[torch.dtype] = None) -> torch.Tensor: return _FromFloat8BlockwiseFunc.apply(self, dequant_dtype) return _FromFloat8BlockwiseFunc.forward(None, self, dequant_dtype) - def detach(self) -> Float8BlockwiseQTensor: - # pylint: disable=missing-function-docstring - return self.__class__.make_like(self) - def clone(self) -> Float8BlockwiseQTensor: # pylint: disable=missing-function-docstring rowwise_data = None diff --git a/transformer_engine/pytorch/tensor/float8_tensor.py b/transformer_engine/pytorch/tensor/float8_tensor.py index eaa546747c..436027732f 100644 --- a/transformer_engine/pytorch/tensor/float8_tensor.py +++ b/transformer_engine/pytorch/tensor/float8_tensor.py @@ -519,10 +519,6 @@ def quantize_( return self.quantize_(tensor.dequantize(), noop_flag=noop_flag) return super().quantize_(tensor, noop_flag=noop_flag) - def detach(self) -> Float8Tensor: - # pylint: disable=missing-function-docstring - return self.__class__.make_like(self) - def clone(self) -> Float8Tensor: # pylint: disable=missing-function-docstring # ``_data`` may be None for columnwise-only sub-storages of a diff --git a/transformer_engine/pytorch/tensor/mxfp8_tensor.py b/transformer_engine/pytorch/tensor/mxfp8_tensor.py index 6cf5fba15c..267806e43e 100644 --- a/transformer_engine/pytorch/tensor/mxfp8_tensor.py +++ b/transformer_engine/pytorch/tensor/mxfp8_tensor.py @@ -321,10 +321,6 @@ def quantize_( return self.quantize_(tensor.dequantize()) return super().quantize_(tensor, noop_flag=noop_flag) - def detach(self) -> MXFP8Tensor: - # pylint: disable=missing-function-docstring - return self.__class__.make_like(self) - def clone(self) -> MXFP8Tensor: # pylint: disable=missing-function-docstring # _rowwise_data may be None for columnwise-only sub-storages (hybrid quantization) diff --git a/transformer_engine/pytorch/tensor/nvfp4_tensor.py b/transformer_engine/pytorch/tensor/nvfp4_tensor.py index 0572f53adb..5e537c3ff0 100644 --- a/transformer_engine/pytorch/tensor/nvfp4_tensor.py +++ b/transformer_engine/pytorch/tensor/nvfp4_tensor.py @@ -528,10 +528,6 @@ def quantize_( self._get_quantizer().update_quantized(tensor, self, noop_flag=noop_flag) return self - def detach(self) -> NVFP4Tensor: - # pylint: disable=missing-function-docstring - return self.__class__.make_like(self) - def clone(self) -> NVFP4Tensor: # pylint: disable=missing-function-docstring assert self._rowwise_data is not None