File size: 11,716 Bytes
961a722
4426514
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
961a722
4426514
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
961a722
4426514
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
"""Load frozen INT8 experts while preserving BF16 vision and decoder compute."""
from collections import defaultdict
import json
from pathlib import Path
import types

from accelerate import init_empty_weights
from safetensors import safe_open
import torch
from torch import nn
from transformers import DiffusionGemmaForBlockDiffusion
from transformers.activations import ACT2FN
from transformers.modeling_outputs import BaseModelOutputWithPooling

RECIPE = {
    "name": "frozen-experts-int8-row-v1",
    "storage": "symmetric INT8 per expert/output row; FP32 scales; round-to-nearest-even",
    "compute": "BF16; dequantize FP32 then cast BF16",
    "scope": "shared encoder/decoder fused experts only; other base weights BF16",
    "calibration": "none",
    "vision_microbatch": 1,
}


def quantize_rows(weight):
    w = weight.float()
    if not torch.isfinite(w).all():
        raise ValueError("Non-finite expert weight in checkpoint")
    scale = w.abs().amax(dim=-1, keepdim=True) / 127
    scale = torch.where(scale == 0, torch.ones_like(scale), scale)
    q = (w / scale).round().clamp(-127, 127).to(torch.int8)
    return q, scale


def dequantize(q, scale, dtype):
    return (q.float() * scale).to(dtype)


