JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
4.35 kB
"""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