atomtanstudio's picture
Add model card, Q6 manifest, loader package and conversion receipts
18e823d verified
Raw History Blame Contribute Delete
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)
@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