"""Lossless storage for LoopQ signed group-32 W4 QDQ weights. This is a storage codec, not an INT4 GEMM kernel. Dequantization materializes one dense weight; resident-memory and runtime-speed claims require integration. """ from __future__ import annotations from dataclasses import dataclass import math import torch from .quantization import quantize_weight _DTYPES = {str(dtype): dtype for dtype in (torch.float16, torch.bfloat16, torch.float32, torch.float64)} @dataclass(frozen=True) class PackedLoopQWeight: codes: torch.Tensor scales: torch.Tensor shape: tuple[int, int] output_dtype: str format_version: int = 2 def validate(self): if len(self.shape) != 2 or any(type(x) is not int or x <= 0 for x in self.shape): raise ValueError('packed weight shape must contain two positive integers') rows, width = self.shape if self.codes.dtype != torch.uint8 or self.codes.shape != (rows, math.ceil(width/2)): raise ValueError('packed code shape/dtype mismatch') if self.scales.dtype != torch.float32 or self.scales.shape != (rows, math.ceil(width/32)): raise ValueError('group-32 FP32 scale shape/dtype mismatch') if not torch.isfinite(self.scales).all() or not (self.scales > 0).all(): raise ValueError('scales must be finite and positive') if self.output_dtype not in _DTYPES: raise ValueError('unsupported output dtype') if width % 2 and (self.codes[:, -1] >> 4).any(): raise ValueError('nonzero unused high nibble') if self.format_version not in (1, 2): raise ValueError('unsupported packed-weight version') if self.format_version == 1 and ( ((self.codes & 15) == 8).any() or ((self.codes >> 4) == 8).any()): raise ValueError('reserved -8 code is outside LoopQ narrow signed range') @property def payload_bytes(self): """Tensor payload only; excludes serialization/container overhead.""" return self.codes.numel() + self.scales.numel() * 4 def dequantize(self, *, device=None): self.validate() return self._dequantize_validated(device=device) def _dequantize_validated(self, *, device=None): """Internal immutable-dispatch path; the owner validates before use.""" codes = self.codes.to(device=device) if device is not None else self.codes pairs = torch.stack((codes & 15, codes >> 4), dim=-1).flatten(-2) values = pairs[:, :self.shape[1]].to(torch.int16) values = torch.where(values >= 8, values - 16, values).float() scales = self.scales.to(device=values.device).repeat_interleave(32, dim=-1)[:, :self.shape[1]] return (values * scales).to(_DTYPES[self.output_dtype]) def state_dict(self): self.validate() return dict(format='loopq_packed_w4_group32', format_version=self.format_version, nibble_order='even_column_low', shape=list(self.shape), output_dtype=self.output_dtype, codes=self.codes.detach().cpu(), scales=self.scales.detach().cpu()) @classmethod def from_state_dict(cls, state): if state.get('format') != 'loopq_packed_w4_group32' or state.get('format_version') not in (1, 2) \ or state.get('nibble_order') != 'even_column_low': raise ValueError('unsupported packed-weight contract') result = cls(state['codes'], state['scales'], tuple(state['shape']), state['output_dtype'], state['format_version']) result.validate() return result def pack_weight(weight: torch.Tensor) -> PackedLoopQWeight: if weight.ndim != 2 or min(weight.shape) <= 0: raise ValueError('weight must be a nonempty matrix') if not torch.isfinite(weight).all(): raise ValueError('weight must be finite') quantized = quantize_weight(weight.detach()) integers = quantized.integers unsigned = (integers.to(torch.int16) & 15).to(torch.uint8) if weight.shape[1] % 2: unsigned = torch.nn.functional.pad(unsigned, (0, 1)) codes = unsigned[:, 0::2] | (unsigned[:, 1::2] << 4) result = PackedLoopQWeight(codes, quantized.scales.detach().float(), tuple(weight.shape), str(weight.dtype)) result.validate() return result