UI-testjev / quantized.py
Shelter's picture
Introduce UI-TestJev with concise results and usage; remove development artifacts
961a722 verified
Raw History Blame Contribute Delete
11.7 kB
"""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