Image-Text-to-Text
PEFT
Safetensors
English
lora
diffusion-gemma
ui-testing
visual-regression
accessibility
custom-inference
conversational
Instructions to use Shelter/UI-testjev with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Shelter/UI-testjev with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
Download quantized.py from Shelter/UI-testjev: direct link, hf CLI and curl.
- Browser
- Download file 11.7 kB
-
https://huggingface.co/Shelter/UI-testjev/resolve/main/quantized.py
- Command line
-
hf download hf://Shelter/UI-testjev/quantized.py
-
curl -L -o quantized.py https://huggingface.co/Shelter/UI-testjev/resolve/main/quantized.py
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 | |