class Int8Experts(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.num_experts = config.num_experts
        self.act_fn = ACT2FN[config.hidden_activation]
        for name in ("gate_up_q", "gate_up_scale", "down_q", "down_scale"):
            self.register_buffer(name, None)

    def load_projection(self, name, source_slice, device):
        shape = tuple(source_slice.get_shape())
        if len(shape) != 3 or shape[0] != self.num_experts:
            raise ValueError(f"Invalid expert projection shape: {shape}")
        prefix = {"gate_up_proj": "gate_up", "down_proj": "down"}[name]
        q = torch.empty(shape, device=device, dtype=torch.int8)
        scale = torch.empty((*shape[:-1], 1), device=device, dtype=torch.float32)
        # CPU conversion makes the recipe independent of the GPU. Memory-mapped
        # safetensors are read in expert-sized slices, never a full BF16 bank on GPU.
        for i in range(shape[0]):
            qi, si = quantize_rows(source_slice[i])
            q[i].copy_(qi); scale[i].copy_(si)
        setattr(self, prefix+"_q", q)
        setattr(self, prefix+"_scale", scale)

    def _one(self, x, qg, sg, qd, sd):
        gate, up = nn.functional.linear(x, dequantize(qg, sg, x.dtype)).chunk(2, dim=-1)
        return nn.functional.linear(self.act_fn(gate) * up, dequantize(qd, sd, x.dtype))

    def forward(self, hidden_states, top_k_index, top_k_weights):
        result = torch.zeros_like(hidden_states)
        # Match HF's expert and (slot, token) accumulation order. Transfer only
        # the small list of active expert IDs, not weights or activations.
        with torch.no_grad():
            mask = nn.functional.one_hot(top_k_index, num_classes=self.num_experts).permute(2, 1, 0)
            active = (mask.sum(dim=(-1, -2)) > 0).nonzero().flatten().tolist()
        for i in active:
            slot, token = torch.where(mask[i])
            args = (hidden_states[token], self.gate_up_q[i], self.gate_up_scale[i], self.down_q[i], self.down_scale[i])
            value = self._one(*args)
            value = value * top_k_weights[token, slot, None]
            result.index_add_(0, token, value.to(result.dtype))
        return result


def _set_tensor(model, name, tensor, parameter=False):
    parent, attr = name.rsplit(".", 1)
    module = model.get_submodule(parent)
    setattr(module, attr, nn.Parameter(tensor, requires_grad=False) if parameter else tensor)


def tensor_storage_bytes(model):
    stores = {}
    for t in list(model.parameters()) + list(model.buffers()):
        if t.device.type == "meta":
            raise RuntimeError("Unloaded meta tensor remains")
        stores[(str(t.device), t.untyped_storage().data_ptr())] = t.untyped_storage().nbytes()
    return sum(stores.values())


def load_streamed_int8(directory, config, device="cuda:0", dtype=torch.bfloat16):
    """Load a local HF safetensors checkpoint without a full BF16 CPU/GPU copy."""
    directory = Path(directory)
    index = directory/"model.safetensors.index.json"
    if index.exists():
        weight_map = json.loads(index.read_text())["weight_map"]
    else:
        with safe_open(directory/"model.safetensors", framework="pt", device="cpu") as f:
            weight_map = {k: "model.safetensors" for k in f.keys()}
    # Parameters are meta; small nonpersistent rotary/embedding buffers initialize
    # normally. Avoid to_empty(), which would discard those buffer values.
    config._attn_implementation = "sdpa"
    with init_empty_weights(include_buffers=False):
        model = DiffusionGemmaForBlockDiffusion(config)
    model.tie_weights()
    aliases = defaultdict(list)
    expected_shapes = {}
    for name, p in model.named_parameters(remove_duplicate=False):
        aliases[id(p)].append(name)
        expected_shapes[name] = tuple(p.shape)
    tasks, consumed = defaultdict(list), set()
    banks = {}
    for i, (encoder, decoder) in enumerate(zip(model.model.encoder.language_model.layers, model.model.decoder.layers)):
        bank = Int8Experts(config.text_config)
        banks[i] = bank
        encoder.experts = decoder.experts = bank
    for names in aliases.values():
        candidates = [n for n in names if n in weight_map]
        if not candidates:
            raise RuntimeError(f"Missing checkpoint parameter: {names}")
        # The pinned checkpoint saves only one tensor per tied group. Refuse an
        # ambiguous conversion instead of silently discarding another base tensor.
        if len(candidates) != 1:
            raise RuntimeError(f"Checkpoint unexpectedly duplicates tied parameters: {candidates}")
        source = candidates[0]
        consumed.add(source)
        if ".experts." in source:
            parts = source.split(".")
            layer = int(parts[parts.index("layers")+1])
            tasks[weight_map[source]].append((source, "expert", (banks[layer], parts[-1])))
        else:
            tasks[weight_map[source]].append((source, "parameter", names))
    # Include all persistent buffers from the original checkpoint, e.g. layer
    # scalars and vision standardization; initialize no learned values by guessing.
    persistent = set(model.state_dict())
    for name, b in model.named_buffers(remove_duplicate=False):
        if name in persistent:
            expected_shapes[name] = tuple(b.shape)
            if name not in weight_map:
                raise RuntimeError(f"Missing checkpoint buffer: {name}")
            tasks[weight_map[name]].append((name, "buffer", [name]))
            consumed.add(name)
        else:
            _set_tensor(model, name, b.to(device))
    if consumed != set(weight_map):
        raise RuntimeError(f"Unexpected checkpoint keys: {sorted(set(weight_map)-consumed)[:10]}")
    for shard, entries in sorted(tasks.items()):
        if Path(shard).name != shard or not shard.endswith(".safetensors"):
            raise ValueError("Unsafe checkpoint shard path")
        print(f"Streaming {shard}: {len(entries)} tensors; quantizing frozen experts on CPU", flush=True)
        with safe_open(directory/shard, framework="pt", device="cpu") as f:
            for source, kind, dest in entries:
                if tuple(f.get_slice(source).get_shape()) != expected_shapes[source]:
                    raise ValueError(f"Checkpoint tensor shape mismatch: {source}")
                if kind == "expert":
                    dest[0].load_projection(dest[1], f.get_slice(source), device)
                else:
                    value = f.get_tensor(source)
                    # HF keeps vision standardization buffers in their checkpoint
                    # precision; float32 cancellation here affects visual evidence.
                    target_dtype = value.dtype if kind == "buffer" else dtype
                    value = value.to(device=device, dtype=target_dtype)
                    if kind == "parameter":
                        p = nn.Parameter(value, requires_grad=False)
                        for name in dest:
                            parent, attr = name.rsplit(".", 1)
                            setattr(model.get_submodule(parent), attr, p)
                    else:
                        _set_tensor(model, dest[0], value)
    model.requires_grad_(False)
    if any(bank.gate_up_q is None or bank.down_q is None for bank in banks.values()):
        raise RuntimeError("Unloaded expert bank")
    model.jev_storage_bytes = tensor_storage_bytes(model)
    print(f"Loaded tensor storage: {model.jev_storage_bytes/2**30:.2f}GiB (excludes runtime activations).", flush=True)
    return model.eval()


def enable_sequential_vision(model):
    """Same pixel patches, roles and order; at most one image in vision attention."""
    encoder = model.model.encoder
    if getattr(encoder, "_jev_sequential_vision", False):
        return
    def image_features(self, pixel_values, image_position_ids=None, **kwargs):
        if image_position_ids is None or pixel_values.shape[0] != image_position_ids.shape[0]:
            raise ValueError("Image batch/position association mismatch")
        features, projected = [], []
        cache = getattr(self, "_jev_feature_cache", None)
        keys = getattr(self, "_jev_image_keys", None)
        if keys is not None and len(keys) != pixel_values.shape[0]:
            raise ValueError("Vision cache image association mismatch")
        if cache is not None and (self.vision_tower.training or self.embed_vision.training or
                any(p.requires_grad for m in (self.vision_tower, self.embed_vision) for p in m.parameters())):
            raise RuntimeError("Feature reuse requires frozen, evaluation-mode vision and projection")
        for i in range(pixel_values.shape[0]):
            hit = cache.get(keys[i], pixel_values.device) if cache is not None and keys is not None else None
            if hit is not None:
                h, projection = hit
            else:
                pos = image_position_ids[i:i+1]
                valid = (pos[0] >= 0).all(dim=-1)
                n = int(valid.sum())
                if n == 0 or n % 9 or not bool(valid[:n].all()) or bool(valid[n:].any()):
                    raise ValueError("Expected contiguous real 3x3 vision patches followed by padding")
                # Remove only processor padding, never pixels belonging to the image.
                h = self.vision_tower(pixel_values=pixel_values[i:i+1, :n],
                    pixel_position_ids=pos[:, :n], **kwargs).last_hidden_state
                projection = self.embed_vision(inputs_embeds=h)
                if cache is not None and keys is not None:
                    cache.put(keys[i], h, projection)
            expected = int((image_position_ids[i] >= 0).all(dim=-1).sum())//9
            if h.shape[0] != expected or projection.shape[0] != expected:
                raise ValueError("Cached feature/image token association mismatch")
            features.append(h)
            projected.append(projection)
        return BaseModelOutputWithPooling(last_hidden_state=torch.cat(features, dim=0),
                                         pooler_output=torch.cat(projected, dim=0))
    encoder.get_image_features = types.MethodType(image_features, encoder)
    encoder._jev_sequential_vision = True