Zero-Shot Classification
Safetensors
PEFT
English
openjev
classification
decision-model
listwise
gemma4
research
Instructions to use bambamdevs/openjev-e4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use bambamdevs/openjev-e4b with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
Download openjev/modeling.py from bambamdevs/openjev-e4b: direct link, hf CLI and curl.
- Browser
- Download file 15 kB
-
https://huggingface.co/bambamdevs/openjev-e4b/resolve/main/openjev/modeling.py
- Command line
-
hf download hf://bambamdevs/openjev-e4b/openjev/modeling.py
-
curl -L -o modeling.py https://huggingface.co/bambamdevs/openjev-e4b/resolve/main/openjev/modeling.py
15 kB
| 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}") | |
| 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 | |
| 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) | |
| 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) | |