From 5f7d53dd0df9e519769f5256dba40203bba8b41f Mon Sep 17 00:00:00 2001 From: Xin He Date: Mon, 7 Sep 2026 23:58:39 -0700 Subject: [PATCH 1/2] refactor: reorganize tensor utility functions and update imports for model optimization Signed-off-by: Xin He --- auto_round/compressors/shard_writer.py | 2 +- auto_round/utils/__init__.py | 8 +- auto_round/utils/missing_tensors.py | 309 +---------- auto_round/utils/model_free_utils.py | 503 +++++++++++++++++- auto_round/utils/offload.py | 2 +- .../export/test_qlinear_pack_clamp.py | 2 +- 6 files changed, 494 insertions(+), 332 deletions(-) diff --git a/auto_round/compressors/shard_writer.py b/auto_round/compressors/shard_writer.py index 6ac0c8b06b..3aaf720581 100644 --- a/auto_round/compressors/shard_writer.py +++ b/auto_round/compressors/shard_writer.py @@ -238,7 +238,7 @@ def _expand_fused_experts(self, name: str, tensor: torch.Tensor) -> list[tuple[s List of (key, 2D tensor) pairs, or None if not a fused expert param. """ from auto_round.modeling.fused_moe.replace_modules import MOE_SKIP_PREFIXES - from auto_round.utils.missing_tensors import split_fused_expert_tensors + from auto_round.utils.model_free_utils import split_fused_expert_tensors parts = name.rsplit(".", 1) if len(parts) != 2: diff --git a/auto_round/utils/__init__.py b/auto_round/utils/__init__.py index 4e8b9c71d0..67987b294f 100644 --- a/auto_round/utils/__init__.py +++ b/auto_round/utils/__init__.py @@ -21,7 +21,13 @@ detect_weight_type, is_quantized_input_module, ) -from auto_round.utils.missing_tensors import copy_missing_tensors_from_source + + +def copy_missing_tensors_from_source(*args, **kwargs): + from auto_round.utils.missing_tensors import copy_missing_tensors_from_source as _copy_missing_tensors_from_source + + return _copy_missing_tensors_from_source(*args, **kwargs) + import transformers from packaging.version import Version diff --git a/auto_round/utils/missing_tensors.py b/auto_round/utils/missing_tensors.py index 8af39c9707..2188c60af5 100644 --- a/auto_round/utils/missing_tensors.py +++ b/auto_round/utils/missing_tensors.py @@ -39,153 +39,14 @@ import json import os -import re -from typing import Optional, Tuple import torch from auto_round.logger import logger from auto_round.utils.common import compress_layer_names +from auto_round.utils.model_free_utils import quantize_weight_rtn, split_fused_expert_tensors from auto_round.utils.weight_handler import _dequant_fp8_linear_weight -# ------------------------------------------------------------------ # -# Fused expert projection patterns # -# ------------------------------------------------------------------ # -# Maps a fused projection name to its constituent split names. -# When a 3D tensor ``*.experts.`` (shape [num_experts, ...]) -# is encountered, it is split along dimension 0 of the per-expert 2-D -# slice to produce one tensor per split name per expert. -_FUSED_EXPERT_PROJ_PATTERNS: dict[str, list[str]] = { - "gate_up_proj": ["gate_proj", "up_proj"], - "w13": ["w1", "w3"], -} - -_AUTOROUND_ISSUE_URL = "https://github.com/intel/auto-round/issues" -_WARNING_INDEX_PLACEHOLDER = "" - - -def _normalize_tensor_name_for_warning(name: str, numeric_replacement: str = _WARNING_INDEX_PLACEHOLDER) -> str: - """Normalize tensor names for warning_once deduplication. - - Replace standalone numeric path segments (e.g. ``layers.12.experts.3``) - and bracket indices (e.g. ``layers[12]``) with a fixed placeholder - (```` by default) so warning keys are stable across different - layer/expert ids. - """ - parts = name.split(".") - normalized_parts = [] - for part in parts: - if part.isdigit(): - normalized_parts.append(numeric_replacement) - continue - normalized_parts.append(re.sub(r"\[(\d+)\]", f"[{numeric_replacement}]", part)) - return ".".join(normalized_parts) - - -def split_fused_expert_tensors( - tensors_dict: dict[str, torch.Tensor], -) -> dict[str, torch.Tensor]: - """Split 3-D fused expert tensors into per-expert 2-D tensors. - - Many MoE checkpoints store expert weights as fused 3-D parameters - under either ``*.experts.`` or ``*.moe.`` - (for example ``gate_up_proj`` with shape - ``[num_experts, 2*intermediate, hidden]``, or ``down_proj`` with shape - ``[num_experts, out, in]``). After unfusing, each expert gets its own - 2-D weight tensor, e.g. - ``experts.0.gate_proj.weight [intermediate, hidden]``. - - Splitting rules: - - * **gate_up_proj** ``[N, 2*inter, hidden]`` → - ``experts.{i}.gate_proj.weight`` + ``experts.{i}.up_proj.weight`` - * **up_gate_proj** ``[N, 2*inter, hidden]`` → - ``experts.{i}.up_proj.weight`` + ``experts.{i}.gate_proj.weight`` - * **Other** stacked projections (e.g. ``down_proj``) ``[N, out, in]`` → - ``experts.{i}..weight`` - - For ``*.moe.`` tensors, outputs use - ``*.moe.experts.{i}..weight`` so they align with the sequential - expert module layout used by quantized MoE replacements. - - Non-3-D or non-expert tensors pass through unchanged. - - Args: - tensors_dict: Mapping of tensor names to tensors. - - Returns: - New dict with fused expert tensors replaced by per-expert 2-D tensors. - """ - result: dict[str, torch.Tensor] = {} - split_count = 0 - - for tensor_name, tensor in tensors_dict.items(): - if tensor.dim() != 3: - result[tensor_name] = tensor - continue - - warning_tensor_name = _normalize_tensor_name_for_warning(tensor_name) - - # Strip optional .weight suffix for pattern matching - stripped = tensor_name - stripped = stripped.removesuffix(".weight") # len(".weight") == 7 - - # Expect: .experts. - dot_idx = stripped.rfind(".") - if dot_idx < 0: - result[tensor_name] = tensor - continue - - parent = stripped[:dot_idx] - proj_name = stripped[dot_idx + 1 :] - - parent_last = parent.split(".")[-1] - is_experts_parent = parent.endswith("experts") or parent_last == "experts" - is_moe_parent = parent.endswith("moe") or parent_last == "moe" - - # The immediate parent must be "experts" or "moe" - if not is_experts_parent and not is_moe_parent: - logger.warning_once( - "Found 3-D tensor '%s' while splitting expert tensors; " - "it will be kept unchanged. If this is an MoE/expert weight that should be split/quantized, " - "please open an issue at %s.", - warning_tensor_name, - _AUTOROUND_ISSUE_URL, - ) - result[tensor_name] = tensor - continue - - target_prefix = parent if is_experts_parent else f"{parent}.experts" - - num_experts = tensor.shape[0] - - if proj_name in _FUSED_EXPERT_PROJ_PATTERNS: - split_names = _FUSED_EXPERT_PROJ_PATTERNS[proj_name] - logger.warning_once( - f"Splitting fused expert tensor '{warning_tensor_name}' " - f"(shape={list(tensor.shape)}, num_experts={num_experts}) " - f"into {split_names}" - ) - for i in range(num_experts): - expert_2d = tensor[i] # [fused_out, in_features] - chunks = expert_2d.chunk(len(split_names), dim=0) - for split_name, chunk in zip(split_names, chunks): - out_key = f"{target_prefix}.{i}.{split_name}.weight" - result[out_key] = chunk.contiguous() - else: - logger.warning_once( - f"Splitting stacked expert tensor '{warning_tensor_name}' " - f"(shape={list(tensor.shape)}, num_experts={num_experts})" - ) - for i in range(num_experts): - expert_2d = tensor[i] # [out, in] - out_key = f"{target_prefix}.{i}.{proj_name}.weight" - result[out_key] = expert_2d.contiguous() - - split_count += 1 - - return result - def copy_missing_tensors_from_source( source_dir: str, @@ -515,174 +376,6 @@ def _is_truly_missing(name: str) -> bool: ) -def quantize_weight_rtn( - weight: torch.Tensor, - bits: int, - group_size: int, - sym: bool = True, - device: Optional[torch.device] = None, - disable_opt_rtn: bool = True, - *, - packing: str, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Quantize a 2-D weight tensor and pack into auto_gptq format. - - Parameters - ---------- - weight : Tensor [out_features, in_features] - bits : target bit-width (e.g. 4, 8) - group_size : quantization group size along in_features - sym : use symmetric quantisation - device : compute device (cuda / cpu). Results are always returned on CPU. - disable_opt_rtn : when False and sym=True, use the optimised-RTN scale - search (``quant_tensor_opt_rtn_sym``) which evaluates three E8M0 - candidates per group and picks the best MSE. Defaults to True - (plain RTN) to preserve backward-compatible behaviour. - packing : the packing_format the artifact will declare (required); the - zero-point convention follows the inference backend registry. GPTQ_FORMAT - entries ("auto_round:auto_gptq", qlinear_torch_zp) store zp - 1 per - nibble and unpack with +1; GPTQ_FORMAT_NO_ZP entries ("auto_round", - "auto_round:gptqmodel", qlinear_torch) store the zero point directly. - Any other value raises. - - Returns - ------- - qweight : [in_features // pack_factor, out_features] int32 - qzeros : [num_groups, out_features // pack_factor] int32 - scales : [num_groups, out_features] float16 - """ - assert weight.dim() == 2, f"Expected 2-D weight, got {weight.dim()}-D" - from auto_round.inference.backend import GPTQ_FORMAT, GPTQ_FORMAT_NO_ZP - - if packing in GPTQ_FORMAT: - zp_minus_one = True - elif packing in GPTQ_FORMAT_NO_ZP: - zp_minus_one = False - else: - raise ValueError( - f"unsupported packing format '{packing}': zero-point conventions are defined only for " - f"{GPTQ_FORMAT} (zp - 1, unpacked with +1) and {GPTQ_FORMAT_NO_ZP} (zp stored directly)" - ) - out_features, in_features = weight.shape - if device is None: - device = weight.device - # Single-step transfer + cast avoids an intermediate BF16 copy on CUDA - # (``weight.to(device).float()`` would briefly allocate both BF16 and - # float32 buffers on the target device). - weight = weight.to(device=device, dtype=torch.float32) - - # --- pad in_features to multiple of group_size --- - if in_features % group_size != 0: - pad = group_size - (in_features % group_size) - weight = torch.nn.functional.pad(weight, (0, pad)) - in_features = weight.shape[1] - - num_groups = in_features // group_size - pack_factor = 32 // bits # values per int32 - - # --- pad out_features to multiple of pack_factor (needed for qzeros) --- - out_pad = 0 - if out_features % pack_factor != 0: - out_pad = pack_factor - (out_features % pack_factor) - weight = torch.nn.functional.pad(weight, (0, 0, 0, out_pad)) - padded_out = weight.shape[0] - - from auto_round.data_type.utils import get_quant_func, reshape_pad_tensor_by_group_size - - # Use get_quant_func (same as WrapperLinear) so all data types and opt_rtn - # variants are resolved via the QUANT_FUNC_WITH_DTYPE registry uniformly. - quant_func, _ = get_quant_func("int", bits, sym=sym, disable_opt_rtn=disable_opt_rtn, iters=0) - # quant_func returns (qdq_result, scale, zp_or_maxq) - _, scale, zp_val = quant_func(weight, bits=bits, group_size=group_size) - - if sym: - maxq = 1 << (bits - 1) # e.g. 8 for 4-bit - zero_point = maxq # unsigned offset for packing - - # scale shape: [padded_out * num_groups, 1] - # Reshape weight for group-wise quantization: [padded_out * num_groups, group_size] - w_grouped, _, _ = reshape_pad_tensor_by_group_size(weight, group_size) - w_grouped = w_grouped.to(device=device, dtype=torch.float32) - del weight - - # Compute integer values for packing - q = (w_grouped / scale).round_().clamp_(-maxq, maxq - 1) - del w_grouped - q += zero_point # shift to unsigned [0, 2*maxq - 1] - q = q.to(torch.int32) - - # scale → [num_groups, padded_out] (float16) - scales_out = scale.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.float16) - del scale - - zp = torch.full((num_groups, padded_out), zero_point, dtype=torch.int32, device=device) - else: - # Asymmetric quantization - max_int = (1 << bits) - 1 - - # scale shape: [padded_out * num_groups, 1], zp_val shape: [padded_out * num_groups, 1] - - # Reshape weight for group-wise quantization - w_grouped, _, _ = reshape_pad_tensor_by_group_size(weight, group_size) - w_grouped = w_grouped.to(device=device, dtype=torch.float32) - del weight - - # Compute integer values for packing - q = (w_grouped / scale).round_() - del w_grouped - q += zp_val - q.clamp_(0, max_int) - q = q.to(torch.int32) - - # scale → [num_groups, padded_out] (float16) - scales_out = scale.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.float16) - del scale - - # zp → [num_groups, padded_out] - zp = zp_val.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.int32) - del zp_val - - # q → [in_features, padded_out] - q = q.reshape(padded_out, in_features).t().contiguous() - - # ---- Pack qweight: [in_features // pack_factor, padded_out] ---- - # Vectorised: reshape → broadcast shift → int64 sum (≡ bitwise-OR for - # non-overlapping bit lanes) avoids a Python loop per bit-lane. - _shifts = torch.arange(pack_factor, dtype=torch.int64, device=device) * bits - q_packed = q.reshape(in_features // pack_factor, pack_factor, padded_out).to(torch.int64) - del q - qweight = (q_packed << _shifts[None, :, None]).sum(dim=1).to(torch.int32) - del q_packed - - # ---- Pack qzeros: [num_groups, padded_out // pack_factor] ---- - # The zero-point convention is a property of the declared packing format - # (see the backend registry): the auto_round:auto_gptq format - # (qlinear_torch_zp) adds +1 to zeros after unpacking, so we subtract 1 - # before packing to compensate. The (zp-1) packing cannot represent - # zp=0: a -1 left-shifts into a negative word and corrupts every nibble - # sharing it. Clamp the stored value so only that zero-point degrades - # (decodes as +1). Formats in GPTQ_FORMAT_NO_ZP (plain qlinear_torch) - # read the nibble as the zero point directly, so it is stored as-is - # (symmetric packing stores the offset-binary constant the same way: - # 7 for the gptq family, 8 for the direct family). - if zp_minus_one: - zp = (zp - 1).clamp_(0, (1 << bits) - 1) - else: - zp = zp.clamp_(0, (1 << bits) - 1) - zp_packed = zp.reshape(num_groups, padded_out // pack_factor, pack_factor).to(torch.int64) - del zp - qzeros = (zp_packed << _shifts[None, None, :]).sum(dim=2).to(torch.int32) - del zp_packed, _shifts - - # Remove output padding from qweight / scales (qzeros stays in pack units) - if out_pad > 0: - qweight = qweight[:, :out_features] - scales_out = scales_out[:, :out_features] - - # Always return CPU tensors (safetensors requires CPU) - return qweight.cpu(), qzeros.cpu(), scales_out.cpu() - - def _woq_quantize_missing_tensors(target_dir: str, missing_tensors_dict: dict) -> dict: """Apply WOQ (Weight-Only Quantization) to missing Linear weight tensors. diff --git a/auto_round/utils/model_free_utils.py b/auto_round/utils/model_free_utils.py index b9e7d9811b..a18b8f0f85 100644 --- a/auto_round/utils/model_free_utils.py +++ b/auto_round/utils/model_free_utils.py @@ -37,7 +37,6 @@ from auto_round.schemes import PRESET_SCHEMES, QuantizationScheme, preset_name_to_scheme from auto_round.utils.common import to_standard_regex from auto_round.utils.device import clear_memory, compile_func -from auto_round.utils.missing_tensors import quantize_weight_rtn, split_fused_expert_tensors _NVFP4_E5M3_DATA_TYPE = "nvfp4_v2" _BLOCK_NAME_TO_IGNORE = ("shared_expert_gate.", ".gate.", "embed", "conv") @@ -63,6 +62,381 @@ ) +# --------------------------------------------------------------------------- +# Fused Expert Tensor Splitting (shared with auto_round.utils.missing_tensors) +# --------------------------------------------------------------------------- +# Maps a fused projection name to its constituent split names. +# When a 3D tensor ``*.experts.`` (shape [num_experts, ...]) +# is encountered, it is split along dimension 0 of the per-expert 2-D +# slice to produce one tensor per split name per expert. +_FUSED_EXPERT_PROJ_PATTERNS: dict[str, list[str]] = { + "gate_up_proj": ["gate_proj", "up_proj"], + "w13": ["w1", "w3"], +} + +# Native vLLM/compressed-tensors FusedMoE parameter naming encodes the +# projection *and* its companion scale directly in the tensor name, with no +# dot separator (e.g. ``experts.w13_weight``, ``experts.w2_weight_scale``, +# ``experts.w2_weight_scale_2``, ``experts.w13_input_scale``), unlike the +# dotted ``experts..weight`` convention used elsewhere. Order matters +# only for readability here: none of these suffixes is itself a suffix of +# another, so matching is unambiguous. +_MOE_NATIVE_SUFFIXES: list[tuple[str, str]] = [ + ("_weight_scale_2", "weight_scale_2"), + ("_weight_scale", "weight_scale"), + ("_weight", "weight"), + ("_input_scale", "input_scale"), +] + +_AUTOROUND_ISSUE_URL = "https://github.com/intel/auto-round/issues" +_WARNING_INDEX_PLACEHOLDER = "" + + +def _normalize_tensor_name_for_warning(name: str, numeric_replacement: str = _WARNING_INDEX_PLACEHOLDER) -> str: + """Normalize tensor names for warning_once deduplication. + + Replace standalone numeric path segments (e.g. ``layers.12.experts.3``) + and bracket indices (e.g. ``layers[12]``) with a fixed placeholder + (```` by default) so warning keys are stable across different + layer/expert ids. + """ + parts = name.split(".") + normalized_parts = [] + for part in parts: + if part.isdigit(): + normalized_parts.append(numeric_replacement) + continue + normalized_parts.append(re.sub(r"\[(\d+)\]", f"[{numeric_replacement}]", part)) + return ".".join(normalized_parts) + + +def _parse_fused_proj_token(token: str) -> tuple[str, str]: + """Split a raw expert-projection token into ``(base_proj, out_suffix)``. + + Handles both the plain dotted convention (token is just the projection + name, e.g. ``"down_proj"`` / ``"gate_up_proj"`` -> ``out_suffix="weight"``) + and the native FusedMoE convention where the companion tensor kind is + encoded in the token itself (e.g. ``"w2_weight_scale"`` -> + ``base_proj="w2"``, ``out_suffix="weight_scale"``). + """ + for raw_suffix, out_suffix in _MOE_NATIVE_SUFFIXES: + if token.endswith(raw_suffix) and len(token) > len(raw_suffix): + return token[: -len(raw_suffix)], out_suffix + return token, "weight" + + +def split_fused_expert_tensors( + tensors_dict: dict[str, torch.Tensor], +) -> dict[str, torch.Tensor]: + """Split 3-D fused expert tensors into per-expert 2-D tensors. + + Many MoE checkpoints store expert weights as fused 3-D parameters + under either ``*.experts.`` or ``*.moe.`` + (for example ``gate_up_proj`` with shape + ``[num_experts, 2*intermediate, hidden]``, or ``down_proj`` with shape + ``[num_experts, out, in]``). After unfusing, each expert gets its own + 2-D weight tensor, e.g. + ``experts.0.gate_proj.weight [intermediate, hidden]``. + + Splitting rules: + + * **gate_up_proj** ``[N, 2*inter, hidden]`` → + ``experts.{i}.gate_proj.weight`` + ``experts.{i}.up_proj.weight`` + * **up_gate_proj** ``[N, 2*inter, hidden]`` → + ``experts.{i}.up_proj.weight`` + ``experts.{i}.gate_proj.weight`` + * **Other** stacked projections (e.g. ``down_proj``) ``[N, out, in]`` → + ``experts.{i}..weight`` + + For ``*.moe.`` tensors, outputs use + ``*.moe.experts.{i}..weight`` so they align with the sequential + expert module layout used by quantized MoE replacements. + + Native vLLM/compressed-tensors FusedMoE naming (e.g. ``experts.w13_weight``, + ``experts.w2_weight_scale``, ``experts.w13_weight_scale_2``, + ``experts.w2_input_scale``) is also recognized: the companion tensor kind + is decoded from the token itself (see :func:`_parse_fused_proj_token`) so + weight/weight_scale/weight_scale_2/input_scale tensors for the same + projection are split consistently and end up under the same + ``experts.{i}.`` prefix, instead of scale tensors being mistaken for + independent weights. + + Non-3-D tensors are only split when recognized as one of these native + scale companions; other non-3-D or non-expert tensors pass through + unchanged. + + Args: + tensors_dict: Mapping of tensor names to tensors. + + Returns: + New dict with fused expert tensors replaced by per-expert 2-D tensors. + """ + result: dict[str, torch.Tensor] = {} + split_count = 0 + + for tensor_name, tensor in tensors_dict.items(): + if tensor.dim() not in (1, 2, 3): + result[tensor_name] = tensor + continue + + warning_tensor_name = _normalize_tensor_name_for_warning(tensor_name) + + # Strip optional .weight suffix (dotted convention) for pattern matching + stripped = tensor_name.removesuffix(".weight") # len(".weight") == 7 + + # Expect: .experts. + dot_idx = stripped.rfind(".") + if dot_idx < 0: + result[tensor_name] = tensor + continue + + parent = stripped[:dot_idx] + raw_token = stripped[dot_idx + 1 :] + proj_name, out_suffix = _parse_fused_proj_token(raw_token) + + parent_last = parent.split(".")[-1] + is_experts_parent = parent.endswith("experts") or parent_last == "experts" + is_moe_parent = parent.endswith("moe") or parent_last == "moe" + + # The immediate parent must be "experts" or "moe" + if not is_experts_parent and not is_moe_parent: + if tensor.dim() == 3: + logger.warning_once( + "Found 3-D tensor '%s' while splitting expert tensors; " + "it will be kept unchanged. If this is an MoE/expert weight that should be split/quantized, " + "please open an issue at %s.", + warning_tensor_name, + _AUTOROUND_ISSUE_URL, + ) + result[tensor_name] = tensor + continue + + # 1-D/2-D tensors under an "experts"/"moe" parent are only touched here + # when they are a recognized native FusedMoE scale companion (e.g. + # ``w13_weight_scale_2`` [num_experts, 2] or ``w2_weight_scale_2`` + # [num_experts]). Anything else (e.g. a plain 1-D bias/norm) passes + # through unchanged. + if tensor.dim() != 3 and out_suffix == "weight": + result[tensor_name] = tensor + continue + + target_prefix = parent if is_experts_parent else f"{parent}.experts" + num_experts = tensor.shape[0] + split_names = _FUSED_EXPERT_PROJ_PATTERNS.get(proj_name, [proj_name]) + + if tensor.dim() == 3: + if len(split_names) > 1: + logger.warning_once( + f"Splitting fused expert tensor '{warning_tensor_name}' " + f"(shape={list(tensor.shape)}, num_experts={num_experts}) " + f"into {split_names}" + ) + else: + logger.warning_once( + f"Splitting stacked expert tensor '{warning_tensor_name}' " + f"(shape={list(tensor.shape)}, num_experts={num_experts})" + ) + for i in range(num_experts): + expert_2d = tensor[i] # [fused_out, in_features] or [out, in] + chunks = expert_2d.chunk(len(split_names), dim=0) + for split_name, chunk in zip(split_names, chunks): + out_key = f"{target_prefix}.{i}.{split_name}.{out_suffix}" + result[out_key] = chunk.contiguous() + else: + # Native per-expert scale companion (e.g. weight_scale_2 / + # input_scale), shape [num_experts] or [num_experts, len(split_names)]. + logger.warning_once( + f"Splitting fused expert scale tensor '{warning_tensor_name}' " + f"(shape={list(tensor.shape)}, num_experts={num_experts}) " + f"into {split_names}" + ) + flat = tensor.reshape(num_experts, -1) + n_splits = len(split_names) + for i in range(num_experts): + row = flat[i] + if row.numel() == n_splits: + per_split_vals = [row[j].reshape(1) for j in range(n_splits)] + else: + # Single value shared across all splits (e.g. a scalar + # global scale not broken out per gate/up projection). + shared = row.reshape(-1)[:1] + per_split_vals = [shared for _ in range(n_splits)] + for split_name, val in zip(split_names, per_split_vals): + out_key = f"{target_prefix}.{i}.{split_name}.{out_suffix}" + result[out_key] = val.contiguous() + + split_count += 1 + + return result + + +def quantize_weight_rtn( + weight: torch.Tensor, + bits: int, + group_size: int, + sym: bool = True, + device: Optional[torch.device] = None, + disable_opt_rtn: bool = True, + *, + packing: str, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Quantize a 2-D weight tensor and pack into auto_gptq format. + + Parameters + ---------- + weight : Tensor [out_features, in_features] + bits : target bit-width (e.g. 4, 8) + group_size : quantization group size along in_features + sym : use symmetric quantisation + device : compute device (cuda / cpu). Results are always returned on CPU. + disable_opt_rtn : when False and sym=True, use the optimised-RTN scale + search (``quant_tensor_opt_rtn_sym``) which evaluates three E8M0 + candidates per group and picks the best MSE. Defaults to True + (plain RTN) to preserve backward-compatible behaviour. + packing : the packing_format the artifact will declare (required); the + zero-point convention follows the inference backend registry. GPTQ_FORMAT + entries ("auto_round:auto_gptq", qlinear_torch_zp) store zp - 1 per + nibble and unpack with +1; GPTQ_FORMAT_NO_ZP entries ("auto_round", + "auto_round:gptqmodel", qlinear_torch) store the zero point directly. + Any other value raises. + + Returns + ------- + qweight : [in_features // pack_factor, out_features] int32 + qzeros : [num_groups, out_features // pack_factor] int32 + scales : [num_groups, out_features] float16 + """ + assert weight.dim() == 2, f"Expected 2-D weight, got {weight.dim()}-D" + from auto_round.inference.backend import GPTQ_FORMAT, GPTQ_FORMAT_NO_ZP + + if packing in GPTQ_FORMAT: + zp_minus_one = True + elif packing in GPTQ_FORMAT_NO_ZP: + zp_minus_one = False + else: + raise ValueError( + f"unsupported packing format '{packing}': zero-point conventions are defined only for " + f"{GPTQ_FORMAT} (zp - 1, unpacked with +1) and {GPTQ_FORMAT_NO_ZP} (zp stored directly)" + ) + out_features, in_features = weight.shape + if device is None: + device = weight.device + # Single-step transfer + cast avoids an intermediate BF16 copy on CUDA + # (``weight.to(device).float()`` would briefly allocate both BF16 and + # float32 buffers on the target device). + weight = weight.to(device=device, dtype=torch.float32) + + # --- pad in_features to multiple of group_size --- + if in_features % group_size != 0: + pad = group_size - (in_features % group_size) + weight = torch.nn.functional.pad(weight, (0, pad)) + in_features = weight.shape[1] + + num_groups = in_features // group_size + pack_factor = 32 // bits # values per int32 + + # --- pad out_features to multiple of pack_factor (needed for qzeros) --- + out_pad = 0 + if out_features % pack_factor != 0: + out_pad = pack_factor - (out_features % pack_factor) + weight = torch.nn.functional.pad(weight, (0, 0, 0, out_pad)) + padded_out = weight.shape[0] + + from auto_round.data_type.utils import get_quant_func, reshape_pad_tensor_by_group_size + + # Use get_quant_func (same as WrapperLinear) so all data types and opt_rtn + # variants are resolved via the QUANT_FUNC_WITH_DTYPE registry uniformly. + quant_func, _ = get_quant_func("int", bits, sym=sym, disable_opt_rtn=disable_opt_rtn, iters=0) + # quant_func returns (qdq_result, scale, zp_or_maxq) + _, scale, zp_val = quant_func(weight, bits=bits, group_size=group_size) + + if sym: + maxq = 1 << (bits - 1) # e.g. 8 for 4-bit + zero_point = maxq # unsigned offset for packing + + # scale shape: [padded_out * num_groups, 1] + # Reshape weight for group-wise quantization: [padded_out * num_groups, group_size] + w_grouped, _, _ = reshape_pad_tensor_by_group_size(weight, group_size) + w_grouped = w_grouped.to(device=device, dtype=torch.float32) + del weight + + # Compute integer values for packing + q = (w_grouped / scale).round_().clamp_(-maxq, maxq - 1) + del w_grouped + q += zero_point # shift to unsigned [0, 2*maxq - 1] + q = q.to(torch.int32) + + # scale → [num_groups, padded_out] (float16) + scales_out = scale.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.float16) + del scale + + zp = torch.full((num_groups, padded_out), zero_point, dtype=torch.int32, device=device) + else: + # Asymmetric quantization + max_int = (1 << bits) - 1 + + # scale shape: [padded_out * num_groups, 1], zp_val shape: [padded_out * num_groups, 1] + + # Reshape weight for group-wise quantization + w_grouped, _, _ = reshape_pad_tensor_by_group_size(weight, group_size) + w_grouped = w_grouped.to(device=device, dtype=torch.float32) + del weight + + # Compute integer values for packing + q = (w_grouped / scale).round_() + del w_grouped + q += zp_val + q.clamp_(0, max_int) + q = q.to(torch.int32) + + # scale → [num_groups, padded_out] (float16) + scales_out = scale.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.float16) + del scale + + # zp → [num_groups, padded_out] + zp = zp_val.squeeze(-1).reshape(padded_out, num_groups).t().contiguous().to(torch.int32) + del zp_val + + # q → [in_features, padded_out] + q = q.reshape(padded_out, in_features).t().contiguous() + + # ---- Pack qweight: [in_features // pack_factor, padded_out] ---- + # Vectorised: reshape → broadcast shift → int64 sum (≡ bitwise-OR for + # non-overlapping bit lanes) avoids a Python loop per bit-lane. + _shifts = torch.arange(pack_factor, dtype=torch.int64, device=device) * bits + q_packed = q.reshape(in_features // pack_factor, pack_factor, padded_out).to(torch.int64) + del q + qweight = (q_packed << _shifts[None, :, None]).sum(dim=1).to(torch.int32) + del q_packed + + # ---- Pack qzeros: [num_groups, padded_out // pack_factor] ---- + # The zero-point convention is a property of the declared packing format + # (see the backend registry): the auto_round:auto_gptq format + # (qlinear_torch_zp) adds +1 to zeros after unpacking, so we subtract 1 + # before packing to compensate. The (zp-1) packing cannot represent + # zp=0: a -1 left-shifts into a negative word and corrupts every nibble + # sharing it. Clamp the stored value so only that zero-point degrades + # (decodes as +1). Formats in GPTQ_FORMAT_NO_ZP (plain qlinear_torch) + # read the nibble as the zero point directly, so it is stored as-is + # (symmetric packing stores the offset-binary constant the same way: + # 7 for the gptq family, 8 for the direct family). + if zp_minus_one: + zp = (zp - 1).clamp_(0, (1 << bits) - 1) + else: + zp = zp.clamp_(0, (1 << bits) - 1) + zp_packed = zp.reshape(num_groups, padded_out // pack_factor, pack_factor).to(torch.int64) + del zp + qzeros = (zp_packed << _shifts[None, None, :]).sum(dim=2).to(torch.int32) + del zp_packed, _shifts + + # Remove output padding from qweight / scales (qzeros stays in pack units) + if out_pad > 0: + qweight = qweight[:, :out_features] + scales_out = scales_out[:, :out_features] + + # Always return CPU tensors (safetensors requires CPU) + return qweight.cpu(), qzeros.cpu(), scales_out.cpu() + + # --------------------------------------------------------------------------- # Generic Utility Helpers # --------------------------------------------------------------------------- @@ -258,10 +632,22 @@ def _reciprocal_global_scale(scale: torch.Tensor) -> torch.Tensor: def _handle_nvfp4_source_tensors( raw_tensors: dict[str, torch.Tensor], matcher: Any, + device: str = "cpu", + shard_name: str | None = None, ) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor], list[str]]: - """Passthrough NVFP4 source tensors when target scheme for the layer is NVFP4.""" + """Passthrough NVFP4 source tensors when target scheme for the layer is NVFP4. + + Layers whose target scheme is *not* NVFP4 (e.g. converting an NVFP4 donor + checkpoint to MXFP4) are dequantized to bfloat16 via + :func:`_dequant_nvfp4_tensors` instead of being left as raw + ``weight_packed``/``weight_scale`` tensors — otherwise those companion + scale tensors would fall through to the generic RTN/MXFP quantization path + and be mistaken for ordinary weights, corrupting both the weight and its + (unrelated) scale. + """ passthrough_tensors: dict[str, torch.Tensor] = {} passthrough_layers: list[str] = [] + dequant_layers: list[str] = [] for name, tensor in list(raw_tensors.items()): if not name.endswith(".weight_packed") or tensor.dtype not in (torch.uint8, torch.int8): @@ -277,30 +663,104 @@ def _handle_nvfp4_source_tensors( continue scheme_bits = scheme.get("bits") scheme_data_type = (scheme.get("data_type") or "").lower() - if not (scheme_bits == 4 and (is_nv_fp(scheme_data_type) or scheme_data_type == _NVFP4_E5M3_DATA_TYPE)): - continue - - keys_to_move = [name, scale_key] - weight_global_scale_key = f"{layer_name}.weight_global_scale" - input_global_scale_key = f"{layer_name}.input_global_scale" - if weight_global_scale_key in raw_tensors: - keys_to_move.append(weight_global_scale_key) - if input_global_scale_key in raw_tensors: - keys_to_move.append(input_global_scale_key) - - for key in keys_to_move: - value = raw_tensors.pop(key).to("cpu") - if key.endswith("_global_scale"): - value = value.reshape([1]) - passthrough_tensors[key] = value - passthrough_layers.append(layer_name) + if scheme_bits == 4 and (is_nv_fp(scheme_data_type) or scheme_data_type == _NVFP4_E5M3_DATA_TYPE): + keys_to_move = [name, scale_key] + weight_global_scale_key = f"{layer_name}.weight_global_scale" + input_global_scale_key = f"{layer_name}.input_global_scale" + if weight_global_scale_key in raw_tensors: + keys_to_move.append(weight_global_scale_key) + if input_global_scale_key in raw_tensors: + keys_to_move.append(input_global_scale_key) + + for key in keys_to_move: + value = raw_tensors.pop(key).to("cpu") + if key.endswith("_global_scale"): + value = value.reshape([1]) + passthrough_tensors[key] = value + passthrough_layers.append(layer_name) + else: + dequant_layers.append(layer_name) if passthrough_layers: logger.info(f"Handling NVFP4 source tensor(s): {len(passthrough_layers)} passthrough layer(s).") + if dequant_layers: + raw_tensors = _dequant_nvfp4_tensors(raw_tensors, dequant_layers, device=device, shard_name=shard_name) + return raw_tensors, passthrough_tensors, passthrough_layers +def _dequant_nvfp4_tensors( + raw_tensors: dict[str, torch.Tensor], + layer_names: list[str], + device: str = "cpu", + shard_name: str | None = None, +) -> dict[str, torch.Tensor]: + """Dequantize NVFP4-packed weight tensors to bfloat16. + + Reverses the NVFP4 packing (see :func:`auto_round.data_type.nvfp.nv_fp4`): + ``real_weight = unpack(weight_packed) * weight_scale.float() * (1 / weight_global_scale)`` + where ``weight_scale`` is the per-block (group_size=16) FP8 E4M3 scale and + ``weight_global_scale`` is the per-tensor FP32 global scale (already + reciprocal-normalized by :func:`_normalize_nvfp4_source_tensors`). + + The dequantized weight is written back under ``.weight`` (removing + ``weight_packed``/``weight_scale``/``weight_global_scale``/ + ``input_global_scale``) so the downstream RTN/MXFP path can requantize it + to the requested target scheme. + """ + from auto_round.data_type.nvfp import get_reciprocal + from auto_round.experimental.qmodules.fp4_utils import unpack_fp4_from_uint8 + + group_size = 16 + dequant_device = str(device or "cpu") + shard_prefix = f"[{shard_name}] " if shard_name else "" + logger.info(f"{shard_prefix}Dequantizing {len(layer_names)} NVFP4 tensor(s) to bfloat16 on {dequant_device}.") + + for layer_name in layer_names: + packed_key = f"{layer_name}.weight_packed" + scale_key = f"{layer_name}.weight_scale" + if packed_key not in raw_tensors or scale_key not in raw_tensors: + continue + packed = raw_tensors.pop(packed_key) + block_scale = raw_tensors.pop(scale_key) + global_scale = raw_tensors.pop(f"{layer_name}.weight_global_scale", None) + # Input scale has no meaning once the weight is dequantized back to a + # plain high-precision tensor; drop it. + raw_tensors.pop(f"{layer_name}.input_global_scale", None) + + def _do_dequant( + packed: torch.Tensor = packed, + block_scale: torch.Tensor = block_scale, + global_scale: torch.Tensor | None = global_scale, + ) -> torch.Tensor: + out_features, half_in_features = packed.shape + in_features = half_in_features * 2 + unpacked = unpack_fp4_from_uint8(packed, out_features, in_features, dtype=torch.float32) + n_groups = in_features // group_size + scale_f = block_scale.to(torch.float32).reshape(out_features, n_groups, 1) + dq = (unpacked.reshape(out_features, n_groups, group_size) * scale_f).reshape(out_features, in_features) + if global_scale is not None: + dq = dq * get_reciprocal(global_scale.to(torch.float32).reshape(())) + return dq.to(torch.bfloat16) + + dq_weight = _dequantize_with_device_fallback( + dequant_device=dequant_device, + shard_prefix=shard_prefix, + op_name="NVFP4 dequant", + tensor_label=layer_name, + on_device=lambda: _do_dequant( + packed=packed.to(dequant_device, non_blocking=True), + block_scale=block_scale.to(dequant_device, non_blocking=True), + global_scale=(global_scale.to(dequant_device, non_blocking=True) if global_scale is not None else None), + ).to("cpu"), + on_cpu=_do_dequant, + ) + raw_tensors[f"{layer_name}.weight"] = dq_weight + + return raw_tensors + + def _is_out_of_memory_error(exc: Exception) -> bool: if isinstance(exc, torch.OutOfMemoryError): return True @@ -1142,10 +1602,13 @@ def _process_shard( output_tensors.update(passthrough_tensors) quantized_layers.extend(passthrough_layers) - # 3) NVFP4 passthrough for layers already stored in packed format. + # 3) NVFP4 passthrough for layers already stored in packed format (or + # dequant to bfloat16 when the target scheme requires re-quantization). raw_tensors, nvfp_passthrough_tensors, nvfp_passthrough_layers = _handle_nvfp4_source_tensors( raw_tensors, matcher, + device=device, + shard_name=shard_name, ) output_tensors.update(nvfp_passthrough_tensors) quantized_layers.extend(nvfp_passthrough_layers) diff --git a/auto_round/utils/offload.py b/auto_round/utils/offload.py index 0cea78e8cd..c705f43bca 100644 --- a/auto_round/utils/offload.py +++ b/auto_round/utils/offload.py @@ -101,7 +101,7 @@ def _maybe_split_fused_expert_keys(state_dict: dict, module: torch.nn.Module) -> continue # original fused module still in the tree; assign as-is to_split[key] = state_dict.pop(key) if to_split: - from auto_round.utils.missing_tensors import split_fused_expert_tensors + from auto_round.utils.model_free_utils import split_fused_expert_tensors state_dict.update(split_fused_expert_tensors(to_split)) return state_dict diff --git a/test/unit/test_cpu/export/test_qlinear_pack_clamp.py b/test/unit/test_cpu/export/test_qlinear_pack_clamp.py index 29aeda6e69..8d9e0819e8 100644 --- a/test/unit/test_cpu/export/test_qlinear_pack_clamp.py +++ b/test/unit/test_cpu/export/test_qlinear_pack_clamp.py @@ -194,7 +194,7 @@ def test_missing_tensors_gptq_pack_clamps_zero_zero_point(): """quantize_weight_rtn packs zp with the same (zp-1) convention; an all-positive weight yields zp=0 everywhere, which must clamp instead of smearing -1 across the packed zero-point words.""" - from auto_round.utils.missing_tensors import quantize_weight_rtn + from auto_round.utils.model_free_utils import quantize_weight_rtn torch.manual_seed(0) bits, group = 4, 32 From 41f5fd5e3c8b9197b9e68ce5b0301bcfab531490 Mon Sep 17 00:00:00 2001 From: Xin He Date: Wed, 9 Sep 2026 11:09:48 +0800 Subject: [PATCH 2/2] fix: remove redundant up_gate_proj documentation and ensure tensor contiguity in dequantization Signed-off-by: Xin He --- auto_round/utils/model_free_utils.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/auto_round/utils/model_free_utils.py b/auto_round/utils/model_free_utils.py index 66190e9f8c..6cf82a2da2 100644 --- a/auto_round/utils/model_free_utils.py +++ b/auto_round/utils/model_free_utils.py @@ -150,8 +150,6 @@ def split_fused_expert_tensors( * **gate_up_proj** ``[N, 2*inter, hidden]`` → ``experts.{i}.gate_proj.weight`` + ``experts.{i}.up_proj.weight`` - * **up_gate_proj** ``[N, 2*inter, hidden]`` → - ``experts.{i}.up_proj.weight`` + ``experts.{i}.gate_proj.weight`` * **Other** stacked projections (e.g. ``down_proj``) ``[N, out, in]`` → ``experts.{i}..weight`` @@ -891,6 +889,7 @@ def _dequant_nvfp4_tensors( if packed_key not in raw_tensors or scale_key not in raw_tensors: continue packed = raw_tensors.pop(packed_key) + packed = packed.contiguous().view(torch.uint8) block_scale = raw_tensors.pop(scale_key) global_scale = raw_tensors.pop(f"{layer_name}.weight_global_scale", None) # Input scale has no meaning once the weight is dequantized back to a