"""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