"""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) @torch.no_grad() 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() @torch.no_grad() 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 @classmethod 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