from __future__ import annotations import gc import json from pathlib import Path from typing import Any import torch from safetensors.torch import load_file from peft import PeftModel from transformers import AutoModelForMultimodalLM, AutoTokenizer, BitsAndBytesConfig from .decision_head import PointerDecisionHead, SpanDecisionHead, HybridDecisionHead from .formatting import pack_question, collate_packed, encode_piece def _find_text_model(module): seen = set() queue = [module] while queue: obj = queue.pop(0) if obj is None or id(obj) in seen: continue seen.add(id(obj)) if hasattr(obj, "layers") and hasattr(obj, "embed_tokens"): return obj for attr in ("model", "language_model", "base_model"): child = getattr(obj, attr, None) if child is not None and child is not obj: queue.append(child) if hasattr(obj, "get_base_model"): try: child = obj.get_base_model() except Exception: child = None if child is not None and child is not obj: queue.append(child) raise RuntimeError("Could not locate Gemma 4 text transformer") def _extract_text_backbone(full_model, model_id: str): try: text = full_model.model.language_model except AttributeError as exc: raise RuntimeError( f"Expected model.language_model in {model_id}, got {full_model.__class__.__name__}" ) from exc expected = int(getattr(text.config, "num_hidden_layers", 0)) actual = len(getattr(text, "layers", [])) if expected <= 0 or actual != expected: raise RuntimeError(f"Backbone layer check failed: actual={actual} expected={expected}") full_model.model.language_model = None del full_model gc.collect() text.config.use_cache = False return text def _apply_backbone_delta(backbone, delta, load_mode: str): if load_mode == "nf4": # Quantized weights cannot receive the accepted dense base tensors via # copy_. Restore only the adapted linear modules in dense FP16. modules = dict(backbone.named_modules()) for name in delta: if not name.endswith(".weight"): continue module_path = name[:-len(".weight")] module = modules.get(module_path) if module is None: continue weight = getattr(module, "weight", None) if module.__class__.__name__ != "Linear4bit" and getattr(weight, "__class__", type(None)).__name__ != "Params4bit": continue parent_path, attr = module_path.rsplit(".", 1) replacement = torch.nn.Linear( module.in_features, module.out_features, bias=module.bias is not None, device=weight.device, dtype=torch.float16, ) if module.bias is not None: with torch.no_grad(): replacement.bias.copy_(module.bias.to(device=weight.device, dtype=torch.float16)) replacement.requires_grad_(False) setattr(backbone.get_submodule(parent_path), attr, replacement) named = dict(backbone.named_parameters()) missing = [name for name in delta if name not in named] if missing: raise RuntimeError(f"Backbone delta incompatible with base/adapter; missing={missing[:8]}") bad_shapes = [ name for name, value in delta.items() if tuple(named[name].shape) != tuple(value.shape) ] if bad_shapes: raise RuntimeError(f"Backbone delta shape mismatch: {bad_shapes[:8]}") with torch.no_grad(): for name, value in delta.items(): target = named[name] target.copy_(value.to(device=target.device, dtype=target.dtype)) class _DecisionModel(torch.nn.Module): def __init__(self, backbone, hidden_size: int, head_dim: int, head_type: str): super().__init__() self.backbone = backbone self.head_type = head_type if head_type == "pointer": self.head = PointerDecisionHead(hidden_size, head_dim=head_dim, normalize=True) elif head_type == "span": self.head = SpanDecisionHead(hidden_size, head_dim=head_dim) elif head_type == "hybrid": self.head = HybridDecisionHead(hidden_size, head_dim=head_dim, pointer_dim=head_dim, normalize=True) else: raise ValueError(f"Unsupported head_type={head_type}") @staticmethod def _option_means(h, batch): starts = batch["option_starts"] ends = batch["option_ends"] bsz, nopt = batch["option_mask"].shape out = torch.zeros((bsz, nopt, h.shape[-1]), device=h.device, dtype=h.dtype) for i in range(bsz): for j in range(nopt): if not bool(batch["option_mask"][i, j]): continue a, b = int(starts[i, j]), int(ends[i, j]) out[i, j] = h[i, a:max(a + 1, b)].mean(0) return out def forward(self, batch): core = _find_text_model(self.backbone) out = core( input_ids=batch["input_ids"], attention_mask=batch["attention_mask"], use_cache=False, return_dict=True, ) return self.score(out.last_hidden_state, batch) def score(self, h, batch): b = torch.arange(h.shape[0], device=h.device) decide_h = h[b, batch["decide_positions"]] option_last_h = h[b[:, None], batch["option_positions"]] if self.head_type == "pointer": return self.head(decide_h, option_last_h, batch["option_mask"]) option_mean_h = self._option_means(h, batch) if self.head_type == "span": return self.head(decide_h, option_mean_h, batch["option_mask"]) return self.head(decide_h, option_last_h, option_mean_h, batch["option_mask"]) def _temperature(calibration: Any, primitive: str = "choice") -> float: if isinstance(calibration, (float, int)): return float(calibration) if not isinstance(calibration, dict): return 1.0 groups = calibration.get("groups") or {} for k, v in groups.items(): if str(k).lower() == primitive.lower(): if isinstance(v, dict): for kk in ("temperature", "T", "t"): if kk in v: return float(v[kk]) if isinstance(v, (float, int)): return float(v) for k in ("global_temperature", "temperature", "Tglobal"): if k in calibration: return float(calibration[k]) return 1.0 def _insert_media(tok, packed: dict, media_ids: list[int]) -> dict: """Place media soft tokens right after "State:\\n" and shift every position.""" prefix = ([tok.bos_token_id] if tok.bos_token_id is not None else []) + encode_piece(tok, "State:\n") ids = packed["input_ids"] if ids[:len(prefix)] != prefix: raise RuntimeError("Unexpected packed prefix; cannot insert media") k, m = len(prefix), len(media_ids) return { **packed, "input_ids": ids[:k] + media_ids + ids[k:], "option_positions": [p + m for p in packed["option_positions"]], "option_spans": [(a + m, b + m) for a, b in packed["option_spans"]], "decide_position": packed["decide_position"] + m, } class OpenJEVR7: def __init__(self, model, tokenizer, config, calibration, device, mm_model=None, media=None): self.model = model self.tokenizer = tokenizer self.config = config self.calibration = calibration self.device = device # Set only when loaded with multimodal=True (experimental). self.mm_model = mm_model self.media = media @classmethod def from_pretrained( cls, repo_dir: str | Path, device: str = "cuda", dtype=torch.bfloat16, load_mode: str = "bf16", multimodal: bool = False, ): """Load the release. multimodal=True keeps Gemma's vision and audio encoders so choice() accepts image= and audio= (experimental; trained on text only).""" if load_mode not in {"bf16", "nf4"}: raise ValueError("load_mode must be 'bf16' or 'nf4'") if load_mode == "nf4" and not str(device).startswith("cuda"): raise ValueError("NF4 loading requires a CUDA device") repo_dir = Path(repo_dir) model_dir = repo_dir / "model" cfg = json.loads((model_dir / "openjev_config.json").read_text(encoding="utf-8")) model_id = cfg["model_id"] revision = cfg["base_revision"] tok = AutoTokenizer.from_pretrained(model_id, revision=revision, use_fast=True) if tok.pad_token_id is None: tok.pad_token = tok.eos_token load_kwargs = {"revision": revision, "low_cpu_mem_usage": True, "attn_implementation": "sdpa"} if load_mode == "nf4": load_kwargs.update({ "dtype": torch.float16, "device_map": {"": device}, "quantization_config": BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype=torch.float16, ), }) if multimodal: from .multimodal import quantization_skip_modules # Gemma's audio encoder cannot run with 4-bit weights; keep both encoders dense. load_kwargs["quantization_config"].llm_int8_skip_modules = quantization_skip_modules(model_id, revision) else: load_kwargs["dtype"] = dtype full = AutoModelForMultimodalLM.from_pretrained(model_id, **load_kwargs) mm_model = media = None if multimodal: from .multimodal import MediaEncoder mm_model = full.model text = mm_model.language_model expected = int(getattr(text.config, "num_hidden_layers", 0)) if expected <= 0 or len(getattr(text, "layers", [])) != expected: raise RuntimeError(f"Backbone layer check failed for {model_id}") text.config.use_cache = False # Wrap the text model in place so the vision/audio path runs through the adapter. backbone = PeftModel.from_pretrained(text, model_dir / "adapter", is_trainable=False) mm_model.language_model = backbone media = MediaEncoder(model_id, revision, tok) else: backbone = _extract_text_backbone(full, model_id) del full gc.collect() backbone = PeftModel.from_pretrained(backbone, model_dir / "adapter", is_trainable=False) delta = load_file(str(model_dir / "backbone_delta.safetensors"), device="cpu") _apply_backbone_delta(backbone, delta, load_mode) del delta core = _find_text_model(backbone) hidden = int(core.config.hidden_size) head_type = str(cfg.get("head_type", "pointer")) model = _DecisionModel(backbone, hidden, int(cfg["head_dim"]), head_type=head_type) head_state = load_file(str(model_dir / "decision_head.safetensors"), device="cpu") model.head.load_state_dict(head_state) if load_mode == "nf4": model.head.to(device=device, dtype=torch.float32) else: if mm_model is not None: mm_model.to(device) model.to(device) model.eval() if mm_model is not None: mm_model.eval() cal_path = model_dir / "calibration.json" calibration = json.loads(cal_path.read_text(encoding="utf-8")) if cal_path.exists() else 1.0 return cls(model, tok, cfg, calibration, device, mm_model=mm_model, media=media) @torch.inference_mode() def choice( self, state, instruction: str, options: list[str], max_length: int | None = None, image=None, audio=None, sampling_rate: int | None = None, ): """Score a closed set of options. image (path or PIL image) and audio (.wav path, or an array with sampling_rate) are experimental and need multimodal=True.""" if len(options) < 2: raise ValueError("choice requires at least two options") has_media = image is not None or audio is not None if has_media and self.media is None: raise ValueError("image= and audio= need OpenJEV.from_pretrained(..., multimodal=True)") max_length = int(max_length or self.config.get("max_length", 8192)) q = { "instruction": instruction, "options": [{"text": str(x)} for x in options], } media_ids, media_inputs = self.media.encode(image, audio, sampling_rate) if has_media else ([], {}) packed = pack_question(self.tokenizer, state, q, max_length=max_length - len(media_ids)) if packed is None: raise ValueError("Input could not be packed within max_length") if has_media: packed = _insert_media(self.tokenizer, packed, media_ids) batch = collate_packed(self.tokenizer, [packed]) batch = {k: v.to(self.device) for k, v in batch.items()} if has_media: logits = self._media_logits(batch, media_inputs)[0, :len(options)] else: logits = self.model(batch)[0, :len(options)] t = _temperature(self.calibration, "choice") probs = torch.softmax(logits.float() / max(t, 1e-6), dim=-1).cpu().tolist() idx = int(max(range(len(probs)), key=probs.__getitem__)) return { "type": "choice", "probabilities": probs, "selected_index": idx, "selected_option": options[idx], "temperature": t, "was_truncated": bool(packed.get("was_truncated", False)), } def _media_logits(self, batch, media_inputs): ids = batch["input_ids"] mm_types = (ids == self.media.image_token_id).long() + 3 * (ids == self.media.audio_token_id).long() extra = {} for key, value in media_inputs.items(): value = value.to(self.device) if value.is_floating_point(): # Match the encoder's float dtype (NF4 towers keep float16 norms/embeddings). tower = self.mm_model.vision_tower if key == "pixel_values" else self.mm_model.audio_tower value = value.to(next(p.dtype for p in tower.parameters() if p.is_floating_point())) extra[key] = value out = self.mm_model( input_ids=ids, attention_mask=batch["attention_mask"], mm_token_type_ids=mm_types, use_cache=False, return_dict=True, **extra, ) return self.model.score(out.last_hidden_state, batch)