Download loopq_quantization/scripts/loopq/packed_weight.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 4.35 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/packed_weight.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/packed_weight.py
-
curl -L -o packed_weight.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/packed_weight.py
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)} | |
| 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') | |
| 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()) | |
| 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 | |