yongyizang's picture
TinyMOSS-Diarize: 2.911-bit packed weights, runtime, and model card
7ccb33d verified
Raw
History Blame Contribute Delete
8.9 kB
"""True int8/int4/int3 packing for canonical RTN fake-quant weights."""
from __future__ import annotations
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from .rtn_quant import _validate_rtn_args
_DTYPE_NAMES = {
torch.float16: "float16",
torch.bfloat16: "bfloat16",
torch.float32: "float32",
torch.float64: "float64",
}
_NAME_DTYPES = {name: dtype for dtype, name in _DTYPE_NAMES.items()}
@dataclass(frozen=True)
class PackedRTN:
"""Integer payload, fp16 scales, and the metadata needed for decoding."""
data: torch.Tensor
scales: torch.Tensor
shape: tuple[int, int]
bits: int
granularity: str
group_size: int
original_dtype: str
@property
def num_weights(self) -> int:
return self.shape[0] * self.shape[1]
@property
def effective_bits(self) -> float:
"""Actual bits/weight including fp16 scales and payload byte padding."""
return (self.data.numel() * 8 + self.scales.numel() * 16) / self.num_weights
def to_bytes(self) -> bytes:
if self.data.dtype == torch.int8:
return self.data.cpu().numpy().tobytes()
return bytes(self.data.cpu().tolist())
def _canonical_codes_and_scales(
weight: torch.Tensor,
bits: int,
granularity: str,
group_size: int,
) -> tuple[torch.Tensor, torch.Tensor]:
_validate_rtn_args(weight, bits, granularity, group_size)
if weight.dtype not in _DTYPE_NAMES:
raise ValueError(f"unsupported fake-quant dtype: {weight.dtype}")
cpu_weight = weight.detach().cpu().contiguous()
rows, columns = cpu_weight.shape
actual_group_size = columns if granularity == "per_channel" else group_size
number_of_groups = (columns + actual_group_size - 1) // actual_group_size
padded_columns = number_of_groups * actual_group_size
padded = F.pad(cpu_weight, (0, padded_columns - columns))
groups = padded.reshape(rows, number_of_groups, actual_group_size)
qmax = 2 ** (bits - 1) - 1
scales = (groups.float().abs().amax(dim=-1) / qmax).to(torch.float16)
safe_scales = torch.where(scales == 0, torch.ones_like(scales), scales).float()
codes = torch.round(groups.float() / safe_scales[..., None]).clamp(-qmax, qmax)
reconstructed = (codes * scales.float()[..., None]).reshape(rows, -1)[:, :columns]
reconstructed = reconstructed.to(cpu_weight.dtype)
if not torch.equal(reconstructed, cpu_weight):
raise ValueError(
"w_fakequant is not exactly representable by RTN integer codes and fp16 "
"scales; fake-quantize an fp16 tensor before packing"
)
return codes.to(torch.int8).reshape(-1), scales.contiguous()
def _pack_w3(codes: torch.Tensor) -> torch.Tensor:
"""Pack signed W3 codes little-endian, eight codes per three bytes."""
logical_codes = codes.numel()
if logical_codes % 8:
codes = F.pad(codes, (0, 8 - logical_codes % 8))
values = (codes.to(torch.int16) & 0x7).reshape(-1, 8)
payload = torch.empty((values.shape[0], 3), dtype=torch.int16)
payload[:, 0] = values[:, 0] | (values[:, 1] << 3) | (values[:, 2] << 6)
payload[:, 1] = (
(values[:, 2] >> 2)
| (values[:, 3] << 1)
| (values[:, 4] << 4)
| (values[:, 5] << 7)
)
payload[:, 2] = (values[:, 5] >> 1) | (values[:, 6] << 2) | (values[:, 7] << 5)
payload_bytes = (logical_codes * 3 + 7) // 8
return payload.reshape(-1)[:payload_bytes].to(torch.uint8).contiguous()
def _unpack_w3(payload: torch.Tensor, logical_codes: int) -> torch.Tensor:
"""Unpack the little-endian W3 stream to signed two's-complement codes."""
padding = (-payload.numel()) % 3
if padding:
payload = F.pad(payload, (0, padding))
packed = payload.to(torch.int16).reshape(-1, 3)
values = torch.empty((packed.shape[0], 8), dtype=torch.int16)
values[:, 0] = packed[:, 0] & 0x7
values[:, 1] = (packed[:, 0] >> 3) & 0x7
values[:, 2] = ((packed[:, 0] >> 6) | (packed[:, 1] << 2)) & 0x7
values[:, 3] = (packed[:, 1] >> 1) & 0x7
values[:, 4] = (packed[:, 1] >> 4) & 0x7
values[:, 5] = ((packed[:, 1] >> 7) | (packed[:, 2] << 1)) & 0x7
values[:, 6] = (packed[:, 2] >> 2) & 0x7
values[:, 7] = (packed[:, 2] >> 5) & 0x7
values = values.reshape(-1)[:logical_codes]
return torch.where(values >= 4, values - 8, values)
def pack_rtn(
w_fakequant: torch.Tensor,
bits: int = 8,
granularity: str = "per_channel",
group_size: int = 128,
) -> PackedRTN:
"""Pack canonical RTN values into int8, int4, or dense int3 payloads."""
codes, scales = _canonical_codes_and_scales(
w_fakequant, bits, granularity, group_size
)
rows, columns = w_fakequant.shape
actual_group_size = columns if granularity == "per_channel" else group_size
number_of_groups = (columns + actual_group_size - 1) // actual_group_size
padded_columns = number_of_groups * actual_group_size
logical_codes = codes.reshape(rows, padded_columns)[:, :columns].reshape(-1)
if bits == 8:
data = logical_codes.contiguous()
elif bits == 4:
if logical_codes.numel() % 2:
logical_codes = F.pad(logical_codes, (0, 1))
nibbles = logical_codes.to(torch.int16) & 0xF
data = (nibbles[0::2] | (nibbles[1::2] << 4)).to(torch.uint8)
else:
data = _pack_w3(logical_codes)
return PackedRTN(
data=data,
scales=scales,
shape=(w_fakequant.shape[0], w_fakequant.shape[1]),
bits=bits,
granularity=granularity,
group_size=group_size,
original_dtype=_DTYPE_NAMES[w_fakequant.dtype],
)
def unpack_rtn(packed: PackedRTN) -> torch.Tensor:
"""Decode a packed RTN tensor exactly to its canonical fake-quant values."""
rows, columns = packed.shape
if packed.bits not in {3, 4, 8}:
raise ValueError("packed bits must be 3, 4, or 8")
if packed.granularity not in {"per_channel", "per_group"}:
raise ValueError("invalid packed granularity")
if packed.group_size <= 0 or rows <= 0 or columns <= 0:
raise ValueError("invalid packed shape or group size")
if packed.scales.dtype != torch.float16 or packed.scales.ndim != 2:
raise ValueError("packed scales must be a two-dimensional float16 tensor")
if packed.original_dtype not in _NAME_DTYPES:
raise ValueError(f"unsupported original dtype metadata: {packed.original_dtype}")
actual_group_size = columns if packed.granularity == "per_channel" else packed.group_size
number_of_groups = (columns + actual_group_size - 1) // actual_group_size
expected_scale_shape = (rows, number_of_groups)
if tuple(packed.scales.shape) != expected_scale_shape:
raise ValueError("scale shape does not match packed metadata")
logical_weights = rows * columns
if packed.bits == 8:
if packed.data.dtype != torch.int8 or packed.data.ndim != 1:
raise ValueError("W8 packed data must be a one-dimensional int8 tensor")
if packed.data.numel() != logical_weights:
raise ValueError("W8 payload size does not match packed shape")
logical_codes = packed.data.cpu()
elif packed.bits == 4:
if packed.data.dtype != torch.uint8 or packed.data.ndim != 1:
raise ValueError("W4 packed data must be a one-dimensional uint8 tensor")
if packed.data.numel() != (logical_weights + 1) // 2:
raise ValueError("W4 payload size does not match packed shape")
payload = packed.data.cpu()
nibbles = torch.empty(payload.numel() * 2, dtype=torch.int16)
nibbles[0::2] = payload.to(torch.int16) & 0xF
nibbles[1::2] = payload.to(torch.int16) >> 4
logical_codes = torch.where(nibbles >= 8, nibbles - 16, nibbles)[:logical_weights]
else:
if packed.data.dtype != torch.uint8 or packed.data.ndim != 1:
raise ValueError("W3 packed data must be a one-dimensional uint8 tensor")
if packed.data.numel() != (logical_weights * 3 + 7) // 8:
raise ValueError("W3 payload size does not match packed shape")
logical_codes = _unpack_w3(packed.data.cpu(), logical_weights)
# Packing omits row-tail padding codes; restore it before applying group scales.
padded_columns = number_of_groups * actual_group_size
codes_by_row = logical_codes.reshape(rows, columns)
codes_by_row = F.pad(codes_by_row, (0, padded_columns - columns))
groups = codes_by_row.reshape(rows, number_of_groups, actual_group_size).float()
decoded = groups * packed.scales.cpu().float()[..., None]
return decoded.reshape(rows, padded_columns)[:, :columns].to(
_NAME_DTYPES[packed.original_dtype]
)
# Module-local conventional names are convenient without shadowing STQ1 exports
# from quantlib's package root.
pack = pack_rtn
unpack = unpack_rtn