openjev-e4b / openjev /modeling.py
bambamdevs's picture
Publish OpenJEV E4B 1.0
03223d7
Raw History Blame Contribute Delete
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}")
@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)