diff --git a/paroquant/cli/convert.py b/paroquant/cli/convert.py index 8055cec..f9edbbc 100644 --- a/paroquant/cli/convert.py +++ b/paroquant/cli/convert.py @@ -15,6 +15,8 @@ from transformers import AutoModelForCausalLM, AutoModelForImageTextToText from paroquant.optim.qexperts import PseudoQuantizedMoEExperts, get_named_moe_experts, is_fused_moe_experts from paroquant.optim.util import get_named_linears, set_module_by_name +from paroquant.optim.quant import pow2_project, pow2_scales_enabled, quant_format + _AWQ_REORDER = (0, 2, 4, 6, 1, 3, 5, 7) _LAYER_PATHS = ["model.layers", "model.language_model.layers", "language_model.layers"] @@ -191,6 +193,29 @@ def _quantize_rotated_weight( return quantized, scales_2d, zeros_2d +def _quantize_rotated_weight_mxfp4( + *, + weight: torch.Tensor, + pairs: torch.Tensor, + theta: torch.Tensor, + channel_scales: torch.Tensor, + exp_bias: torch.Tensor | None, + group_size: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Rotate exactly as the optimizer did, then emit MXFP4 buffers. + + The rotation must go through the same kernel the optimizer used, or the codes will describe + a different weight than the one that was calibrated. + """ + from paroquant.kernels.cuda import scaled_pairwise_rotation + from paroquant.optim.mxfp4 import quantize_to_codes + + if channel_scales.ndim == 1: + channel_scales = channel_scales.unsqueeze(0) + rotated = scaled_pairwise_rotation(weight * channel_scales, pairs, theta, None, group_size) + return quantize_to_codes(rotated, exp_bias) + + def _to_awq_buffers( quantized: torch.Tensor, scales_2d: torch.Tensor, @@ -236,6 +261,31 @@ def _convert_pseudo(model: torch.nn.Module, result_dir: Path) -> int: return count +class _MXFP4BufferHolder(torch.nn.Module): + """Holds MXFP4 buffers so save_pretrained writes them under the module's own names. + + Deliberately not RotateQuantizedLinear: that class registers AWQ buffers (qweight/qzeros/ + scales) and imports autoawq's GEMM at module import, neither of which applies to MXFP4. It + also means a load_state_dict(..., strict=False) would silently DROP every MXFP4 buffer and + leave the AWQ ones zeroed, writing a checkpoint of zeros. + """ + + def __init__(self, in_features: int, out_features: int, bias: bool, krot: int): + super().__init__() + self.in_features, self.out_features = in_features, out_features + self.register_buffer("weight", torch.zeros(out_features, in_features // 2, + dtype=torch.uint8)) + self.register_buffer("weight_scale", torch.zeros(out_features, in_features // 32, + dtype=torch.uint8)) + self.register_buffer("theta", torch.zeros(krot, in_features // 2, dtype=torch.float16)) + self.register_buffer("pairs", torch.zeros(krot, in_features, dtype=torch.int16)) + self.register_buffer("channel_scales", torch.ones(1, in_features, dtype=torch.float16)) + if bias: + self.register_buffer("bias", torch.zeros(out_features, dtype=torch.float16)) + else: + self.bias = None + + @torch.no_grad() def _quantize_layer(state_dict: dict, device: str) -> tuple[dict[str, torch.Tensor], int, int, int]: weight = state_dict["weight"].to(device=device, dtype=torch.float32) @@ -247,7 +297,38 @@ def _quantize_layer(state_dict: dict, device: str) -> tuple[dict[str, torch.Tens theta = _stack_if_numbered(state_dict, "angles_grouped").to(device=device, dtype=torch.float32) channel_scales_opt = state_dict["channel_scales"].to(device=device, dtype=torch.float32) + + if quant_format() == "mxfp4": + exp_bias = state_dict.get("quantizer.exp_bias") + if exp_bias is not None: + exp_bias = exp_bias.to(device=device, dtype=torch.float32) + packed, e8m0 = _quantize_rotated_weight_mxfp4( + weight=weight, + pairs=pairs, + theta=theta, + channel_scales=channel_scales_opt, + exp_bias=exp_bias, + group_size=group_size, + ) + channel_scales_mx = (1.0 / channel_scales_opt).to(torch.float16).cpu() + if channel_scales_mx.ndim == 1: + channel_scales_mx = channel_scales_mx.unsqueeze(0) + buffers = { + "weight": packed.cpu(), + "weight_scale": e8m0.cpu(), + "theta": theta.to(torch.float16).cpu(), + "pairs": pairs.cpu(), + "channel_scales": channel_scales_mx, + } + if "bias" in state_dict and state_dict["bias"] is not None: + buffers["bias"] = state_dict["bias"].to(torch.float16).cpu() + return buffers, bits, group_size, int(theta.shape[0]) + scales_flat = state_dict["quantizer.scale"].to(device=device, dtype=torch.float32).reshape(-1, 1) + if pow2_scales_enabled(): + # Must match the projection the optimizer trained under, or the codes written + # below will not correspond to the scales written beside them. + scales_flat = pow2_project(scales_flat) zp_flat = state_dict["quantizer.zero_point_float"].to(device=device, dtype=torch.float32).reshape(-1, 1) quantized, scales_2d, zeros_2d = _quantize_rotated_weight( weight=weight, @@ -410,7 +491,9 @@ def _convert_real( model: torch.nn.Module, result_dir: Path, ) -> tuple[int, dict[str, Any] | None, dict[str, torch.Tensor]]: - from paroquant.inference.backends.transformers.modules import RotateQuantizedLinear + if quant_format() != "mxfp4": + # imports autoawq's GEMM at module scope; irrelevant (and absent) for MXFP4 + from paroquant.inference.backends.transformers.modules import RotateQuantizedLinear blocks = _get_blocks(model) count = 0 @@ -426,15 +509,26 @@ def _convert_real( sd = torch.load(pt_file, map_location="cpu", weights_only=False) buffers, bits, group_size, krot = _quantize_layer(sd, device="cuda") - rl = RotateQuantizedLinear( - module.in_features, - module.out_features, - bias=module.bias is not None, - group_size=group_size, - bits=bits, - krot=krot, - ) - rl.load_state_dict(buffers, strict=False) + if quant_format() == "mxfp4": + rl = _MXFP4BufferHolder( + module.in_features, + module.out_features, + bias=module.bias is not None, + krot=krot, + ) + # strict: every buffer we emitted must land somewhere, or the checkpoint is + # silently written as zeros. + rl.load_state_dict(buffers, strict=True) + else: + rl = RotateQuantizedLinear( + module.in_features, + module.out_features, + bias=module.bias is not None, + group_size=group_size, + bits=bits, + krot=krot, + ) + rl.load_state_dict(buffers, strict=False) set_module_by_name(layer, name, rl) count += 1 @@ -455,7 +549,11 @@ def _convert_real( torch.cuda.empty_cache() quant_config = { - "quant_method": "paroquant", + # A distinct method name rather than a "format" discriminator inside "paroquant": the + # buffers are completely different (packed e2m1 + e8m0 vs AWQ int4), so the shipped + # int4 serving path should not have to branch at load time. + "quant_method": "paroquant_mxfp4" if quant_format() == "mxfp4" else "paroquant", + "format": quant_format(), "bits": bits, "group_size": group_size, "krot": krot, @@ -481,7 +579,10 @@ def main() -> None: raise FileNotFoundError(f"Result directory not found: {result_dir}") source_dir = _resolve_source_dir(args.model) - model = _load_model(str(source_dir), device_map="cpu" if args.mode == "real" else "cuda") + # Both modes keep the model on CPU. _convert_pseudo already moves each block to GPU and back + # itself, so device_map="cuda" only forces the whole fp16 model resident up front -- 55.6 GiB + # for a 27B, which does not fit a 32 GiB card. + model = _load_model(str(source_dir), device_map="cpu") quant_config: dict[str, Any] | None = None save_state_dict: dict[str, torch.Tensor] | None = None diff --git a/paroquant/cli/optimize.py b/paroquant/cli/optimize.py index 9469078..51803e7 100644 --- a/paroquant/cli/optimize.py +++ b/paroquant/cli/optimize.py @@ -1,6 +1,7 @@ from __future__ import annotations import json +import os from copy import deepcopy from dataclasses import dataclass, field from pathlib import Path @@ -18,6 +19,7 @@ from paroquant.optim.qexperts import PseudoQuantizedMoEExperts, get_named_moe_ex from paroquant.optim.util import ( set_module_by_name, load_model, + move_capture_modules, move_embed, load_tokenizer, get_blocks, @@ -102,6 +104,61 @@ def setup_wandb(args: Config) -> wandb.Run | None: return wandb_run + + +def _release_module_memory(module: nn.Module) -> None: + """Drop a module's parameter storage once the optimizer will never touch it again. + + The layer-wise loop only moves forward -- a finished block's outputs are cached and the + block itself is never revisited -- yet the whole fp16 model stays CPU-resident for the run. + On a 60 GiB box holding a 55.6 GiB model that means permanent swapping and, past ~18 layers, + thrashing. Freeing each finished block (~0.9 GiB) as we go brings the footprint under RAM. + Parameters only (buffers in these blocks are tiny) and no explicit gc: a first version that + also walked buffers and called gc.collect() after capture never reached the layer loop. + Memory only: the saved .pt files are what convert reads, so results are unaffected. + """ + with torch.no_grad(): + for p_ in module.parameters(recurse=True): + p_.data = torch.empty(0, dtype=p_.dtype) + + +def _trained_rotations_index(path: str | None) -> dict[tuple[int, str], str] | None: + """(layer_idx, module_name) -> key prefix in a trained ParoQuant checkpoint, or None. + + PARO_INIT_ROTATIONS= initializes every linear's pairs/theta/ + channel_scales from an already-trained checkpoint instead of identity. Rationale: from + identity, with the calibration budget a 60 GiB box can afford, the rotation stage stops + contributing past ~layer 10 and leaves angles at exactly zero on many layers; a trained + rotation transfers across weight grids (it tames per-group outliers, which helps e2m1 as + much as int4). Pair it with a params spec that omits angles/channel_scales to freeze them. + """ + if not path: + return None + import re + from safetensors import safe_open + idx: dict[tuple[int, str], str] = {} + with safe_open(path, framework="pt") as f: + for k in f.keys(): + m = re.match(r"^(.*\.layers\.(\d+)\.(.+))\.theta$", k) + if m: + idx[(int(m.group(2)), m.group(3))] = m.group(1) + logger.info(f"trained rotations: {len(idx)} modules from {path}") + return idx + + +def _load_trained_rotation(path: str, prefix: str, device, weight_dtype): + from safetensors import safe_open + with safe_open(path, framework="pt") as f: + pairs = f.get_tensor(f"{prefix}.pairs").to(device=device, dtype=torch.int16) + theta = f.get_tensor(f"{prefix}.theta").to(device=device, dtype=weight_dtype) + cs = f.get_tensor(f"{prefix}.channel_scales").to(device=device, dtype=torch.float32) + mask = torch.zeros_like(theta, dtype=torch.bool) + # stored pre-inverted (the serving prologue multiplies activations by it); the optimizer's + # PseudoQuantizedLinear multiplies the WEIGHT, i.e. wants the reciprocal + cs_opt = (1.0 / cs).view(1, -1).to(weight_dtype) + return [pairs, theta, mask], cs_opt + + def main(): args = simple_parsing.parse(Config, add_option_string_dash_variants=simple_parsing.DashVariant.DASH) print(args) @@ -132,6 +189,9 @@ def main(): with open(output_dir / "args.json", "w") as f: json.dump(vars(args), f, indent=2) + trained_path = os.environ.get("PARO_INIT_ROTATIONS") or None + trained_index = _trained_rotations_index(trained_path) + # Load model. model = load_model(args.model, device_map="cpu", dtype=torch.float16).half() move_embed(model, device) @@ -161,7 +221,7 @@ def main(): # Capture per-batch positional args and layer kwargs. logger.info("Capturing layer positional args and kwargs...") - model.to(device) + move_capture_modules(model, blocks, device) ( og_layer_input_batches, kwargs_list, @@ -366,13 +426,22 @@ def main(): old_module.to(device) if isinstance(old_module, nn.Linear): weight = old_module.weight.float() - rotation_pairs = init_rotation_data( - weight, - seed=args.seed + layer_idx, - group_size=args.group_size, - num_rotations=args.num_rotations, - ) - channel_scales = torch.ones(1, weight.shape[1], dtype=torch.float16, device=device) + trained_key = (layer_idx, name) + if trained_index is not None and trained_key in trained_index: + rotation_pairs, channel_scales = _load_trained_rotation( + trained_path, trained_index[trained_key], device, torch.float16) + assert rotation_pairs[0].shape[0] == args.num_rotations, ( + rotation_pairs[0].shape, args.num_rotations) + else: + if trained_index is not None: + logger.warning(f"no trained rotation for layer {layer_idx} {name}; using identity") + rotation_pairs = init_rotation_data( + weight, + seed=args.seed + layer_idx, + group_size=args.group_size, + num_rotations=args.num_rotations, + ) + channel_scales = torch.ones(1, weight.shape[1], dtype=torch.float16, device=device) new_module = PseudoQuantizedLinear( old_module, @@ -541,6 +610,7 @@ def main(): if all_files_exist: layer.cpu() + _release_module_memory(layer) continue # Save the optimized result @@ -552,6 +622,7 @@ def main(): ) layer.cpu() + _release_module_memory(layer) if wandb_run is not None: wandb_run.finish() diff --git a/paroquant/kernels/cuda/__init__.py b/paroquant/kernels/cuda/__init__.py index b2eca39..dd23dee 100644 --- a/paroquant/kernels/cuda/__init__.py +++ b/paroquant/kernels/cuda/__init__.py @@ -1,3 +1,4 @@ +import os import shutil import sys from pathlib import Path @@ -13,7 +14,7 @@ def _rotation_build_directory() -> Path: abi_tag = ( f"py{sys.version_info.major}{sys.version_info.minor}_" f"torch{torch.__version__}_" - f"cu{torch.version.cuda or 'none'}" + f"cu{torch.version.cuda or 'none'}_hip{torch.version.hip or 'none'}" ) abi_tag = "".join(c if c.isalnum() else "_" for c in abi_tag) build_dir = cache_root / "paroquant_rotation" / abi_tag @@ -21,7 +22,74 @@ def _rotation_build_directory() -> Path: return build_dir +def _load_rotation_extension_rocm(): + """Build the rotation kernel for ROCm/HIP. + + torch.utils.cpp_extension.load() only hipifies the files named in `sources`, so rotation.cuh + is left as CUDA and the compile fails on the include. Hipify the whole directory instead, + fix up the one bf16 intrinsic HIP does not provide, and build with hipcc directly. The + kernel registers itself through TORCH_LIBRARY, so load_library is enough -- pybind.cpp is an + empty module upstream and is not needed here. + """ + import shutil + import subprocess + import sysconfig + + from torch.utils.hipify import hipify_python + + build_dir = _rotation_build_directory() + so = build_dir / "paroquant_rotation.so" + newest_src = max(f.stat().st_mtime for f in _dir.iterdir() if f.suffix in (".cu", ".cuh")) + if not (so.exists() and so.stat().st_mtime >= newest_src): + work = build_dir / "src" + shutil.rmtree(work, ignore_errors=True) + work.mkdir(parents=True, exist_ok=True) + for f in _dir.iterdir(): + if f.suffix in (".cu", ".cuh"): + shutil.copy(f, work / f.name) + hipify_python.hipify( + project_directory=str(work), + output_directory=str(work), + includes=[str(work / "*")], + is_pytorch_extension=True, + show_detailed=False, + ) + # HIP's bf16 header has no __floats2bfloat162_rn; build the pair explicitly. + header = work / "rotation_hip.cuh" + text = header.read_text() + needle = "return __floats2bfloat162_rn(a, b);" + if needle in text: + header.write_text( + text.replace( + needle, + "__hip_bfloat162 r; r.x = __float2bfloat16(a); " + "r.y = __float2bfloat16(b); return r;", + ) + ) + torch_dir = Path(torch.__file__).parent + cmd = [ + os.environ.get("HIPCC", "/opt/rocm/bin/hipcc"), + "-O3", "-std=c++17", "-fPIC", "-shared", "-ffast-math", "-w", + f"--offload-arch={os.environ.get('PYTORCH_ROCM_ARCH', 'gfx1201')}", + f"-I{torch_dir}/include", + f"-I{torch_dir}/include/torch/csrc/api/include", + f"-I{sysconfig.get_paths()['include']}", + "-I/opt/rocm/include", + "-D__HIP_PLATFORM_AMD__=1", "-DUSE_ROCM=1", + str(work / "rotation.hip"), "-o", str(so), + f"-L{torch_dir}/lib", + "-ltorch", "-ltorch_cpu", "-ltorch_hip", "-lc10", "-lc10_hip", + ] + r = subprocess.run(cmd, capture_output=True, text=True) + if r.returncode != 0: + raise RuntimeError(f"hipcc failed building the rotation kernel:\n{r.stderr[-4000:]}") + torch.ops.load_library(str(so)) + return None + + def _load_rotation_extension(): + if torch.version.hip is not None: + return _load_rotation_extension_rocm() build_dir = _rotation_build_directory() load_kwargs = dict( name="paroquant_rotation", diff --git a/paroquant/optim/mxfp4.py b/paroquant/optim/mxfp4.py new file mode 100644 index 0000000..4789052 --- /dev/null +++ b/paroquant/optim/mxfp4.py @@ -0,0 +1,173 @@ +"""MXFP4 (OCP microscaling) fake-quantization for the ParoQuant optimizer. + +Motivation: the radiance serving GEMM's prefill cost is linear in VALU per tile-group, and the +two FMAs are (a) the fp16 group scale and (b) the asymmetric zero point. MXFP4 removes both by +construction -- its e8m0 scale is a power of two, so it folds at weight staging, and e2m1 is a +signed float with no zero point. A ParoQuant checkpoint in MXFP4 therefore serves through the +existing MXFP4 kernel's zero-VALU loop while keeping the learned rotations that make 4-bit work. + +Format: blocks of 32 along the input dimension. One shared e8m0 scale per block (a power of two), +each element an e2m1 float with magnitudes {0, .5, 1, 1.5, 2, 3, 4, 6}. +""" +from __future__ import annotations + +import os + +import torch +import torch.nn as nn + +from .quant import clamp_ste, round_ste + +BLOCK = 32 +# e2m1 magnitudes, and the midpoints between them for round-to-nearest. +_GRID = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) +_MIDS = (0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0) +_EMAX = 6.0 +# e8m0 spans 2^-127..2^127; nothing here goes near the ends, but clamp so a degenerate block +# cannot produce an unrepresentable scale. +_EXP_MIN, _EXP_MAX = -127.0, 127.0 +_cache: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {} + + +def scale_rule() -> str: + """Shared-exponent rule: "ocp" (default, AMD/Quark-compatible) or "noclip". + + Measured on real Qwen3.8-27B weights (layers 0-1, down/gate/out_proj): ocp beats noclip by + ~1.9% RMSE on every tensor. Using all eight grid levels is worth more than never clipping + the block maximum. It also matches AMD's release, so a comparison against their MXFP4 + isolates the rotations instead of confounding them with the scale convention. + """ + rule = os.environ.get("PARO_MXFP4_SCALE_RULE", "ocp").lower() + if rule not in ("noclip", "ocp"): + raise ValueError(f"PARO_MXFP4_SCALE_RULE must be 'noclip' or 'ocp', got {rule!r}") + return rule + + +def _tables(device, dtype): + key = (device, dtype) + if key not in _cache: + _cache[key] = ( + torch.tensor(_GRID, device=device, dtype=dtype), + torch.tensor(_MIDS, device=device, dtype=dtype), + ) + return _cache[key] + + +def round_to_e2m1(v: torch.Tensor) -> torch.Tensor: + """Round to the nearest representable e2m1 value (ties away from zero, as bucketize gives).""" + grid, mids = _tables(v.device, v.dtype) + mag = v.abs().clamp(max=_EMAX) + return torch.sign(v) * grid[torch.bucketize(mag, mids)] + + +def shared_exponent(xs: torch.Tensor, exp_bias: torch.Tensor | None = None) -> torch.Tensor: + """The e8m0 exponent for each block of `xs` ([n_blocks, BLOCK]). + + See `scale_rule()` for the two conventions and which one measured better. + """ + amax = xs.abs().amax(dim=1, keepdim=True) + tiny = torch.finfo(torch.float32).tiny + if scale_rule() == "ocp": + # What AMD's Quark MXFP4 release uses (confirmed from the checkpoint: per-block maxima + # are only 4.0 and 6.0, never 3.0, which is this rule's fingerprint). amax/scale lands + # in [4, 8) so the block's largest elements clip to 6, in exchange for using all eight + # grid levels. + e = torch.floor(torch.log2(amax.clamp_min(tiny))) - 2.0 + else: + # Never clip: the smallest power of two with amax/scale <= 6, so the ratio sits in + # [3, 6]. Costs the top of the grid, buys exact representation of the block maximum. + e = torch.ceil(torch.log2((amax / _EMAX).clamp_min(tiny))) + if exp_bias is not None: + # A learned per-block nudge, kept small: one exponent step is a factor of two, so this + # trades headroom against resolution and should not wander. + e = e + clamp_ste(exp_bias, -2.0, 2.0) + e = round_ste(e) + return e.clamp(_EXP_MIN, _EXP_MAX) + + +def mxfp4_fake_quant(x: torch.Tensor, exp_bias: torch.Tensor | None = None) -> torch.Tensor: + """Quantize to MXFP4 and back, straight-through.""" + dtype = x.dtype + xf = x.float() + assert xf.shape[-1] % BLOCK == 0, xf.shape + xs = xf.reshape(-1, BLOCK) + scale = torch.exp2(shared_exponent(xs, exp_bias)) + v = xs / scale + q = round_to_e2m1(v) + q = (q - v).detach() + v # STE: gradient flows to the weight and to exp_bias + return (q * scale).reshape(xf.shape).to(dtype) + + +class MXFP4Quantizer(nn.Module): + """Drop-in for UniformAffineQuantizer with an MXFP4 grid. + + MXFP4's scale is derived from the block rather than learned, so the only optimizable + parameter is a per-block exponent bias. Keeping one means the second optimization stage + ("quantizer:") still has a parameter group and the surrounding plumbing is unchanged. + """ + + def __init__(self, weight: torch.Tensor, n_bits, group_size): + super().__init__() + del n_bits, group_size # MXFP4 fixes both: 4 bits, blocks of 32 + assert weight.dim() == 2, weight.shape + assert weight.shape[-1] % BLOCK == 0, weight.shape + n_blocks = weight.numel() // BLOCK + self.exp_bias = nn.Parameter(torch.zeros(n_blocks, 1, device=weight.device, + dtype=torch.float32)) + self.enable_checkpoint = False + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return mxfp4_fake_quant(x, self.exp_bias) + + def optim_params(self) -> list[nn.Parameter]: + return [self.exp_bias] + + def set_optim_enabled(self, enabled: bool): + for param in self.optim_params(): + param.requires_grad = enabled + + +def quantize_to_codes( + x: torch.Tensor, exp_bias: torch.Tensor | None = None +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize to the on-disk MXFP4 buffers. + + Returns (packed_codes, e8m0) where packed_codes is uint8 [out, in/2] holding two e2m1 nibbles + per byte -- element 2i in the LOW nibble, 2i+1 in the high one, matching AMD's Quark release -- + and e8m0 is uint8 [out, in/32] holding the biased shared exponent (value = 2^(E-127)). + + e2m1 bit layout: bit3 sign, bits2:1 exponent, bit0 mantissa, so the magnitude's index into + _GRID is exactly the low three bits. + """ + assert x.dim() == 2, x.shape + out_f, in_f = x.shape + assert in_f % BLOCK == 0, x.shape + xf = x.float() + xs = xf.reshape(-1, BLOCK) + + e = shared_exponent(xs, exp_bias) + scale = torch.exp2(e) + v = (xs / scale).reshape(out_f, in_f) + + grid, mids = _tables(v.device, v.dtype) + idx = torch.bucketize(v.abs().clamp(max=_EMAX), mids).to(torch.uint8) # 0..7 == low 3 bits + sign = (v < 0).to(torch.uint8) << 3 + codes = (idx | sign).reshape(out_f, in_f) + # -0 is representable but pointless; normalize it to +0 so the buffer is canonical. + codes = torch.where(idx == 0, torch.zeros_like(codes), codes) + + packed = (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous() + e8m0 = (e.reshape(out_f, in_f // BLOCK) + 127.0).round().clamp(0, 254).to(torch.uint8) + return packed, e8m0.contiguous() + + +def dequantize_from_codes(packed: torch.Tensor, e8m0: torch.Tensor) -> torch.Tensor: + """Inverse of `quantize_to_codes`, for verification.""" + out_f, half = packed.shape + lo, hi = packed & 0x0F, (packed >> 4) & 0x0F + codes = torch.stack([lo, hi], dim=-1).reshape(out_f, half * 2) + grid, _ = _tables(packed.device, torch.float32) + mag = grid[(codes & 0x7).long()] + val = torch.where((codes & 0x8) != 0, -mag, mag) + scale = torch.exp2(e8m0.float() - 127.0) + return (val.reshape(-1, BLOCK) * scale.reshape(-1, 1)).reshape(out_f, half * 2) diff --git a/paroquant/optim/qlinear.py b/paroquant/optim/qlinear.py index 384b66e..85be61d 100644 --- a/paroquant/optim/qlinear.py +++ b/paroquant/optim/qlinear.py @@ -5,6 +5,7 @@ from torch import nn from torch.utils.checkpoint import checkpoint from .util import get_named_linears +from .quant import quant_format from .quantizer import UniformAffineQuantizer from paroquant.kernels.cuda import scaled_pairwise_rotation @@ -153,7 +154,13 @@ class PseudoQuantizedLinear(nn.Module): self.group_size, ) - self.quantizer = UniformAffineQuantizer( + if quant_format() == "mxfp4": + from .mxfp4 import MXFP4Quantizer + + quantizer_cls = MXFP4Quantizer + else: + quantizer_cls = UniformAffineQuantizer + self.quantizer = quantizer_cls( weight, n_bits=self.n_bits, group_size=self.group_size, @@ -214,8 +221,10 @@ class PseudoQuantizedLinear(nn.Module): num_rotations=num_rotations, ) - # Initialize the quantizer - if "quantizer.scale" in state_dict: + # Initialize the quantizer. MXFP4 saves quantizer.exp_bias where the uniform-affine + # quantizer saves quantizer.scale; without this the submodule is never built and + # load_state_dict below rejects the key. + if any(k.startswith("quantizer.") for k in state_dict): qlinear.set_optim_enabled(quantizer=True) qlinear.load_state_dict(state_dict) diff --git a/paroquant/optim/quant.py b/paroquant/optim/quant.py index 4dea8dd..5a8022a 100644 --- a/paroquant/optim/quant.py +++ b/paroquant/optim/quant.py @@ -1,5 +1,7 @@ from __future__ import annotations +import os + import torch @@ -11,3 +13,41 @@ def round_ste(x: torch.Tensor) -> torch.Tensor: def clamp_ste(x: torch.Tensor, min: float | None = None, max: float | None = None) -> torch.Tensor: """Straight-through estimator for clamping.""" return (x.clamp(min, max) - x).detach() + x + + +def pow2_scales_enabled() -> bool: + """Constrain group scales to powers of two (PARO_POW2_SCALES=1). + + Read from the environment rather than threaded through the CLI so that the optimizer and + `paroquant.cli.convert` cannot disagree: if they projected differently the exported codes + would not match the exported scales. + """ + return os.environ.get("PARO_POW2_SCALES", "0") == "1" + + +def pow2_project(x: torch.Tensor) -> torch.Tensor: + """Round each scale to the nearest power of two. + + The exponent is clamped to the fp16 normal range so that the value survives the fp16 + checkpoint export exactly -- which is the whole point: an exactly-representable pow2 scale + lets an inference kernel fold it at weight-staging time instead of paying a multiply per + output element per group. + """ + return torch.exp2(torch.clamp(torch.round(torch.log2(x)), -14.0, 15.0)) + + +def pow2_ste(x: torch.Tensor) -> torch.Tensor: + """Straight-through estimator for the pow2 projection.""" + return (pow2_project(x) - x).detach() + x + + +def quant_format() -> str: + """Weight grid: "int" (upstream uniform-affine) or "mxfp4" (OCP microscaling e2m1/e8m0). + + Env-selected rather than threaded through the CLI so that the optimizer and + `paroquant.cli.convert` cannot disagree about the grid the codes were trained on. + """ + fmt = os.environ.get("PARO_QUANT_FORMAT", "int").lower() + if fmt not in ("int", "mxfp4"): + raise ValueError(f"PARO_QUANT_FORMAT must be 'int' or 'mxfp4', got {fmt!r}") + return fmt diff --git a/paroquant/optim/quantizer.py b/paroquant/optim/quantizer.py index 89d83b4..b1a9ac8 100644 --- a/paroquant/optim/quantizer.py +++ b/paroquant/optim/quantizer.py @@ -4,7 +4,7 @@ import torch import torch.nn as nn from torch.utils.checkpoint import checkpoint -from .quant import clamp_ste, round_ste +from .quant import clamp_ste, pow2_scales_enabled, pow2_ste, quant_format, round_ste def _calc_scales_and_zero_points( @@ -96,11 +96,21 @@ class UniformAffineQuantizer(nn.Module): x = x.float() assert torch.isnan(x).sum() == 0, x + if quant_format() == "mxfp4": + # MXFP4 derives its scale per 32-block and has no zero point, so scale/zero_point + # arguments do not apply; qlinear takes this path whenever no learnable quantizer + # has been constructed yet. + from .mxfp4 import mxfp4_fake_quant + + return mxfp4_fake_quant(x).to(dtype) + qmin, qmax = 0, 2**n_bits - 1 if scale is None or zero_point is None: scale, zero_point = _calc_scales_and_zero_points(x, group_size, qmin, qmax) scale = clamp_ste(scale, min=1e-5, max=1e5) + if pow2_scales_enabled(): + scale = pow2_ste(scale) round_zero_point = clamp_ste(-round_ste(zero_point), qmin, qmax) dim1, dim2 = x.shape x = x.reshape(-1, group_size) diff --git a/paroquant/optim/util.py b/paroquant/optim/util.py index b5b7d1a..567ad99 100644 --- a/paroquant/optim/util.py +++ b/paroquant/optim/util.py @@ -103,6 +103,32 @@ def move_embed(model, device): _move("per_layer_projection_norm") +def move_capture_modules(model, blocks, device) -> None: + """Move only the modules `capture_layer_inputs_and_args` actually executes. + + Every decoder block is replaced by a Catcher that returns an empty tensor without ever + calling the real module, so block weights are never read on device during capture. Moving + the whole model instead costs the full fp16 footprint -- 55.6 GiB for a 27B, which does not + fit a 32 GiB card. Swap the blocks (and any vision tower, unused for text calibration) out + of the tree, move what is left, and swap them back on CPU. + """ + saved = [blocks[i] for i in range(len(blocks))] + placeholder = nn.Identity() + for i in range(len(blocks)): + blocks[i] = placeholder + inner = getattr(model, "model", model) + visual = getattr(inner, "visual", None) + if visual is not None: + inner.visual = nn.Identity() + try: + model.to(device) + finally: + for i, block in enumerate(saved): + blocks[i] = block + if visual is not None: + inner.visual = visual + + def empty_cache(): gc.collect() torch.cuda.empty_cache() @@ -155,7 +181,9 @@ def get_calib_dataset( dataset = load_dataset("mit-han-lab/pile-val-backup", split="validation") dataset = dataset.shuffle(seed=seed) elif data == "wikitext2": - dataset = load_dataset("wikitext", "wikitext-2-raw-v1", split=split) + # Canonical id since the dataset was moved under a namespace; the bare "wikitext" + # is rejected outright by newer huggingface_hub URI parsing. + dataset = load_dataset("Salesforce/wikitext", "wikitext-2-raw-v1", split=split) dataset = dataset.shuffle(seed=seed) elif data == "c4": if split == "train": @@ -173,11 +201,8 @@ def get_calib_dataset( dataset = dataset.shuffle(seed=seed) elif data == "redpajama": test_split, val_split = 0.2, 0.1 - dataset = load_dataset( - "liang2kl/RedPajama-Data-1T-Sample-Backup", - split="train", - trust_remote_code=True, - ) + # trust_remote_code was removed in datasets 4.x; this mirror is plain data anyway. + dataset = load_dataset("liang2kl/RedPajama-Data-1T-Sample-Backup", split="train") dataset = dataset.shuffle(seed=seed) test_size = int(len(dataset) * test_split) val_size = int(len(dataset) * val_split)