# -*- coding: utf-8 -*- """Nunchaku parameter packer for QwenImage transformer blocks. Encapsulates NVFP4 (SVDQW4A4Linear) and AWQ INT4 (AWQW4A16Linear) quantization and memory-layout transformation for Nunchaku inference. """ import sys import os import torch # Ensure packages/deepcompressor is available PACKAGE_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "packages", "deepcompressor")) if PACKAGE_ROOT not in sys.path: sys.path.insert(0, PACKAGE_ROOT) from deepcompressor.data.dtype import QDType from deepcompressor.quantizer.config.base import QuantizerConfig from deepcompressor.quantizer.processor import Quantizer from deepcompressor.backend.nunchaku.convert import ( convert_to_nunchaku_w4x4y16_linear_state_dict, convert_to_nunchaku_w4x16_adanorm_zero_state_dict, ) def quantize_and_pack_nvfp4_linear( weight: torch.Tensor, bias: torch.Tensor | None = None, smooth: torch.Tensor | None = None, lora: tuple[torch.Tensor, torch.Tensor] | None = None, per_channel: bool = False, device: str = "cuda", ) -> dict[str, torch.Tensor]: """Quantize a linear layer to NVFP4 and pack into Nunchaku SVDQ format. Args: weight: Original weight tensor [out_features, in_features] in BF16/FP16. bias: Optional bias tensor [out_features]. smooth: Optional channel smoothing factors [in_features]. lora: Optional tuple of (lora_down, lora_up) where: lora_down has shape [rank, in_features] and lora_up has shape [out_features, rank]. per_channel: If True, uses per-channel scales (wcscales); otherwise per-tensor (wtscale). device: Device to perform quantization on. Returns: dict containing: - qweight: [out_features, in_features // 2] int8 - wscales: [in_features // 16, out_features] float8_e4m3fn - wcscales (if per_channel) or wtscale (if not per_channel) - bias: [out_features] - smooth_factor: [in_features] - smooth_factor_orig: [in_features] - proj_down: [in_features, rank] - proj_up: [out_features, rank] """ weight = weight.detach().clone().to(device=device) if bias is not None: bias = bias.detach().clone().to(device=device) else: bias = torch.zeros(weight.shape[0], dtype=weight.dtype, device=device) if smooth is not None: smooth = smooth.detach().clone().to(device=device) else: smooth = torch.ones(weight.shape[1], dtype=weight.dtype, device=device) if lora is not None: lora_down, lora_up = lora lora = (lora_down.detach().clone().to(device=device), lora_up.detach().clone().to(device=device)) # Configure NVFP4 quantizer # Group shapes: 1st level is per-channel ([1, -1]) or per-tensor ([-1, -1]), 2nd level is group 16 first_level = [1, -1] if per_channel else [-1, -1] cfg = QuantizerConfig( dtype=QDType.sfp4_e2m1_all, group_shapes=[first_level, [1, 16, 1, 1, 1]], scale_dtypes=[None, QDType.sfp8_e4m3_nan], ) q = Quantizer(config=cfg, develop_dtype=torch.float32) # Compute effective residual weight if lora is present if lora is not None: # unsmoothed weight for quantization w_res = weight - (lora[1] @ lora[0]) else: w_res = weight res = q.quantize(w_res, return_with_dequant=True, return_with_quant=True) scale_sd = res.scale.state_dict("scale") sd = convert_to_nunchaku_w4x4y16_linear_state_dict( weight=w_res, scale=scale_sd["scale.0"], subscale=scale_sd["scale.1"], bias=bias, smooth=smooth, lora=lora, float_point=True, ) # Rename keys to match Nunchaku SVDQW4A4Linear parameter names packed: dict[str, torch.Tensor] = {} packed["qweight"] = sd["qweight"].cpu() packed["wscales"] = sd["wscales"].cpu() if per_channel: packed["wcscales"] = sd["wcscales"].cpu() else: packed["wtscale"] = sd["wtscale"].cpu() packed["bias"] = sd["bias"].cpu() packed["smooth_factor"] = sd["smooth"].cpu() packed["smooth_factor_orig"] = sd["smooth_orig"].cpu() if lora is not None: packed["proj_down"] = sd["lora_down"].cpu() packed["proj_up"] = sd["lora_up"].cpu() return packed def quantize_and_pack_qkv_nvfp4( weights: list[torch.Tensor], biases: list[torch.Tensor], lora: tuple[torch.Tensor, torch.Tensor] | None = None, device: str = "cuda", ) -> dict[str, torch.Tensor]: """Quantize concatenated Q, K, V projections to NVFP4 with independent per-tensor Level-0 scales. In DeepCompressor / Nunchaku, QKV projections are concatenated, but each individual projection (e.g. Q, K, V) has its own distinct dynamic range. Applying per-tensor scaling independently to each projection and then concatenating the Level-0 scales into `wcscales` prevents catastrophic scale underflow and FP8 subscale clipping (which otherwise causes pebbled noise). Args: weights: List of weight tensors [w_q, w_k, w_v] in BF16. biases: List of bias tensors [b_q, b_k, b_v] in BF16. lora: Optional tuple of (lora_down, lora_up) for joint low-rank branch. device: Device to perform quantization on. Returns: dict containing: - qweight: [out_features, in_features // 2] int8 - wscales: [in_features // 16, out_features] float8_e4m3fn - wcscales: [out_features] bfloat16 - bias: [out_features] bfloat16 - smooth_factor: [in_features] bfloat16 - smooth_factor_orig: [in_features] bfloat16 - proj_down: [in_features, rank] (if lora is present) - proj_up: [out_features, rank] (if lora is present) """ weights = [w.detach().clone().to(device=device, dtype=torch.bfloat16) for w in weights] biases = [b.detach().clone().to(device=device, dtype=torch.bfloat16) for b in biases] w_total = torch.cat(weights, dim=0) b_total = torch.cat(biases, dim=0) if lora is not None: ld, lu = lora ld = ld.detach().clone().to(device=device, dtype=torch.bfloat16) lu = lu.detach().clone().to(device=device, dtype=torch.bfloat16) w_res = (w_total.float() - (lu.float() @ ld.float())).to(torch.bfloat16) cur = 0 w_res_slices = [] for w in weights: w_res_slices.append(w_res[cur:cur + w.shape[0]]) cur += w.shape[0] else: w_res = w_total w_res_slices = weights ld, lu = None, None cfg = QuantizerConfig( dtype=QDType.sfp4_e2m1_all, group_shapes=[[-1, -1], [1, 16, 1, 1, 1]], scale_dtypes=[None, QDType.sfp8_e4m3_nan], ) q = Quantizer(config=cfg, develop_dtype=torch.float32) scales_0 = [] subscales_1 = [] for w_s in w_res_slices: res = q.quantize(w_s, return_with_quant=True) sd_s = res.scale.state_dict("scale") scales_0.append(sd_s["scale.0"].view(-1).expand(w_s.shape[0]).reshape(w_s.shape[0], 1, 1, 1)) subscales_1.append(sd_s["scale.1"]) scale_cat = torch.cat(scales_0, dim=0) subscale_cat = torch.cat(subscales_1, dim=0) sd = convert_to_nunchaku_w4x4y16_linear_state_dict( weight=w_res, scale=scale_cat, subscale=subscale_cat, bias=b_total, lora=(ld, lu) if lora is not None else None, float_point=True, ) packed: dict[str, torch.Tensor] = { "qweight": sd["qweight"].cpu(), "wscales": sd["wscales"].cpu(), "wcscales": sd["wcscales"].cpu(), "bias": sd["bias"].cpu(), "smooth_factor": sd["smooth"].cpu(), "smooth_factor_orig": sd["smooth_orig"].cpu(), } if lora is not None: packed["proj_down"] = sd["lora_down"].cpu() packed["proj_up"] = sd["lora_up"].cpu() return packed def quantize_and_pack_adanorm_linear( weight: torch.Tensor, bias: torch.Tensor | None = None, device: str = "cuda", ) -> dict[str, torch.Tensor]: """Quantize a modulation linear layer (img_mod.1 or txt_mod.1) to AWQ INT4 with 6 splits. Args: weight: Original weight tensor [18432, 3072] in BF16/FP16. bias: Original bias tensor [18432] in BF16/FP16. device: Device to perform quantization on. Returns: dict containing: - qweight: [4608, 1536] int32 - wscales: [48, 18432] bf16 - wzeros: [48, 18432] bf16 - bias: [18432] bf16 """ weight = weight.detach().clone().to(device=device) if bias is not None: bias = bias.detach().clone().to(device=device) else: bias = torch.zeros(weight.shape[0], dtype=weight.dtype, device=device) cfg = QuantizerConfig( dtype=QDType.sint4, group_shapes=[[1, 64, 1, 1, 1]], scale_dtypes=[None], ) q = Quantizer(config=cfg, develop_dtype=torch.float32) res = q.quantize(weight, return_with_dequant=True, return_with_quant=True) scale_sd = res.scale.state_dict("scale") sd = convert_to_nunchaku_w4x16_adanorm_zero_state_dict( weight=weight, scale=scale_sd["scale.0"], bias=bias, ) packed: dict[str, torch.Tensor] = {} packed["qweight"] = sd["qweight"].cpu() packed["wscales"] = sd["wscales"].cpu() packed["wzeros"] = sd["wzeros"].cpu() packed["bias"] = sd["bias"].cpu() return packed