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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions modelopt/torch/quantization/qtensor/nf4_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,26 +158,26 @@ def dequantize(self, dtype: torch.dtype = None, **kwarg):
cuda_ext = get_cuda_ext()

# get kwargs
scales = kwarg["scale"]
block_sizes = kwarg["block_sizes"]
double_scale = kwarg["double_scale"]
scale_zeros = kwarg["scale_zeros"]

# unpadd the scales if needed
scales = scales.view(-1)[: (self._quantized_data.numel() * 2) // block_sizes[-1]]
# Dequantize the scales while they are still shaped (num_scale_groups, scale_block_size),
# then drop the padding. The stored int8 scales are grouped and each group has its own
# double_scale, so dividing a flattened vector by double_scale.unsqueeze(-1) would
# broadcast against the group axis instead of pairing each scale with its own group.
scales = _dequantize_scalers(kwarg["scale"], double_scale, scale_zeros, dtype).flatten()
scales = scales[: (self._quantized_data.numel() * 2) // block_sizes[-1]]

if cuda_ext and self._quantized_data.is_cuda:
# with a custom cuda kernel
scales = _dequantize_scalers(scales, double_scale, scale_zeros, dtype).flatten()
output = cuda_ext.NF4_dequantize(self._quantized_data, scales, block_sizes[-1])
return (
output.view(-1)[: np.prod(self.metadata["shape"])] # handle padding
.reshape(self.metadata["shape"])
.to(dtype)
)
else:
# de-qauntize scales
scales = _dequantize_scalers(scales, double_scale, scale_zeros, dtype).flatten()
# indexing in torch required long dtype, we may need to optimize this with customized kernels
first_half_idx = (self._quantized_data >> 4).to(torch.long)
second_half_idx = (self._quantized_data & 0x0F).to(torch.long)
Expand Down
131 changes: 131 additions & 0 deletions tests/unit/torch/quantization/test_nf4_tensor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""NF4 double quantization tests that span more than one int8 scale group.

With ``scale_bits=8`` the per-block scales are stored as int8, grouped by
``scale_block_sizes``. Each group carries its own ``double_scale``, so the scales have to be
dequantized group-wise before the zero padding at the tail is trimmed. When the input has
enough blocks to span several groups the dequantize step used to broadcast the flattened
scale vector against the group axis and blow up.
"""

from __future__ import annotations

import pytest
import torch

from modelopt.torch.quantization.config import QuantizerAttributeConfig
from modelopt.torch.quantization.nn import TensorQuantizer

# 4 scales per int8 group, so 8 blocks of 2 elements give 2 groups.
_BLOCK_SIZES = {-1: 2, "scale_bits": 8, "scale_block_sizes": {-1: 4}}


def _dequantize(x: torch.Tensor, block_sizes: dict) -> torch.Tensor:
quantizer = TensorQuantizer(
QuantizerAttributeConfig(num_bits=4, block_sizes=block_sizes, fake_quant=False)
)
# The first call does the real quantize, the second one dequantizes back.
return quantizer(quantizer(x))


@pytest.mark.parametrize(
("test_input", "test_output"),
[
# 16 elements -> 8 blocks -> 2 scale groups.
(
torch.arange(16, dtype=torch.bfloat16).view(1, 16),
torch.tensor(
[
[
0.0,
1.0,
2.1875,
3.0312,
3.6094,
5.0,
5.0625,
7.0,
9.0,
9.0,
11.0,
11.0,
13.0,
13.0,
15.0,
15.0,
]
],
dtype=torch.bfloat16,
),
),
# Same but the tail has to be padded: 15 elements -> 8 blocks -> 2 scale groups.
(
torch.arange(15, dtype=torch.bfloat16).view(1, 15),
torch.tensor(
[
[
0.0,
1.0,
2.1719,
3.0,
3.6094,
5.0,
5.0625,
7.0,
9.0,
9.0,
11.0,
11.0,
13.0,
13.0,
14.0,
]
],
dtype=torch.bfloat16,
),
),
],
)
def test_nf4_double_quant_multiple_scale_groups(test_input, test_output):
"""Dequantizing a tensor with several int8 scale groups matches the expected values."""
quantizer = TensorQuantizer(
QuantizerAttributeConfig(num_bits=4, block_sizes=_BLOCK_SIZES, fake_quant=False)
)
deq_x = quantizer(quantizer(test_input))

# Guard the premise of this test: a single group would not exercise the group axis.
assert quantizer._scale.shape[0] > 1

assert deq_x.shape == test_input.shape
assert torch.equal(deq_x, test_output)


def test_nf4_double_quant_roundtrip_wide_input():
"""A wide input spanning many scale groups round-trips without blowing up."""
torch.manual_seed(0)
block_sizes = {-1: 16, "scale_bits": 8, "scale_block_sizes": {-1: 4}}
x = torch.rand(256, 32, dtype=torch.bfloat16)

quantizer = TensorQuantizer(
QuantizerAttributeConfig(num_bits=4, block_sizes=block_sizes, fake_quant=False)
)
deq_x = quantizer(quantizer(x))

# 512 blocks / 4 scales per group -> 128 groups.
assert quantizer._scale.shape[0] == 128
assert deq_x.shape == x.shape
assert torch.allclose(deq_x, x, rtol=1e-1, atol=1e-1)