Download loader/prism_quant/quant.py from atomtanstudio/Prism-Q6: direct link, hf CLI and curl.
- Browser
- Download file 8.06 kB
-
https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/quant.py
- Command line
-
hf download hf://atomtanstudio/Prism-Q6/loader/prism_quant/quant.py
-
curl -L -o quant.py https://huggingface.co/atomtanstudio/Prism-Q6/resolve/main/loader/prism_quant/quant.py
8.06 kB
| """Prism custom Q6: symmetric row-group quantization, not GGUF. | |
| Only the packed representation persists. Inference reconstructs one Linear's | |
| weight at a time; device placement of that representation is owned by the caller. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import re | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| FORMAT = "prism-custom-q6" | |
| VERSION = 1 | |
| PROFILE = "blocks-q6-v1" | |
| GROUP_SIZE = 64 | |
| _BLOCK_WEIGHT = re.compile( | |
| r"^(?:fusion_blocks\.\d+\.(?:video_block|audio_block|a2v_conditioner|v2a_conditioner)" | |
| r"|remaining_video_blocks\.\d+|video_dit_2\.blocks\.\d+)\..+\.weight$" | |
| ) | |
| def is_quantized_weight(name: str, shape: tuple | list) -> bool: | |
| return len(shape) == 2 and bool(_BLOCK_WEIGHT.fullmatch(name)) | |
| def _check_group_size(group_size: int) -> None: | |
| if group_size != GROUP_SIZE: | |
| raise ValueError("Version 1 requires group_size=64") | |
| def packed_shape(out_features: int, in_features: int, group_size: int = GROUP_SIZE) -> tuple: | |
| _check_group_size(group_size) | |
| if out_features <= 0 or in_features <= 0: | |
| raise ValueError("Linear dimensions must be positive") | |
| return (out_features, math.ceil(in_features / group_size), group_size * 3 // 4) | |
| def pack_codes(codes: torch.Tensor) -> torch.Tensor: | |
| """Pack unsigned codes 0..63, four codes into three little-endian bytes.""" | |
| if codes.dtype != torch.uint8 or codes.shape[-1] % 4: | |
| raise ValueError("Codes must be uint8 with a final dimension divisible by four") | |
| if bool((codes > 63).any()): | |
| raise ValueError("Q6 code is outside 0..63") | |
| c = codes.reshape(*codes.shape[:-1], -1, 4) | |
| result = torch.empty((*c.shape[:-1], 3), dtype=torch.uint8, device=codes.device) | |
| result[..., 0] = c[..., 0] | ((c[..., 1] & 3) << 6) | |
| result[..., 1] = (c[..., 1] >> 2) | ((c[..., 2] & 15) << 4) | |
| result[..., 2] = (c[..., 2] >> 4) | (c[..., 3] << 2) | |
| return result.flatten(-2) | |
| def unpack_codes(packed: torch.Tensor) -> torch.Tensor: | |
| if packed.dtype != torch.uint8 or packed.shape[-1] % 3: | |
| raise ValueError("Packed codes must be uint8 with a final dimension divisible by three") | |
| b = packed.reshape(*packed.shape[:-1], -1, 3) | |
| result = torch.empty((*b.shape[:-1], 4), dtype=torch.uint8, device=packed.device) | |
| result[..., 0] = b[..., 0] & 63 | |
| result[..., 1] = (b[..., 0] >> 6) | ((b[..., 1] & 15) << 2) | |
| result[..., 2] = (b[..., 1] >> 4) | ((b[..., 2] & 3) << 4) | |
| result[..., 3] = b[..., 2] >> 2 | |
| return result.flatten(-2) | |
| def quantize_weight(weight: torch.Tensor, *, group_size: int = GROUP_SIZE) -> tuple: | |
| """Quantize one matrix or a bounded row slice; callers stream large inputs.""" | |
| _check_group_size(group_size) | |
| if weight.ndim != 2 or min(weight.shape) <= 0 or not weight.is_floating_point(): | |
| raise ValueError("Expected a nonempty floating point matrix") | |
| w = weight.detach().to(dtype=torch.float32) | |
| if not bool(torch.isfinite(w).all()): | |
| raise ValueError("Cannot quantize nonfinite weights") | |
| out_features, in_features = w.shape | |
| groups = math.ceil(in_features / group_size) | |
| if in_features % group_size: | |
| w = F.pad(w, (0, groups * group_size - in_features)) | |
| w = w.reshape(out_features, groups, group_size) | |
| maximum = w.abs().amax(dim=-1) | |
| scales = maximum / 31.0 | |
| scales = torch.where(maximum == 0, torch.ones_like(scales), scales) | |
| if not bool(torch.isfinite(scales).all()) or not bool((scales > 0).all()): | |
| raise ValueError("Scale overflow or underflow") | |
| codes = (torch.round(w / scales.unsqueeze(-1)).clamp_(-31, 31) + 32).to(torch.uint8) | |
| return pack_codes(codes), scales.contiguous() | |
| def dequantize_weight( | |
| qweight: torch.Tensor, scales: torch.Tensor, shape: tuple | list, | |
| *, dtype: torch.dtype = torch.bfloat16, row_chunk: int = 1024, | |
| ) -> torch.Tensor: | |
| """Expand exactly one matrix with bounded unpacking/FP32 scratch space.""" | |
| if len(shape) != 2 or tuple(qweight.shape) != packed_shape(*shape): | |
| raise ValueError("Packed weight shape does not match original matrix") | |
| if qweight.dtype != torch.uint8 or scales.dtype != torch.float32: | |
| raise ValueError("Expected uint8 packed weights and float32 scales") | |
| if tuple(scales.shape) != tuple(qweight.shape[:2]) or scales.device != qweight.device: | |
| raise ValueError("Scale shape/device mismatch") | |
| if dtype not in (torch.float32, torch.float16, torch.bfloat16) or row_chunk <= 0: | |
| raise ValueError("Unsupported compute dtype or row_chunk") | |
| out_features, in_features = shape | |
| result = torch.empty(tuple(shape), device=qweight.device, dtype=dtype) | |
| for start in range(0, out_features, row_chunk): | |
| stop = min(start + row_chunk, out_features) | |
| codes = unpack_codes(qweight[start:stop]).to(torch.float32).sub_(32) | |
| codes.mul_(scales[start:stop].unsqueeze(-1)) | |
| result[start:stop].copy_(codes.flatten(1)[:, :in_features]) | |
| return result | |
| class Q6Linear(nn.Module): | |
| """Frozen, inference-only Linear with no persistent dense weight cache. | |
| Packed weights may remain on CPU, or an offload controller may stage them. | |
| The temporary weight and bias always match the input's device and dtype. | |
| FP32 scales survive parent module dtype changes without a lossy round trip. | |
| """ | |
| def __init__( | |
| self, in_features: int, out_features: int, bias: bool = True, | |
| *, group_size: int = GROUP_SIZE, device="cpu", dtype=torch.bfloat16, | |
| row_chunk: int = 1024, | |
| ): | |
| super().__init__() | |
| self.in_features = int(in_features) | |
| self.out_features = int(out_features) | |
| self.group_size = group_size | |
| self.row_chunk = row_chunk | |
| shape = packed_shape(out_features, in_features, group_size) | |
| self.register_buffer("qweight", torch.empty(shape, dtype=torch.uint8, device=device)) | |
| self.register_buffer("scales", torch.empty(shape[:2], dtype=torch.float32, device=device)) | |
| self.bias = nn.Parameter(torch.empty(out_features, dtype=dtype, device=device), requires_grad=False) if bias else None | |
| def from_metadata(cls, in_features, out_features, bias=True, group_size=GROUP_SIZE, **kwargs): | |
| return cls(in_features, out_features, bias, group_size=group_size, **kwargs) | |
| def _apply(self, fn, recurse=True): | |
| # A parent .to(dtype=...) recursively invokes _apply, not our .to. | |
| # Move scales according to fn's device without ever casting their data. | |
| scales = self._buffers.pop("scales") | |
| try: | |
| result = super()._apply(fn, recurse=recurse) | |
| probe = fn(torch.empty(0, dtype=torch.float32, device=scales.device)) | |
| self._buffers["scales"] = scales.to(device=probe.device, dtype=torch.float32) | |
| return result | |
| except BaseException: | |
| self._buffers["scales"] = scales | |
| raise | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| if x.requires_grad and torch.is_grad_enabled(): | |
| raise RuntimeError("Q6Linear is inference-only; use torch.inference_mode()") | |
| if x.dtype not in (torch.bfloat16, torch.float16, torch.float32): | |
| raise ValueError("Q6Linear input must be bfloat16, float16 or float32") | |
| if x.shape[-1] != self.in_features: | |
| raise ValueError("Q6Linear input dimension mismatch") | |
| packed = self.qweight.to(device=x.device) | |
| scales = self.scales.to(device=x.device) | |
| weight = dequantize_weight( | |
| packed, scales, (self.out_features, self.in_features), | |
| dtype=x.dtype, row_chunk=self.row_chunk, | |
| ) | |
| bias = self.bias.to(device=x.device, dtype=x.dtype) if self.bias is not None else None | |
| with torch.autocast(device_type=x.device.type, enabled=False): | |
| return F.linear(x, weight, bias) | |
| def extra_repr(self): | |
| return f"in_features={self.in_features}, out_features={self.out_features}, group_size={self.group_size}, bias={self.bias is not None}" | |
| QuantLinear = Q6Linear | |