Files
kefu/wechat_rpa/.venv/Lib/site-packages/onnxruntime/quantization/cuda_quantizer.py
T
2026-09-21 10:34:06 +08:00

644 lines
26 KiB
Python

# -------------------------------------------------------------------------
# Copyright (c) Microsoft Corporation. All rights reserved.
# Licensed under the MIT License. See License.txt in the project root for
# license information.
# --------------------------------------------------------------------------
"""CUDA weight-only quantization helpers.
This module contains small Python utilities for producing the weight layouts
consumed by CUDA weight-only kernels. The blockwise quantizers wrap the same C++
pybind entry points used by runtime prepacking, and the mixed-GEMM weight packer
is a PyTorch reimplementation of the runtime CUDA packing, so tests and model
builders can generate byte-identical quantized weights. The PyTorch packer runs
on CUDA when a device is available and falls back to CPU otherwise, which is the
only option on platforms where the standalone CUDA packer is not built (Windows).
A GPU-gated parity test validates it against that standalone CUDA packer.
Two storage families are exposed:
* raw MatMulNBits blockwise storage, laid out by output channel as ``[N, K/pack]``;
* CUDA mixed-GEMM prepacked storage. MatMulNBits prepacked initializers keep the
schema shape ``[N, K/block_size, block_size*bits/8]`` and require the node
attribute ``weight_prepacked=1``. QMoE/CUTLASS callers use the kernel-facing
shape ``[K, N/pack]``.
All public helpers take one logical expert weight matrix with shape ``[N, K]``.
"""
from __future__ import annotations
import functools
import logging
from typing import TYPE_CHECKING
import numpy as np
_logger = logging.getLogger(__name__)
if TYPE_CHECKING:
import torch
def _get_torch():
"""Import torch lazily so importing onnxruntime.quantization does not require torch."""
try:
import torch # noqa: PLC0415
except ImportError as e:
raise ImportError("CUDA weight-only quantization requires torch. Please install torch to use it.") from e
return torch
def _get_pack_weights_for_cuda_mixed_gemm():
"""Return the standalone CUDA mixed-GEMM weight packer (parity oracle).
Production packing uses the PyTorch implementation (``_pack_weights_for_cuda_mixed_gemm``).
This standalone packer lives in ``onnxruntime.capi.onnxruntime_cuda_quant_preprocess``, a
separate extension module that links the CUDA runtime (built only on non-Windows CUDA
builds). It is imported lazily here (never at ``import onnxruntime`` time) and is used by
the parity test to validate the PyTorch packer byte-for-byte.
"""
try:
from onnxruntime.capi import onnxruntime_cuda_quant_preprocess as _cuda_quant # noqa: PLC0415
except ImportError as e:
raise ImportError(
"The standalone CUDA weight packer (onnxruntime_cuda_quant_preprocess) is unavailable; "
"it is built only on non-Windows onnxruntime-gpu CUDA builds."
) from e
try:
return _cuda_quant.pack_weights_for_cuda_mixed_gemm
except AttributeError as e:
raise ImportError("onnxruntime_cuda_quant_preprocess is missing pack_weights_for_cuda_mixed_gemm.") from e
def has_cuda_weight_prepacking() -> bool:
"""Return True if mixed-GEMM weight prepacking is available.
Prepacking is implemented with PyTorch (CUDA when available, CPU otherwise), so it is
available whenever torch is importable. Callers use this to skip prepack code paths
(and tests) when torch is unavailable.
"""
try:
_get_torch()
except ImportError:
return False
return True
@functools.lru_cache(maxsize=1)
def _warn_cpu_prepack_once() -> None:
_logger.warning(
"CUDA device is not available; packing mixed-GEMM weights on CPU with PyTorch. "
"This is correct but significantly slower for large Mixture-of-Experts models. "
"Pack on a CUDA-enabled machine for best performance."
)
def _prepack_device():
"""Pick the torch device for mixed-GEMM weight packing (CUDA if available, else CPU)."""
torch = _get_torch()
if torch.cuda.is_available():
return torch.device("cuda")
_warn_cpu_prepack_once()
return torch.device("cpu")
def _preprocess_weights_for_mixed_gemm_torch(tensor, bits: int, sm: int):
"""PyTorch port of the runtime CUDA ``preprocess_weights_for_mixed_gemm``.
``tensor`` is a signed int8 weight in ``(K, N/pack)`` packed row-major layout on any
device. Returns the CUTLASS mixed-GEMM layout with the same shape/dtype/device. This
mirrors ``preprocess_weights_for_mixed_gemm_cuda`` (permute_B_rows -> subbyte_transpose
-> interleave_column_major -> add_bias_and_interleave) so its output is byte-identical
to the standalone CUDA packer, for both the SM80 (Ampere) and SM90 (Hopper) layouts.
"""
torch = _get_torch()
bits_a = 16 # fp16/bf16 activations
bits_b = 4 if bits == 4 else 8
if tensor.dim() == 2:
tensor = tensor.unsqueeze(0)
permutation_map = {
"16_8": [0, 1, 8, 9, 2, 3, 10, 11, 4, 5, 12, 13, 6, 7, 14, 15],
"16_4": [
0,
1,
8,
9,
16,
17,
24,
25,
2,
3,
10,
11,
18,
19,
26,
27,
4,
5,
12,
13,
20,
21,
28,
29,
6,
7,
14,
15,
22,
23,
30,
31,
],
}
mma_shape_n = 8
b_rows_per_mma = 8 * 16 // bits_b
num_experts, num_rows, num_cols = tensor.shape[0], tensor.shape[1], tensor.shape[2]
if num_rows % b_rows_per_mma != 0 or num_cols % mma_shape_n != 0:
raise ValueError(
f"weight shape (rows={num_rows}, packed_cols={num_cols}) is incompatible with mixed-GEMM "
f"packing (rows must be a multiple of {b_rows_per_mma}, packed cols a multiple of {mma_shape_n})."
)
# permute_B_rows_for_mixed_gemm
if sm < 100:
pmap = permutation_map[f"{bits_a}_{bits_b}"]
row_idx = [(r // b_rows_per_mma) * b_rows_per_mma + pmap[r % b_rows_per_mma] for r in range(num_rows)]
tensor = tensor[:, row_idx, :]
# subbyte_transpose
original_shape = tensor.shape
if bits_b == 4:
u = tensor.view(torch.uint8)
high = (u >> 4).permute(0, 2, 1).unsqueeze(2)
low = ((u << 4) >> 4).permute(0, 2, 1).unsqueeze(2)
merged = torch.cat([low, high], dim=2).reshape(u.shape[0], -1, u.shape[1])
merged = merged[:, :, 0::2] + merged[:, :, 1::2] * 16
tensor = merged.view(torch.int8).reshape(original_shape)
else:
tensor = tensor.permute(0, 2, 1).reshape(original_shape)
# interleave_column_major_tensor
interleave = bits_a // bits_b
if interleave > 1 and sm < 90:
rows_per_tile = 128 * 8 // bits_a
elts_in_int32 = 32 // bits_b
if num_rows % elts_in_int32 != 0 or num_rows % rows_per_tile != 0:
raise ValueError(f"num_rows ({num_rows}) is incompatible with column-interleave tiling.")
tensor = tensor.reshape(
num_experts, -1, interleave, num_rows // rows_per_tile, rows_per_tile * 4 // elts_in_int32
)
tensor = tensor.permute(0, 1, 3, 2, 4).reshape(original_shape)
# add_bias_and_interleave_quantized_tensor_inplace
if bits_b == 8:
t = tensor.to(torch.int64) # widen so the +128 rebias cannot overflow int8
t += -256 * (t > 127).to(torch.int64) + 128
t = t.reshape(-1, 4)[:, [0, 2, 1, 3]].reshape(original_shape)
tensor = t.to(torch.uint8).view(torch.int8)
else:
u = tensor.view(torch.uint8)
high = (u >> 4).unsqueeze(-1)
low = ((u << 4) >> 4).unsqueeze(-1)
merged = torch.cat([low, high], dim=-1).reshape(u.shape[0], u.shape[1], -1)
merged = merged.reshape(-1, 8)[:, [0, 2, 4, 6, 1, 3, 5, 7]].reshape(merged.shape)
merged = merged.to(torch.int16)
merged += -16 * (merged > 7).to(torch.int16) + 8
merged = merged[:, :, 0::2] + merged[:, :, 1::2] * 16
tensor = merged.to(torch.uint8).view(torch.int8)
return tensor.squeeze(0).contiguous()
def _pack_weights_for_cuda_mixed_gemm(q_weights, n: int, k: int, bits: int, force_arch: int = 80) -> np.ndarray:
"""PyTorch implementation of the CUDA ``pack_weights_for_cuda_mixed_gemm``.
``q_weights`` is ORT's unsigned MatMulNBits/QMoE storage ``(N, K/pack)`` (uint8). Returns
a flat ``int8`` numpy array with the CUTLASS mixed-GEMM layout, byte-identical to the
standalone CUDA packer. Runs on CUDA when available, otherwise on CPU.
"""
torch = _get_torch()
bits = int(bits)
force_arch = int(force_arch)
if bits not in (4, 8):
raise ValueError(f"bits must be 4 or 8, got {bits}.")
if force_arch not in (80, 90):
raise ValueError(f"force_arch must be 80 (SM80) or 90 (SM90), got {force_arch}.")
pack = 8 // bits
device = _prepack_device()
q = torch.as_tensor(np.ascontiguousarray(q_weights)).view(torch.uint8).reshape(n, k // pack).to(device)
# Front-end adaptor: transpose ORT (N, K) -> (K, N) and convert unsigned -> signed int8.
if bits == 4:
low = (q & 0x0F).to(torch.int16)
high = (q >> 4).to(torch.int16)
unpacked = torch.empty((n, k), dtype=torch.int16, device=device)
unpacked[:, 0::2] = low
unpacked[:, 1::2] = high
signed_t = (unpacked - 8).transpose(0, 1).contiguous() # (K, N), zero point 8
packed_t = ((signed_t[:, 0::2] & 0x0F) | ((signed_t[:, 1::2] & 0x0F) << 4)).to(torch.uint8).view(torch.int8)
else:
signed_t = (q.to(torch.int16) - 128).transpose(0, 1).contiguous() # (K, N), zero point 128
packed_t = signed_t.to(torch.uint8).view(torch.int8)
out = _preprocess_weights_for_mixed_gemm_torch(packed_t.contiguous(), bits, force_arch)
return out.reshape(-1).cpu().numpy()
def _get_quantize_matmul_nbits():
"""Return MatMulNBits blockwise quantizers from the ORT pybind module."""
try:
from onnxruntime.capi._pybind_state import ( # noqa: PLC0415
quantize_matmul_4bits,
quantize_matmul_8bits,
)
except ImportError as e:
raise ImportError(
"CUDA blockwise quantization requires quantize_matmul_4bits and quantize_matmul_8bits from onnxruntime."
) from e
return quantize_matmul_4bits, quantize_matmul_8bits
class CudaQuantizer:
"""CUDA quantizer utilities for MoE/QMoE and MatMulNBits-style weight-only kernels.
The methods are stateless; callers may use the class directly without
constructing an object.
"""
@staticmethod
def qmoe_symmetric_per_channel_quantize(
weights: torch.Tensor,
bits: int,
*,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize one QMoE expert with symmetric per-channel storage.
``weights`` has logical shape ``[N, K]``. Returns raw QMoE storage
``[N, K/pack]`` and scales ``[N]``. By default, this emits the ORT CUDA
QMoE storage contract: unsigned bytes/nibbles with an implicit zero-point
offset, so each stored value is ``q + zero_point`` even though the numeric
quantization is symmetric. By default it uses the full ``[-8, 7]`` /
``[-128, 127]`` range. Set ``unsigned_full_range=False`` to use the legacy
``[-7, 7]`` / ``[-127, 127]`` range.
"""
torch = _get_torch()
weights = weights.detach().cpu().to(torch.float32).contiguous()
bits = int(bits)
if bits not in (4, 8):
raise ValueError(f"QMoE per-channel quantization only supports 4 or 8 bits, got {bits}.")
n, k = weights.shape
pack = 8 // bits
if k % pack != 0:
raise ValueError(f"K ({k}) must be divisible by {pack} for QMoE per-channel quantization.")
if bits == 4:
if unsigned_full_range:
qmin, qmax, scale_divisor, zero_point = -8, 7, 8, 8
else:
qmin, qmax, scale_divisor, zero_point = -7, 7, 7, 8
else: # bits == 8, already validated above
if unsigned_full_range:
qmin, qmax, scale_divisor, zero_point = -128, 127, 128, 128
else:
qmin, qmax, scale_divisor, zero_point = -127, 127, 127, 128
scales = weights.abs().amax(dim=1, keepdim=True) / float(scale_divisor)
scales = torch.clamp(scales, min=torch.finfo(torch.float32).eps)
quantized = torch.clamp(torch.round(weights / scales), qmin, qmax).to(torch.int16).contiguous()
quantized = (quantized + zero_point).to(torch.uint8)
if bits == 4:
qweight = (quantized[:, 0::2] & 0xF) | ((quantized[:, 1::2] & 0xF) << 4)
qweight = qweight.to(torch.uint8)
else:
qweight = quantized
return qweight.contiguous(), scales.squeeze(-1).contiguous()
@staticmethod
def qmoe_per_channel_quantize(
weights: torch.Tensor,
bits: int,
prepack: bool,
force_arch: int = 80,
*,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize per-channel QMoE weights and optionally CUTLASS-prepack them.
When ``prepack`` is true, returned weights have shape ``[K, N/pack]``.
Otherwise, returned weights keep raw per-channel storage ``[N, K/pack]``.
Prepacking uses PyTorch (CUDA when available, CPU otherwise).
"""
torch = _get_torch()
qweight, scales = CudaQuantizer.qmoe_symmetric_per_channel_quantize(
weights,
bits,
unsigned_full_range=unsigned_full_range,
)
if not prepack:
return qweight, scales
n, k = weights.shape
pack = 8 // int(bits)
if n % pack != 0:
raise ValueError(f"N ({n}) must be divisible by {pack} for CUDA QMoE prepacked weights.")
packed = _pack_weights_for_cuda_mixed_gemm(qweight.numpy(), n, k, int(bits), force_arch)
packed = np.asarray(packed).view(np.uint8).reshape(k, n // pack)
return torch.from_numpy(np.ascontiguousarray(packed)), scales
@staticmethod
def _matmulnbits_blockwise_quantize_impl(
weights: torch.Tensor,
bits: int,
block_size: int,
*,
symmetric: bool,
abs_scales: bool,
unsigned_full_range: bool,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Quantize ``weights`` with MatMulNBits pybinds and return unflattened storage."""
torch = _get_torch()
bits = int(bits)
block_size = int(block_size)
w = weights.detach().cpu().to(torch.float32).contiguous().numpy()
n, k = w.shape
if bits not in (4, 8):
raise ValueError(f"Blockwise quantization only supports 4 or 8 bits, got {bits}.")
if block_size <= 0:
raise ValueError(f"Blockwise quantization requires a positive block_size, got {block_size}.")
num_blocks = (k + block_size - 1) // block_size
pack = 8 // bits
blob_size = (block_size + pack - 1) // pack
if symmetric:
if bits == 4:
qmin, qmax, scale_divisor, zero_point = (-8, 7, 8, 8) if unsigned_full_range else (-7, 7, 7, 8)
else:
qmin, qmax, scale_divisor, zero_point = (
(-128, 127, 128, 128) if unsigned_full_range else (-127, 127, 127, 128)
)
padded_k = num_blocks * block_size
if padded_k != k:
w = np.pad(w, ((0, 0), (0, padded_k - k)), "constant")
blocked = w.reshape(n, num_blocks, block_size)
scales = np.max(np.abs(blocked), axis=2).astype(np.float32) / np.float32(scale_divisor)
scales = np.maximum(scales, np.finfo(np.float32).eps)
quantized = np.clip(np.rint(blocked / scales[:, :, np.newaxis]), qmin, qmax).astype(np.int16)
quantized = (quantized + zero_point).astype(np.uint8)
if bits == 4:
qweight = np.zeros((n, num_blocks, blob_size), dtype=np.uint8)
qweight[:, :, : quantized[:, :, 0::2].shape[2]] = quantized[:, :, 0::2] & 0xF
qweight[:, :, : quantized[:, :, 1::2].shape[2]] |= (quantized[:, :, 1::2] & 0xF) << 4
else:
qweight = quantized
zero_points = np.zeros((n, (num_blocks + 1) // 2 if bits == 4 else num_blocks), dtype=np.uint8)
return torch.from_numpy(qweight), torch.from_numpy(scales), torch.from_numpy(zero_points)
w_t = np.ascontiguousarray(w.T)
qweight = np.zeros((n, num_blocks, blob_size), dtype=np.uint8)
scales = np.zeros((n, num_blocks), dtype=np.float32)
zero_points = np.zeros((n, (num_blocks + 1) // 2 if bits == 4 else num_blocks), dtype=np.uint8)
quantize_matmul_4bits, quantize_matmul_8bits = _get_quantize_matmul_nbits()
quantize = quantize_matmul_4bits if bits == 4 else quantize_matmul_8bits
quantize(qweight, w_t, scales, zero_points, block_size, n, k, symmetric)
if abs_scales:
scales = np.abs(scales)
return torch.from_numpy(qweight), torch.from_numpy(scales), torch.from_numpy(zero_points)
@staticmethod
def matmulnbits_blockwise_quantize(
weights: torch.Tensor,
bits: int,
block_size: int,
*,
symmetric: bool = True,
return_zero_points: bool = False,
abs_scales: bool = True,
flatten_qweight: bool = True,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor] | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Quantize one expert with ONNX Runtime's MatMulNBits blockwise encoding.
``weights`` has logical shape ``[N, K]``. Returns raw flattened storage
``[N, ceil(K/block_size)*ceil(block_size/pack)]`` and block scales
``[N, ceil(K/block_size)]`` by default. Set ``flatten_qweight=False`` for
the MatMulNBits initializer shape
``[N, ceil(K/block_size), ceil(block_size/pack)]``.
Set ``return_zero_points=True`` to also return packed block zero-points.
Symmetric quantization uses the full ``[-8, 7]`` / ``[-128, 127]`` range
by default. Set ``unsigned_full_range=False`` to use the legacy
``[-7, 7]`` / ``[-127, 127]`` range.
"""
qweight, scales, zero_points = CudaQuantizer._matmulnbits_blockwise_quantize_impl(
weights,
bits,
block_size,
symmetric=symmetric,
abs_scales=abs_scales,
unsigned_full_range=unsigned_full_range,
)
if flatten_qweight:
qweight = qweight.reshape(qweight.shape[0], -1).contiguous()
if return_zero_points:
return qweight, scales, zero_points
return qweight, scales
@staticmethod
def matmulnbits_prepacked_blockwise_quantize(
weights: torch.Tensor,
bits: int,
block_size: int,
force_arch: int = 80,
*,
symmetric: bool = True,
return_zero_points: bool = False,
abs_scales: bool = True,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor] | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Quantize and CUDA-prepack one MatMulNBits weight initializer.
``weights`` has logical shape ``[N, K]``. Returns ``B`` with the standard
MatMulNBits initializer shape ``[N, K/block_size, block_size*bits/8]``
and scales with shape ``[N, K/block_size]``.
The ``force_arch`` value selects the mixed-GEMM weight layout and must match
the ``weight_prepacked`` attribute set on the MatMulNBits node:
* ``force_arch=80`` (default): SM80/Ampere layout, consumed by the SM80 kernel
(also used on newer GPUs via the compatibility path). Use ``weight_prepacked=1``.
* ``force_arch=90``: SM90/Hopper layout, consumed by the native SM90 TMA/WGMMA
kernel. Use ``weight_prepacked=2``. Requires ``block_size`` in {64, 128}.
"""
torch = _get_torch()
bits = int(bits)
block_size = int(block_size)
force_arch = int(force_arch)
if force_arch not in (80, 90):
raise ValueError(f"force_arch must be 80 (SM80) or 90 (SM90), but got {force_arch}.")
# The native SM90 kernel needs group_size to be a multiple of the 64-element Hopper K tile,
# so block_size=32 is only supported by the SM80/Ampere-class kernel.
allowed_block_sizes = (32, 64, 128) if force_arch == 80 else (64, 128)
if block_size not in allowed_block_sizes:
raise ValueError(
f"block_size must be one of {allowed_block_sizes} for force_arch={force_arch}, but got {block_size}."
)
n, k = weights.shape
if k % block_size != 0:
raise ValueError(f"K ({k}) must be divisible by block_size ({block_size}) for CUDA-prepacked weights.")
qweight, scales, zero_points = CudaQuantizer._matmulnbits_blockwise_quantize_impl(
weights,
bits,
block_size,
symmetric=symmetric,
abs_scales=abs_scales,
unsigned_full_range=unsigned_full_range,
)
pack_weights_for_cuda_mixed_gemm = _pack_weights_for_cuda_mixed_gemm
packed = pack_weights_for_cuda_mixed_gemm(qweight.reshape(n, -1).numpy(), n, k, bits, force_arch)
packed = np.asarray(packed).view(np.uint8).reshape(qweight.shape)
packed = torch.from_numpy(np.ascontiguousarray(packed))
if return_zero_points:
return packed, scales, zero_points
return packed, scales
@staticmethod
def qmoe_prepacked_blockwise_quantize(
weights: torch.Tensor,
bits: int,
block_size: int,
force_arch: int = 80,
*,
symmetric: bool = True,
return_zero_points: bool = False,
abs_scales: bool = False,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor] | tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Quantize one expert and CUTLASS-prepack it for CUDA QMoE fpA_intB GEMM.
``weights`` has logical shape ``[N, K]``. Returns ``qweight`` with shape
``[K, N/pack]`` and block scales with shape ``[N, K/block_size]``.
Set ``return_zero_points=True`` to also return packed block zero-points.
"""
bits = int(bits)
block_size = int(block_size)
n, k = weights.shape
pack = 8 // bits
if k % block_size != 0:
raise ValueError(f"K ({k}) must be divisible by block_size ({block_size}) for CUDA-prepacked weights.")
if n % pack != 0:
raise ValueError(f"N ({n}) must be divisible by {pack} for QMoE blockwise quantization.")
qweight, scales, zero_points = CudaQuantizer._matmulnbits_blockwise_quantize_impl(
weights,
bits,
block_size,
symmetric=symmetric,
abs_scales=abs_scales,
unsigned_full_range=unsigned_full_range,
)
pack_weights_for_cuda_mixed_gemm = _pack_weights_for_cuda_mixed_gemm
packed = pack_weights_for_cuda_mixed_gemm(qweight.reshape(n, -1).numpy(), n, k, bits, force_arch)
packed = np.asarray(packed).view(np.uint8).reshape(k, n // pack)
torch = _get_torch()
packed = torch.from_numpy(np.ascontiguousarray(packed))
if return_zero_points:
return packed, scales, zero_points
return packed, scales
@staticmethod
def symmetric_blockwise_quantize(
weights: torch.Tensor,
bits: int,
block_size: int,
*,
unsigned_full_range: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize one expert with a pure-PyTorch symmetric blockwise encoding.
This helper is useful for non-CUDA reference paths. Unlike the pybind-backed
helpers above, it pads the last dimension when it is not divisible by
``block_size`` and returns storage with the same leading shape as ``weights``.
"""
torch = _get_torch()
weights = weights.detach().cpu().contiguous()
original_shape = weights.shape
bits = int(bits)
block_size = int(block_size)
if bits == 4:
qmin, qmax, scale_divisor = (-8, 7, 8) if unsigned_full_range else (-7, 7, 7)
elif bits == 8:
qmin, qmax, scale_divisor = (-128, 127, 128) if unsigned_full_range else (-127, 127, 127)
else:
raise ValueError(f"CUDA blockwise quantization only supports 4 or 8 bits, got {bits}.")
last_dim = original_shape[-1]
num_blocks = (last_dim + block_size - 1) // block_size
pad_size = num_blocks * block_size - last_dim
if pad_size > 0:
pad_shape = list(original_shape)
pad_shape[-1] = pad_size
padding = torch.zeros(pad_shape, dtype=weights.dtype, device=weights.device)
weights_padded = torch.cat([weights, padding], dim=-1)
else:
weights_padded = weights
reshaped_weights = weights_padded.view(*original_shape[:-1], num_blocks, block_size)
block_max_abs = torch.max(torch.abs(reshaped_weights), dim=-1)[0]
scales = torch.clamp(block_max_abs / scale_divisor, min=1e-8)
quantized = torch.round(reshaped_weights / scales.unsqueeze(-1))
quantized = torch.clamp(quantized, qmin, qmax)
if bits == 4:
quantized_flat = quantized.to(torch.int8).view(*original_shape[:-1], num_blocks * block_size)
if pad_size > 0:
quantized_flat = quantized_flat[..., :-pad_size]
quantized_uint4 = (quantized_flat + 8).to(torch.uint8)
packed_shape = list(original_shape)
packed_shape[-1] = (original_shape[-1] + 1) // 2
qweight = torch.zeros(packed_shape, dtype=torch.uint8, device=weights.device)
qweight[..., :] = quantized_uint4[..., 0::2] & 0xF
if quantized_uint4.shape[-1] > 1:
qweight[..., : quantized_uint4[..., 1::2].shape[-1]] |= (quantized_uint4[..., 1::2] & 0xF) << 4
else:
qweight = quantized.to(torch.int8).view(*original_shape[:-1], num_blocks * block_size)
if pad_size > 0:
qweight = qweight[..., :-pad_size]
qweight = qweight.view(original_shape)
return qweight.cpu(), scales.cpu()