File size: 4,350 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""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)}


@dataclass(frozen=True)
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')

    @property
    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())

    @classmethod
    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