catplusplus's picture
Upload folder using huggingface_hub
b6a8400 verified
Raw History Blame Contribute Delete
9.52 kB
# -*- 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