File size: 5,257 Bytes
f2c0505
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
from __future__ import annotations

from collections import OrderedDict
import weakref
import torch
from torch import nn

from .triton_w4a4 import PackedActivation, a4_pack_triton, w4a4_linear_triton


class ActivationPackCache:
    """Small per-device cache to reuse identical A4 packing across Q/K/V linears."""
    def __init__(self, capacity: int = 4):
        self.capacity = int(capacity)
        self._items: OrderedDict[tuple, tuple[weakref.ReferenceType, PackedActivation]] = OrderedDict()
        self.hits = 0
        self.misses = 0

    @staticmethod
    def _key(x: torch.Tensor, d: int) -> tuple:
        return (
            x.device.type,
            -1 if x.device.index is None else int(x.device.index),
            int(x.data_ptr()),
            tuple(x.shape),
            tuple(x.stride()),
            int(getattr(x, '_version', 0)),
            int(d),
        )

    def get_or_pack(self, x, d, rot):
        key = self._key(x, d)
        hit = self._items.get(key)
        if hit is not None and hit[0]() is x:
            self.hits += 1
            self._items.move_to_end(key)
            return hit[1]
        self.misses += 1
        packed = a4_pack_triton(
            x,
            rot['perm'],
            rot['signs'],
            int(rot['block_size']),
            rot['codebook'],
        )
        self._items[key] = (weakref.ref(x), packed)
        self._items.move_to_end(key)
        while len(self._items) > self.capacity:
            self._items.popitem(last=False)
        return packed

    def clear(self):
        self._items.clear()

    def stats(self):
        return {"hits": int(self.hits), "misses": int(self.misses), "entries": len(self._items)}


class OrbitQuantW4A4Engine:
    def __init__(self, rotation_bank, cache_capacity: int = 4):
        self.bank = rotation_bank
        self.activation_cache = ActivationPackCache(cache_capacity)
        self._device_rot = {}
        self.linear_calls = 0

    def rotation(self, d: int, device: torch.device):
        idx = -1 if device.index is None else int(device.index)
        key = (int(d), device.type, idx)
        if key not in self._device_rot:
            src = self.bank.tensors[int(d)]
            self._device_rot[key] = {
                'perm': src['perm'].to(device=device, dtype=torch.int32),
                'signs': src['signs'].to(device=device, dtype=torch.int8),
                'codebook': src['codebook'].to(device=device, dtype=torch.float32),
                'block_size': int(src['block_size'].item()),
            }
        return self._device_rot[key]

    def linear(self, x, module: 'OrbitQuantPackedLinear'):
        self.linear_calls += 1
        d = module.in_features
        rot = self.rotation(d, x.device)
        a = self.activation_cache.get_or_pack(x, d, rot)
        out = w4a4_linear_triton(
            a,
            module.packed_weight,
            module.row_scale,
            rot['codebook'],
            module.bias,
            out_dtype=x.dtype if x.dtype in (torch.bfloat16, torch.float16) else torch.bfloat16,
        )
        return out.reshape(*x.shape[:-1], module.out_features)


class OrbitQuantPackedLinear(nn.Module):
    """Packed nonuniform OrbitQuant W4 x online OrbitQuant A4 linear.

    There is deliberately no dense `weight` Parameter. Target weight residency is
    uint4 packed storage + one BF16 row norm from the artifact (promoted to FP32 for the kernel multiply).
    """
    def __init__(self, in_features: int, out_features: int, bias: bool, engine: OrbitQuantW4A4Engine):
        super().__init__()
        self.in_features = int(in_features)
        self.out_features = int(out_features)
        self.engine = engine
        # Empty/meta bias is materialized by the streaming checkpoint loader.
        if bias:
            self.bias = nn.Parameter(torch.empty(out_features, device='meta', dtype=torch.bfloat16), requires_grad=False)
        else:
            self.register_parameter('bias', None)
        self.register_buffer('packed_weight', None, persistent=False)
        self.register_buffer('row_scale', None, persistent=False)
        self._orbitquant_w4a4 = True
        self._orbitquant_call_count = 0

    def set_packed(self, packed: torch.Tensor, scale: torch.Tensor):
        expected = (self.in_features // 2, self.out_features)
        if packed.dtype != torch.uint8 or tuple(packed.shape) != expected:
            raise ValueError(
                f'packed shape/dtype mismatch: got {tuple(packed.shape)} {packed.dtype}, '
                f'expected GEMM-native [K/2,N]={expected} uint8'
            )
        if scale.shape != (self.out_features,):
            raise ValueError('row scale shape mismatch')
        self.packed_weight = packed.contiguous()
        self.row_scale = scale.float().contiguous()
        return self

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if self.packed_weight is None or self.row_scale is None:
            raise RuntimeError('packed OrbitQuant weight has not been loaded')
        self._orbitquant_call_count += 1
        return self.engine.linear(x, self)

    def extra_repr(self) -> str:
        return f'in_features={self.in_features}, out_features={self.out_features}, packed=W4, activation=A4'