Download code/mmjev.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 44.1 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/mmjev.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/mmjev.py
-
curl -L -o mmjev.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/mmjev.py
44.1 kB
| """MM-Jev: a Jev-style typed decision model (noul / choice / score) on Gemma 3n E4B with multimodal state. | |
| State = segments (text, image, audio, video = frames + optional soundtrack). Questions carry text only. Nothing is | |
| generated: every option is scored and the distribution is the model's own softmax over the options. | |
| Architecture (see docs/architecture.md) | |
| --------------------------------------- | |
| 1. Option tree in one pass. [state] -> [question_j] -> [option_j1] [option_j2] ... packed in one sequence with a tree | |
| attention mask; every branch sees only its ancestors, position ids restart at the end of the question. The state is | |
| encoded once for all questions and options, and option order cannot matter (structural, like openjev). | |
| 2. KV-share query truncation (exact). Gemma 3n layers 20-34 compute no K/V of their own (they read layers 18/19), so a | |
| token's state in those layers only feeds its own output: they run on the scoring tokens only (1 per option). | |
| 3. Modality exit + latent memory. After layer k all media tokens (image / audio / video soft tokens) are dropped from | |
| the sequence. M learned latent tokens appended to every media segment, and the text tokens after it, have attended to | |
| the media in layers < k and carry what is needed forward (DyVTE / LLaVA-Mini / VoCo-LLaMA; 2512.07580 shows deep-layer | |
| visual tokens are no better than random). | |
| 4. Elastic depth. A second decision head reads the hidden state after layer 19 (the last layer with its own K/V): | |
| exiting there skips 15 of 35 layers of weights (miniReranker-style mid-depth exit). | |
| 5. Visual tokens. Input stays 768x768 (>= 720p). The 16x16 MobileNet-V5 grid is average-pooled before projection | |
| (2x2 -> 64 tokens per image, 4x4 -> 16 per video frame); near-duplicate video frames are dropped. | |
| 6. Head. score = w . h(last token of branch), w initialised to E[Yes] - E[No] (zero-shot "is this answer correct?" | |
| log-odds). Loss: log score + ranked probability score for ordinal questions; per-type temperature post hoc. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass, replace | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| QTYPES = {"noul": 0, "choice": 1, "score": 2} | |
| NOUL_DEFAULT = {"no": "no, the statement does not hold", "yes": "yes, the statement holds"} | |
| SYSTEM = ("You are a decision model. Read the state and the question, then judge whether the proposed answer " | |
| "is correct. Reply Yes or No.") | |
| FIRST_SHARED = 20 # Gemma 3n E4B: 35 layers, the last 15 share K/V | |
| # ------------------------------------------------------------------------------------------ inputs | |
| class Seg: | |
| kind: str # text | image | audio | video | |
| data: object = None # str | PIL.Image | np.ndarray 16 kHz | list[PIL] -- or cached tower features (Tensor) | |
| audio: object = None # optional soundtrack of a video (np.ndarray or cached Tensor) | |
| fps: float = 1.0 # frame rate of `data` for a video | |
| def k_bucket(k: int) -> str: | |
| return "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11-30" if k <= 30 else "31+" | |
| def options_of(q: dict) -> tuple[list[str], list[str]]: | |
| """(labels, branch texts) in label-index order. noul is always [no, yes] so p[1] is the noul probability.""" | |
| t, crit = q["type"], q.get("criteria") | |
| if t == "noul": | |
| crit = {**NOUL_DEFAULT, **(crit or {})} | |
| return ["no", "yes"], [f"no — {crit['no']}", f"yes — {crit['yes']}"] | |
| if t == "score": | |
| crit = list(crit) | |
| return [str(i) for i in range(len(crit))], [f"level {i} of {len(crit) - 1} — {c}" for i, c in enumerate(crit)] | |
| if isinstance(crit, dict): | |
| return list(crit), [f"{k} — {v}" if v else k for k, v in crit.items()] | |
| labels = list(crit if crit is not None else q["options"]) | |
| return labels, labels | |
| class FastConfig: | |
| media_exit: int | None = None # drop media tokens after this many layers (None = keep) | |
| exit_layer: int = 35 # 35 = full depth, 20 = mid-depth head | |
| truncate_shared: bool = True # exact KV-share truncation | |
| image_pool: int = 2 # 16x16 grid -> (16/pool)^2 tokens per image | |
| frame_pool: int = 4 # per video frame | |
| n_latents: int = 8 # latent memory tokens per media segment (0 = none) | |
| dedup_tau: float = 0.985 # drop a video frame whose pooled feature has cos > tau with the last kept one | |
| max_frames: int = 8 | |
| sibling_from: int | None = None # from this layer on, option branches of one question see each other (listwise) | |
| # ------------------------------------------------------------------------------------------ Gemma 3n speed patch | |
| def _patch_gaussian_topk(): | |
| """The HF MLP builds a torch Normal and calls icdf on every forward of the 10 sparse layers. Cache the constant.""" | |
| from transformers.models.gemma3n import modeling_gemma3n as mg | |
| if getattr(mg.Gemma3nTextMLP, "_mmjev_patched", False): | |
| return | |
| cache = {} | |
| def _gaussian_topk(self, inputs): | |
| s = float(self.activation_sparsity) | |
| if s not in cache: | |
| cache[s] = float(torch.distributions.Normal(0, 1).icdf(torch.tensor(s))) | |
| mean = inputs.mean(-1, keepdim=True) | |
| std = inputs.std(-1, keepdim=True, unbiased=False) | |
| return F.relu(inputs - (mean + std * cache[s])) | |
| mg.Gemma3nTextMLP._gaussian_topk = _gaussian_topk | |
| mg.Gemma3nTextMLP._mmjev_patched = True | |
| # ------------------------------------------------------------------------------------------ model | |
| class MMJev(nn.Module): | |
| def __init__(self, base, processor, fast: FastConfig | None = None, n_latent_max: int = 16): | |
| super().__init__() | |
| _patch_gaussian_topk() | |
| self.base = base | |
| self.proc = processor | |
| self.tok = processor.tokenizer | |
| self.fast = fast or FastConfig() | |
| core = self._core() | |
| self.lm = core.language_model | |
| self.cfg = self.lm.config | |
| self.d = self.cfg.hidden_size | |
| self.device = self.lm.embed_tokens.weight.device | |
| self.dtype = torch.float16 | |
| t = self.tok | |
| self.id = {k: t.convert_tokens_to_ids(v) for k, v in dict( | |
| boi="<start_of_image>", eoi="<end_of_image>", boa="<start_of_audio>", eoa="<end_of_audio>", | |
| sot="<start_of_turn>", eot="<end_of_turn>").items()} | |
| self.yes_id = t.encode("Yes", add_special_tokens=False)[0] | |
| self.no_id = t.encode("No", add_special_tokens=False)[0] | |
| with torch.no_grad(): | |
| E = self.lm.embed_tokens.weight | |
| w = (E[self.yes_id].float() - E[self.no_id].float())[None] | |
| self.heads = nn.ModuleDict({str(k): nn.Linear(self.d, 1).to(self.device, torch.float32) | |
| for k in (35, FIRST_SHARED)}) | |
| for h in self.heads.values(): | |
| h.weight.copy_(w.to(self.device)); h.bias.zero_() | |
| # latent memory tokens, initialised near the embedding of a neutral word so they start in-distribution | |
| init = self._embed_ids(torch.tensor(self._t(" summary"))).float().mean(0) | |
| self.latents = nn.Parameter(init[None].repeat(n_latent_max, 1) | |
| + 0.02 * init.std() * torch.randn(n_latent_max, self.d, device=self.device)) | |
| self.softcap = getattr(self.cfg, "final_logit_softcapping", None) | |
| self.register_buffer("temperature", torch.ones(3, device=self.device)) | |
| self.vision_fn = None # optional compiled vision forward (pixel fp16 channels_last -> last_hidden_state) | |
| def _core(self): | |
| m = self.base | |
| while not (hasattr(m, "language_model") and hasattr(m, "vision_tower")): | |
| m = m.model if hasattr(m, "model") else m.base_model | |
| return m | |
| def _t(self, s: str) -> list[int]: | |
| return self.tok.encode(s, add_special_tokens=False) | |
| def _embed_ids(self, ids: torch.Tensor) -> torch.Tensor: | |
| """As Gemma3nModel.forward: text ids from embed_tokens, ids in the vision / audio hard-token ranges | |
| (e.g. <end_of_image> = 262144) from embed_vision / embed_audio.""" | |
| core = self._core() | |
| ev, ea = core.embed_vision, core.embed_audio | |
| ids = ids.to(self.device) | |
| text = ids < ev.vocab_offset | |
| out = self.lm.embed_tokens(torch.where(text, ids, torch.zeros_like(ids))).to(self.dtype) | |
| vis = (ids >= ev.vocab_offset) & (ids < ea.vocab_offset) | |
| if vis.any(): | |
| out[vis] = ev(input_ids=ids[vis][None]).to(self.dtype)[0] | |
| aud = ids >= ea.vocab_offset | |
| if aud.any(): | |
| out[aud] = ea(input_ids=ids[aud][None]).to(self.dtype)[0] | |
| return out | |
| # -------------------------------------------------------------------------------------- media encoders | |
| def vision_tower_features(self, images) -> torch.Tensor: | |
| """list[PIL] -> MobileNet-V5 grid (n, C, 16, 16) at the native 768x768 input.""" | |
| vt = self._core().vision_tower | |
| dt = next(vt.parameters()).dtype | |
| out = [] | |
| for s in range(0, len(images), 8): | |
| pv = self.proc.image_processor(images[s:s + 8], return_tensors="pt")["pixel_values"] | |
| pv = pv.to(self.device, dt).contiguous(memory_format=torch.channels_last) | |
| if self.vision_fn is not None: | |
| h = self.vision_fn(pv) | |
| else: | |
| h = vt(pixel_values=pv, do_pooling=False, return_dict=True).last_hidden_state | |
| out.append(h.to(self.dtype)) | |
| return torch.cat(out, 0) | |
| def embed_vision_grid(self, grid: torch.Tensor, pool: int) -> torch.Tensor: | |
| """(n, C, 16, 16) -> (n, (16/pool)^2, d). Pooled before projection so the soft-embedding norm sees averages.""" | |
| ev = self._core().embed_vision | |
| h = grid.float() | |
| target = max(1, 16 // pool) # grids may be cached already pooled (e.g. 8x8 / 4x4) | |
| if h.shape[-1] > target: | |
| h = F.avg_pool2d(h, h.shape[-1] // target) | |
| n, C = h.shape[:2] | |
| h = h.reshape(n, C, -1).permute(0, 2, 1) * (C ** 0.5) | |
| return ev(inputs_embeds=h.to(next(ev.parameters()).dtype)).to(self.dtype) | |
| def audio_tower_features(self, clips) -> list: | |
| """list[np 16 kHz] -> list[(T_i, C)] valid conformer outputs (before embed_audio).""" | |
| at = self._core().audio_tower | |
| enc = self.proc.feature_extractor([np.asarray(c, np.float32) for c in clips], sampling_rate=16000, | |
| return_tensors="pt", padding="longest") | |
| feats = enc["input_features"].to(self.device, next(at.parameters()).dtype) | |
| mask = enc["input_features_mask"].to(self.device) | |
| ao = at(feats, ~mask.bool(), return_dict=True) | |
| pad = ao.audio_mel_mask | |
| return [ao.last_hidden_state[i][~pad[i]].to(self.dtype) for i in range(len(clips))] | |
| def embed_audio_feats(self, a: torch.Tensor) -> torch.Tensor: | |
| ea = self._core().embed_audio | |
| return ea(inputs_embeds=a[None].to(next(ea.parameters()).dtype)).to(self.dtype)[0] | |
| def dedup_frames(grid: torch.Tensor, tau: float) -> list[int]: | |
| """Keep frame i unless its mean-pooled feature nearly copies the last kept frame (LongVU-style).""" | |
| v = F.normalize(grid.float().mean((2, 3)), dim=-1) | |
| keep = [0] | |
| for i in range(1, len(v)): | |
| if float((v[i] * v[keep[-1]]).sum()) < tau: | |
| keep.append(i) | |
| if keep[-1] != len(v) - 1: | |
| keep.append(len(v) - 1) # always keep the last frame: the "now" of the state | |
| return keep | |
| def encode_state(self, state: list[Seg], fc: FastConfig | None = None): | |
| """-> list of encoded segments, running the frozen towers unless cached features are given.""" | |
| fc = fc or self.fast | |
| out = [] | |
| for seg in state: | |
| if seg.kind == "text": | |
| out.append(("text", seg.data)) | |
| elif seg.kind == "image": | |
| g = seg.data if torch.is_tensor(seg.data) else self.vision_tower_features([seg.data])[0] | |
| out.append(("image", self.embed_vision_grid(g[None].to(self.device), fc.image_pool)[0])) | |
| elif seg.kind == "audio": | |
| a = seg.data if torch.is_tensor(seg.data) else self.audio_tower_features([seg.data])[0] | |
| out.append(("audio", self.embed_audio_feats(a.to(self.device)))) | |
| elif seg.kind == "video": | |
| frames = seg.data | |
| n = len(frames) | |
| idx = np.arange(n) | |
| if n > fc.max_frames: | |
| idx = np.linspace(0, n - 1, fc.max_frames).round().astype(int) | |
| g = frames[torch.as_tensor(idx)] if torch.is_tensor(frames) else \ | |
| self.vision_tower_features([frames[i] for i in idx]) | |
| g = g.to(self.device) | |
| keep = self.dedup_frames(g, fc.dedup_tau) if fc.dedup_tau < 1 else list(range(len(g))) | |
| emb = self.embed_vision_grid(g[keep], fc.frame_pool) | |
| out.append(("video", emb, [float(idx[k]) / seg.fps for k in keep])) | |
| if seg.audio is not None: | |
| a = seg.audio if torch.is_tensor(seg.audio) else self.audio_tower_features([seg.audio])[0] | |
| out.append(("soundtrack", self.embed_audio_feats(a.to(self.device)))) | |
| return out | |
| # -------------------------------------------------------------------------------------- tree plan | |
| def plan(self, enc, questions: list[dict], fc: FastConfig | None = None): | |
| """Nodes of the tree; a node holds pieces ('ids', list) | ('emb', Tensor, is_media) | ('lat', M).""" | |
| M = (fc or self.fast).n_latents | |
| pieces = [("ids", [self.tok.bos_token_id, self.id["sot"]] + self._t("user\n" + SYSTEM + "\n\nState:\n"))] | |
| for e in enc: | |
| kind = e[0] | |
| if kind == "text": | |
| pieces.append(("ids", self._t(e[1] + "\n"))) | |
| elif kind == "image": | |
| pieces += [("ids", self._t("Image: ") + [self.id["boi"]]), ("emb", e[1], True)] | |
| if M: | |
| pieces.append(("lat", M)) | |
| pieces.append(("ids", [self.id["eoi"]] + self._t("\n"))) | |
| elif kind in ("audio", "soundtrack"): | |
| lead = "Audio: " if kind == "audio" else "Video soundtrack: " | |
| pieces += [("ids", self._t(lead) + [self.id["boa"]]), ("emb", e[1], True)] | |
| if M: | |
| pieces.append(("lat", M)) | |
| pieces.append(("ids", [self.id["eoa"]] + self._t("\n"))) | |
| elif kind == "video": | |
| emb, times = e[1], e[2] | |
| pieces.append(("ids", self._t(f"Video, {len(times)} key frames:\n"))) | |
| for j, t in enumerate(times): | |
| pieces += [("ids", self._t(f"{t:.1f}s ") + [self.id["boi"]]), ("emb", emb[j], True), | |
| ("ids", [self.id["eoi"]])] | |
| if M: | |
| pieces.append(("lat", M)) | |
| pieces.append(("ids", self._t("\n"))) | |
| nodes = [dict(parent=-1, pieces=pieces)] | |
| q_opts = [] | |
| for q in questions: | |
| _, texts = options_of(q) | |
| nodes.append(dict(parent=0, pieces=[("ids", self._t( | |
| f"\nQuestion ({q['type']}): {q['instructions'].strip()}{self.option_context(q)}\nProposed answer: "))])) | |
| qn = len(nodes) - 1 | |
| ids = [] | |
| for txt in texts: | |
| nodes.append(dict(parent=qn, pieces=[("ids", self._t(txt) + [self.id["eot"]] + self._t("\n") | |
| + [self.id["sot"]] + self._t("model\n"))])) | |
| ids.append(len(nodes) - 1) | |
| q_opts.append(ids) | |
| return nodes, q_opts | |
| option_context_on = True | |
| option_context_max = 8 | |
| def option_context(self, q) -> str: | |
| """Canonical option context: every branch sees the whole answer space, listed in a canonical order (sorted | |
| labels for `choice`, the intrinsic level order for `score`), so the prefix -- and hence every probability -- | |
| is still independent of the order the caller passed the options in.""" | |
| if not self.option_context_on or q["type"] == "noul": | |
| return "" | |
| labels, _ = options_of(q) | |
| if len(labels) > self.option_context_max: | |
| return "" # large label spaces: siblings compare in-attention instead (sibling_from) | |
| if q["type"] == "score": | |
| crit = list(q["criteria"]) | |
| return "\nScale: " + "; ".join(f"level {i} = {str(c)[:80]}" for i, c in enumerate(crit)) | |
| return "\nOptions: " + "; ".join(sorted(labels, key=lambda x: x.lower())) | |
| def isolate(plan): | |
| """One linear plan per option (root -> question -> that option): the naive K-pass baseline.""" | |
| nodes, q_opts = plan | |
| out = [] | |
| for opts in q_opts: | |
| for o in opts: | |
| qn = nodes[o]["parent"] | |
| out.append(([dict(nodes[0]), dict(nodes[qn], parent=0), dict(nodes[o], parent=1)], [[2]])) | |
| return out | |
| def pack(self, plans): | |
| """Concatenate the node pieces of every plan into right-padded tensors plus the tree mask.""" | |
| B = len(plans) | |
| rows = [] | |
| for nodes, q_opts in plans: | |
| parts, ple, pos, node_of, media = [], [], [], [], [] | |
| start, length, end = {}, {}, {} | |
| cur = 0 | |
| for n, nd in enumerate(nodes): | |
| p0 = 0 if nd["parent"] < 0 else start[nd["parent"]] + length[nd["parent"]] | |
| start[n], ln = p0, 0 | |
| for pc in nd["pieces"]: | |
| if pc[0] == "ids": | |
| ids = torch.tensor(pc[1], dtype=torch.long) | |
| parts.append(("ids", ids)); k = len(ids) | |
| ple.append(torch.where(ids < self.cfg.vocab_size_per_layer_input, ids, torch.zeros_like(ids))) | |
| media += [False] * k | |
| elif pc[0] == "emb": | |
| parts.append(("emb", pc[1])); k = len(pc[1]) | |
| ple.append(torch.zeros(k, dtype=torch.long)); media += [pc[2]] * k | |
| else: | |
| k = pc[1] | |
| parts.append(("lat", k)); ple.append(torch.zeros(k, dtype=torch.long)); media += [False] * k | |
| ln += k | |
| length[n] = ln | |
| pos.append(torch.arange(p0, p0 + ln)); node_of += [n] * ln | |
| cur += ln | |
| end[n] = cur - 1 | |
| A = torch.zeros(len(nodes), len(nodes), dtype=torch.bool) | |
| for i in range(len(nodes)): | |
| j = i | |
| while j >= 0: | |
| A[i, j] = True | |
| j = nodes[j]["parent"] | |
| qof = torch.full((len(nodes),), -1, dtype=torch.long) # option node -> its question node | |
| for opts in q_opts: | |
| for o in opts: | |
| qof[o] = nodes[o]["parent"] | |
| rows.append(dict(parts=parts, ple=torch.cat(ple), pos=torch.cat(pos), node=torch.tensor(node_of), | |
| media=torch.tensor(media), A=A, qof=qof, ends=[[end[o] for o in opts] for opts in q_opts])) | |
| L = max(len(r["ple"]) for r in rows) | |
| ple = torch.zeros(B, L, dtype=torch.long) | |
| pos = torch.zeros(B, L, dtype=torch.long) | |
| valid = torch.zeros(B, L, dtype=torch.bool) | |
| media = torch.zeros(B, L, dtype=torch.bool) | |
| mask = torch.zeros(B, L, L, dtype=torch.bool) | |
| sib = torch.zeros(B, L, L, dtype=torch.bool) | |
| causal = torch.ones(L, L, dtype=torch.bool).tril() | |
| lat = self.latents.to(self.dtype) | |
| xs = [] | |
| for b, r in enumerate(rows): | |
| ids_all = [p[1] for p in r["parts"] if p[0] == "ids"] | |
| tok = self._embed_ids(torch.cat(ids_all)) | |
| seq, off = [], 0 | |
| for p in r["parts"]: | |
| if p[0] == "ids": | |
| seq.append(tok[off:off + len(p[1])]); off += len(p[1]) | |
| elif p[0] == "emb": | |
| seq.append(p[1].to(self.dtype)) | |
| else: | |
| seq.append(lat[:p[1]]) | |
| seq = torch.cat(seq, 0) | |
| n = len(seq) | |
| xs.append(F.pad(seq, (0, 0, 0, L - n))) # keeps the graph to the latent parameters | |
| ple[b, :n], pos[b, :n], valid[b, :n], media[b, :n] = r["ple"], r["pos"], True, r["media"] | |
| pos[b, n:] = r["pos"].max() + 1 | |
| nd = r["node"] | |
| mask[b, :n, :n] = r["A"][nd[:, None], nd[None, :]] & causal[:n, :n] | |
| # sibling bridge (Set-Encoder style): an option-branch token may read the END token of every sibling | |
| # branch of the same question -- one summary per option, positions stay tree positions -> equivariant | |
| qn = r["qof"][nd] # question of each token's option branch (-1: none) | |
| is_end = torch.zeros(n, dtype=torch.bool) | |
| is_end[[i for e in r["ends"] for i in e]] = True | |
| sib[b, :n, :n] = (qn[:, None] == qn[None, :]) & (qn[:, None] >= 0) & is_end[None, :] | |
| mask |= torch.eye(L, dtype=torch.bool)[None] # padded rows see themselves (no all -inf rows) | |
| return dict(x=torch.stack(xs), ple=ple, pos=pos, valid=valid, media=media, mask=mask, sib=sib | mask, | |
| ends=[r["ends"] for r in rows]) | |
| # -------------------------------------------------------------------------------------- decoder loop | |
| def _masks(self, mask, pos_q, pos_k): | |
| """Full and sliding-window boolean masks [B,1,Q,K] from the tree mask and the tree position ids.""" | |
| s = mask & ((pos_q[:, :, None] - pos_k[:, None, :]) < self.cfg.sliding_window) | |
| return {"full_attention": mask[:, None], "sliding_attention": s[:, None]} | |
| def host_indices(self, pk, fc: FastConfig, L_pad: int | None = None, K_pad: int | None = None, | |
| Lk_pad: int | None = None): | |
| """Host-side index tensors: scoring positions in packed coordinates (sidx_full) and after the modality exit | |
| (sidx_kept), and the kept-token gather (kidx, kval). Optional padding to fixed buckets for CUDA graphs.""" | |
| ends, valid, med = pk["ends"], pk["valid"], pk["media"] | |
| B, L = valid.shape | |
| K = K_pad or max(sum(len(e) for e in r) for r in ends) | |
| sidx = torch.zeros(B, K, dtype=torch.long) | |
| for b, r in enumerate(ends): | |
| flat = [i for e in r for i in e] | |
| sidx[b, :len(flat)] = torch.tensor(flat) | |
| out = dict(sidx_full=sidx, sidx_kept=sidx.clone(), kidx=None, kval=None) | |
| if fc.media_exit is not None: | |
| keep = valid & ~med | |
| Lk = Lk_pad or int(keep.sum(1).max()) | |
| kidx = torch.zeros(B, Lk, dtype=torch.long) | |
| kval = torch.zeros(B, Lk, dtype=torch.bool) | |
| remap = torch.zeros(B, L_pad or L, dtype=torch.long) | |
| for b in range(B): | |
| ii = keep[b].nonzero().squeeze(1) | |
| kidx[b, :len(ii)], kval[b, :len(ii)] = ii, True | |
| if len(ii) < Lk: # padding slots point at a padding token of the packed sequence | |
| kidx[b, len(ii):] = (L_pad or L) - 1 | |
| remap[b, ii] = torch.arange(len(ii)) | |
| out.update(sidx_kept=remap.gather(1, sidx), kidx=kidx, kval=kval) | |
| return out | |
| def core(self, fc: FastConfig, x, ple, pos, mask, sidx_full, sidx_kept, kidx=None, kval=None, sib=None): | |
| """Pure-GPU decoder: Gemma3nTextModel.forward re-implemented with the modality exit, the KV-share truncation | |
| and the early exit. No host syncs, so it can be captured in a CUDA graph. Returns scores (B, K) float32.""" | |
| if self.fp32_residual: | |
| # AltUp residual streams reach |h| ~ 700, where fp16 has ~0.5 resolution: keep the streams in fp32 and | |
| # let autocast run every matmul in fp16 (vs a per-layer fp32 reference: max |d logit| 2.5 -> ~0.05) | |
| with torch.autocast("cuda", dtype=torch.float16): | |
| return self._decoder(fc, x.float(), ple.float(), pos, mask, sidx_full, sidx_kept, kidx, kval, sib) | |
| return self._decoder(fc, x, ple, pos, mask, sidx_full, sidx_kept, kidx, kval, sib) | |
| fp32_residual = True | |
| def _decoder(self, fc, x, ple, pos, mask, sidx_full, sidx_kept, kidx=None, kval=None, sib=None): | |
| lm, cfg, dev = self.lm, self.cfg, self.device | |
| B, L = x.shape[:2] | |
| K = sidx_full.shape[1] | |
| per_layer = lm.project_per_layer_inputs(x, ple) | |
| eps = torch.full((), 1e-5, device=dev) | |
| target = torch.mean(x ** 2, dim=-1, keepdim=True) ** 0.5 | |
| hs = [x] | |
| for i in range(1, cfg.altup_num_inputs): | |
| p = lm.altup_projections[i - 1](x).to(x.dtype) | |
| hs.append(p * target / torch.sqrt(torch.maximum(torch.mean(p ** 2, dim=-1, keepdim=True), eps))) | |
| hs = torch.stack(hs, 0) | |
| def rope(p): | |
| return {lt: lm.rotary_emb(hs, p, lt) for lt in set(cfg.layer_types)} | |
| def take(t, idx, dim, bdim): | |
| """Gather positions idx (B, n) along `dim` of t, whose batch axis is `bdim`.""" | |
| shape = list(t.shape); shape[dim] = idx.shape[1] | |
| view = [1] * t.dim(); view[dim] = idx.shape[1]; view[bdim] = B | |
| return t.gather(dim, idx.reshape(view).expand(shape)) | |
| sidx = sidx_full | |
| use_sib = fc.sibling_from is not None and sib is not None | |
| if not use_sib: | |
| sib = mask | |
| pe, cur_pos = rope(pos), pos | |
| masks, smasks = self._masks(mask, pos, pos), self._masks(sib, pos, pos) | |
| qmask = qsib = None | |
| shared = {} | |
| for i in range(fc.exit_layer): | |
| if fc.media_exit is not None and i == fc.media_exit: | |
| Lk = kidx.shape[1] | |
| hs, per_layer, cur_pos = take(hs, kidx, 2, 1), take(per_layer, kidx, 1, 0), cur_pos.gather(1, kidx) | |
| eye = torch.eye(Lk, dtype=torch.bool, device=dev)[None] | |
| mask = (take(take(mask, kidx, 1, 0), kidx, 2, 0) & kval[:, None, :]) | eye | |
| sib = (take(take(sib, kidx, 1, 0), kidx, 2, 0) & kval[:, None, :]) | eye | |
| pe = rope(cur_pos) | |
| masks, smasks = self._masks(mask, cur_pos, cur_pos), self._masks(sib, cur_pos, cur_pos) | |
| sidx = sidx_kept | |
| if i == FIRST_SHARED and fc.truncate_shared: | |
| hs, per_layer = take(hs, sidx, 2, 1), take(per_layer, sidx, 1, 0) | |
| qpos = cur_pos.gather(1, sidx) | |
| pe = rope(qpos) | |
| masks = self._masks(take(mask, sidx, 1, 0), qpos, cur_pos) | |
| smasks = self._masks(take(sib, sidx, 1, 0), qpos, cur_pos) | |
| sidx = torch.arange(K, device=dev)[None].expand(B, K) | |
| lt = cfg.layer_types[i] | |
| m = smasks if (use_sib and i >= fc.sibling_from) else masks | |
| hs = lm.layers[i](hs, pe[lt], per_layer[:, :, i, :], shared_kv_states=shared, | |
| attention_mask=m[lt], position_ids=None) | |
| hs = take(hs, sidx, 2, 1) | |
| target = torch.mean(hs[0] ** 2, dim=-1, keepdim=True) ** 0.5 | |
| outs = [hs[0]] | |
| for i in range(1, cfg.altup_num_inputs): | |
| p = lm.altup_unembed_projections[i - 1](hs[i]).to(x.dtype) | |
| outs.append(p * target / torch.sqrt(torch.maximum(torch.mean(p ** 2, dim=-1, keepdim=True), eps))) | |
| h = lm.norm(torch.stack(outs).mean(0)) | |
| s = self.heads[str(35 if fc.exit_layer >= 35 else FIRST_SHARED)](h.float()).squeeze(-1) | |
| if self.softcap: | |
| s = torch.tanh(s / self.softcap) * self.softcap | |
| return s | |
| def ple_lookup(self, ple_ids): | |
| """Per-layer embeddings gathered on the CPU (the 4.7 GB table lives there), copied to the GPU.""" | |
| lm, cfg = self.lm, self.cfg | |
| w = lm.embed_tokens_per_layer.weight | |
| e = F.embedding(ple_ids.to(w.device), w) * lm.embed_tokens_per_layer.embed_scale.to(w.dtype) | |
| B, L = ple_ids.shape | |
| return e.reshape(B, L, cfg.num_hidden_layers, cfg.hidden_size_per_layer_input) | |
| def split_scores(s, ends): | |
| res = [] | |
| for b, r in enumerate(ends): | |
| out, o = [], 0 | |
| for e in r: | |
| out.append(s[b, o:o + len(e)]); o += len(e) | |
| res.append(out) | |
| return res | |
| def run(self, pk, fc: FastConfig | None = None): | |
| """Eager path (training and reference). Returns per sample, per question, the option logits.""" | |
| fc = fc or self.fast | |
| dev = self.device | |
| ix = self.host_indices(pk, fc) | |
| g = lambda t: None if t is None else t.to(dev, non_blocking=True) | |
| s = self.core(fc, pk["x"], g(self.ple_lookup(pk["ple"])).to(self.dtype), g(pk["pos"]), g(pk["mask"]), | |
| g(ix["sidx_full"]), g(ix["sidx_kept"]), g(ix["kidx"]), g(ix["kval"]), | |
| g(pk.get("sib")) if fc.sibling_from is not None else None) | |
| return self.split_scores(s, pk["ends"]) | |
| # -------------------------------------------------------------------------------------- CUDA graphs | |
| BUCKETS_L = (96, 128, 160, 192, 256, 320, 384, 512, 640, 768, 1024, 1536, 2048) | |
| BUCKETS_K = (4, 8, 16, 32, 64, 128) | |
| def _bucket(n, buckets): | |
| for b in buckets: | |
| if n <= b: | |
| return b | |
| return n | |
| compile_core = False | |
| max_graphs = 8 | |
| def _compiled(self, fc): | |
| """Inductor-fused decoder (elementwise chains of AltUp / LAuReL / PLE / norms fused), one per FastConfig.""" | |
| key = (fc.media_exit, fc.exit_layer, fc.truncate_shared) | |
| if not hasattr(self, "_compiled_fns"): | |
| self._compiled_fns = {} | |
| if key not in self._compiled_fns: | |
| import functools | |
| self._compiled_fns[key] = torch.compile(functools.partial(self.core, fc), dynamic=False, | |
| mode="max-autotune-no-cudagraphs") | |
| return self._compiled_fns[key] | |
| def run_graph(self, pk, fc: FastConfig | None = None): | |
| """Batch-1 latency path: pad to (L, K, Lk) buckets and replay a captured CUDA graph (one per bucket).""" | |
| fc = fc or self.fast | |
| assert pk["x"].shape[0] == 1 | |
| dev = self.device | |
| L0 = pk["x"].shape[1] | |
| L = self._bucket(L0, self.BUCKETS_L) | |
| K = self._bucket(max(sum(len(e) for e in r) for r in pk["ends"]), self.BUCKETS_K) | |
| Lk = None | |
| if fc.media_exit is not None: | |
| Lk = self._bucket(int((pk["valid"] & ~pk["media"]).sum()), self.BUCKETS_L) | |
| Lk = min(Lk, L) | |
| key = (L, K, Lk, fc.media_exit, fc.exit_layer, fc.truncate_shared, fc.sibling_from) | |
| if not hasattr(self, "_graphs"): | |
| self._graphs = {} | |
| # pad the packed inputs to the bucket | |
| pad = L - L0 | |
| x = F.pad(pk["x"], (0, 0, 0, pad)) | |
| ple_ids = F.pad(pk["ple"], (0, pad)) | |
| pos = torch.cat([pk["pos"], pk["pos"].max() + 1 + torch.arange(pad)[None]], 1) | |
| mask = torch.zeros(1, L, L, dtype=torch.bool) | |
| mask[:, :L0, :L0] = pk["mask"] | |
| mask |= torch.eye(L, dtype=torch.bool)[None] | |
| sib = torch.zeros(1, L, L, dtype=torch.bool) | |
| sib[:, :L0, :L0] = pk["sib"] | |
| sib |= torch.eye(L, dtype=torch.bool)[None] | |
| valid = F.pad(pk["valid"], (0, pad)); media = F.pad(pk["media"], (0, pad)) | |
| ix = self.host_indices(dict(ends=pk["ends"], valid=valid, media=media), fc, L_pad=L, K_pad=K, Lk_pad=Lk) | |
| ple = self.ple_lookup(ple_ids).to(self.dtype) | |
| feeds = dict(x=x, ple=ple, pos=pos, mask=mask, sidx_full=ix["sidx_full"], sidx_kept=ix["sidx_kept"]) | |
| if Lk is not None: | |
| feeds.update(kidx=ix["kidx"], kval=ix["kval"]) | |
| if fc.sibling_from is not None: | |
| feeds.update(sib=sib) | |
| if key not in self._graphs: | |
| while len(self._graphs) >= self.max_graphs: # LRU: every graph pins a private memory pool | |
| self._graphs.pop(next(iter(self._graphs))) | |
| torch.cuda.synchronize(); torch.cuda.empty_cache() | |
| static = {k: v.to(dev).clone() for k, v in feeds.items()} | |
| side = torch.cuda.Stream() | |
| side.wait_stream(torch.cuda.current_stream()) | |
| fn = self._compiled(fc) if self.compile_core else (lambda **kw: self.core(fc, **kw)) | |
| with torch.cuda.stream(side): | |
| for _ in range(3): | |
| fn(**static) | |
| torch.cuda.current_stream().wait_stream(side) | |
| graph = torch.cuda.CUDAGraph() | |
| with torch.cuda.graph(graph): | |
| out = fn(**static) | |
| self._graphs[key] = (graph, static, out) | |
| self._graphs[key] = self._graphs.pop(key) # mark as most recently used | |
| graph, static, out = self._graphs[key] | |
| for k, v in feeds.items(): | |
| static[k].copy_(v, non_blocking=True) | |
| graph.replay() | |
| return self.split_scores(out.clone(), pk["ends"]) | |
| # -------------------------------------------------------------------------------------- public API | |
| def prepare(self, state, questions, fc: FastConfig | None = None): | |
| if isinstance(state, str): | |
| state = [Seg("text", state)] | |
| elif isinstance(state, Seg): | |
| state = [state] | |
| return self.plan(self.encode_state(state, fc), questions, fc) | |
| def decide(self, state, questions, calibrated: bool = True, fc: FastConfig | None = None, graph: bool = False): | |
| """Jev-style call. questions: {name: {type, instructions, criteria}} (or a list). | |
| noul -> {"noul": p_yes}; choice -> {"choice", "probabilities"}; score -> {"score" in [0,1], "level", ...}.""" | |
| names = list(questions) if isinstance(questions, dict) else list(range(len(questions))) | |
| qs = [questions[n] for n in names] | |
| self.eval() | |
| pk = self.pack([self.prepare(state, qs, fc)]) | |
| logits = (self.run_graph(pk, fc) if graph else self.run(pk, fc))[0] | |
| out = {} | |
| for n, q, lg in zip(names, qs, logits): | |
| T = float(self.temperature[QTYPES[q["type"]]]) if calibrated else 1.0 | |
| if calibrated and getattr(self, "temps_k", None): | |
| T = self.temps_k.get(f"{q['type']}:{k_bucket(len(lg))}", T) | |
| p = torch.softmax(lg / T, -1).cpu().numpy() | |
| labels, _ = options_of(q) | |
| probs = {l: float(v) for l, v in zip(labels, p)} | |
| conf = 1.0 - float(-(p * np.log(np.clip(p, 1e-12, 1))).sum() / math.log(max(len(p), 2))) | |
| if q["type"] == "noul": | |
| out[n] = {"noul": float(p[1]), "confidence": conf} | |
| elif q["type"] == "score": | |
| ev = float((p * np.arange(len(p))).sum() / max(len(p) - 1, 1)) | |
| out[n] = {"score": ev, "level": int(p.argmax()), "probabilities": probs, "confidence": conf} | |
| else: | |
| out[n] = {"choice": labels[int(p.argmax())], "probabilities": probs, "confidence": conf} | |
| return out | |
| # ------------------------------------------------------------------------------------------ loss | |
| def decision_loss(logits: torch.Tensor, target: torch.Tensor, qtype: str, w_rps: float = 1.0): | |
| """Strictly proper: log score (soft CE) for every type, + ranked probability score for ordinal questions.""" | |
| logp = torch.log_softmax(logits, -1) | |
| loss = -(target * logp).sum() | |
| if qtype == "score" and len(logits) > 1: | |
| p = logp.exp() | |
| loss = loss + w_rps * ((p.cumsum(-1) - target.cumsum(-1)) ** 2).sum() / (len(logits) - 1) | |
| return loss | |
| # ------------------------------------------------------------------------------------------ loading | |
| def _cast(mod, dt): | |
| for p in mod.parameters(): | |
| if p.dtype in (torch.bfloat16, torch.float16, torch.float32): | |
| p.data = p.data.to(dt) | |
| for b in mod.buffers(): | |
| if b.dtype in (torch.bfloat16, torch.float16, torch.float32): | |
| b.data = b.data.to(dt) | |
| def load_gemma3n(path: str, vision_dtype=torch.float32, audio_dtype=torch.float32, ple_on_gpu: bool | None = None): | |
| """Loaded on the CPU (safetensors are mmapped) and moved to the GPU module by module in fp16 -- except the 4.7 GB | |
| per-layer-embedding table, which stays on the CPU (only gathered from; a multi-device device_map would make | |
| accelerate copy it to the GPU). Loaded as bf16 first: the audio tower's 1e10 clamp constant overflows fp16. | |
| MobileNet-V5 overflows fp16 as shipped; call fp16_safe_vision() to run it in fp16.""" | |
| from transformers import AutoProcessor, Gemma3nForConditionalGeneration | |
| model = Gemma3nForConditionalGeneration.from_pretrained(path, dtype=torch.bfloat16, device_map="cpu", | |
| attn_implementation="sdpa") | |
| core = model.model | |
| lm = core.language_model | |
| ple = lm.embed_tokens_per_layer | |
| for _, child in lm.named_children(): | |
| if child is not ple: | |
| child.to("cuda", torch.float16) | |
| for name, b in list(lm.named_buffers(recurse=False)): | |
| setattr(lm, name, b.to("cuda", torch.float16 if b.is_floating_point() else b.dtype)) | |
| if ple_on_gpu is None: # 24 GB+ cards (L4, A100) keep the 4.7 GB table on the GPU | |
| ple_on_gpu = torch.cuda.get_device_properties(0).total_memory > 20 * 2**30 | |
| ple.to("cuda" if ple_on_gpu else "cpu", torch.float16) | |
| core.embed_vision.to("cuda", torch.float16) | |
| core.embed_audio.to("cuda", torch.float16) | |
| core.vision_tower.to("cuda"); _cast(core.vision_tower, vision_dtype) | |
| core.audio_tower.to("cuda"); _cast(core.audio_tower, audio_dtype) | |
| torch.cuda.empty_cache() | |
| return model, AutoProcessor.from_pretrained(path) | |
| # ------------------------------------------------------------------------------------------ MatFormer width | |
| _FFN_ORIG = {} | |
| def set_ffn_width(lm, width: int | None): | |
| """MatFormer elastic width: use the first `width` FFN neurons of every layer (E2B width = 8192). None restores.""" | |
| for i, layer in enumerate(lm.layers): | |
| mods = [getattr(layer.mlp, n) for n in ("gate_proj", "up_proj", "down_proj")] | |
| mods = [getattr(m, "base_layer", m) for m in mods] | |
| if i not in _FFN_ORIG: | |
| _FFN_ORIG[i] = [m.weight for m in mods] | |
| g, u, d = _FFN_ORIG[i] | |
| ws = (g, u, d) if width is None else ( | |
| nn.Parameter(g.data[:width], requires_grad=False), nn.Parameter(u.data[:width], requires_grad=False), | |
| nn.Parameter(d.data[:, :width], requires_grad=False)) | |
| for m, wt in zip(mods, ws): | |
| m.weight = wt | |
| # ------------------------------------------------------------------------------------------ fp16-safe MobileNet-V5 | |
| def _rms_norm2d_fp32(x, normalized_shape, weight=None, eps=1e-5): | |
| """timm's rms_norm2d squares x in its own dtype: in fp16 |x| > 256 overflows. Statistics in fp32 instead.""" | |
| v = x.float().pow(2).mean(dim=1, keepdim=True) | |
| y = (x.float() * torch.rsqrt(v + eps)).to(x.dtype) | |
| if weight is not None: | |
| y = y * weight.reshape(1, -1, 1, 1).to(y.dtype) | |
| return y | |
| def fp16_safe_vision(vision_tower, calib_pixels: torch.Tensor, headroom: float = 4096.0): | |
| """Run MobileNet-V5 in fp16 without overflow, exactly up to eps: | |
| 1. RMSNorm statistics in fp32; | |
| 2. every conv whose output feeds straight into an RMSNorm is divided by a power of two s so that its fp32 | |
| calibration abs-max stays below `headroom` -- RMSNorm(conv(x) / s) == RMSNorm(conv(x)). | |
| Returns the number of rescaled convs.""" | |
| import timm.layers.norm_act as na | |
| import timm.layers.norm as nm | |
| for mod in (na, nm): | |
| mod.rms_norm2d = _rms_norm2d_fp32 | |
| if hasattr(mod, "fast_rms_norm2d"): | |
| mod.fast_rms_norm2d = _rms_norm2d_fp32 | |
| tm = vision_tower.timm_model | |
| for m in tm.modules(): | |
| if hasattr(m, "_fast_norm"): | |
| m._fast_norm = False | |
| pairs = [] | |
| for m in tm.modules(): | |
| if hasattr(m, "conv") and hasattr(m, "bn") and "Rms" in type(m.bn).__name__: | |
| pairs.append(m.conv) | |
| if isinstance(m, nn.Sequential) and hasattr(m, "down_conv") and hasattr(m, "norm"): | |
| pairs.append(m.down_conv) | |
| _cast(vision_tower, torch.float32) | |
| amax = {} | |
| hooks = [c.register_forward_hook(lambda mod, i, o: amax.__setitem__(mod, max(amax.get(mod, 0.0), float(o.abs().max())))) | |
| for c in pairs] | |
| for s in range(0, len(calib_pixels), 4): | |
| vision_tower(pixel_values=calib_pixels[s:s + 4].float().cuda(), do_pooling=False, return_dict=True) | |
| for h in hooks: | |
| h.remove() | |
| n = 0 | |
| for c in pairs: | |
| s = 2.0 ** max(0, math.ceil(math.log2(max(amax.get(c, 0.0), 1e-6) / headroom))) | |
| if s > 1: | |
| c.weight.div_(s); n += 1 | |
| if c.bias is not None: | |
| c.bias.div_(s) | |
| _cast(vision_tower, torch.float16) | |
| vision_tower.to(memory_format=torch.channels_last) | |
| return n | |
| # ------------------------------------------------------------------------------------------ adapter I/O | |
| class LoRALinear(nn.Module): | |
| """y = W x + (B A x) * alpha / r; slices A/B when the wrapped FFN projection is MatFormer-sliced.""" | |
| def __init__(self, base, r=16, alpha=32, dropout=0.05): | |
| super().__init__() | |
| self.base_layer = base | |
| self.lora_A = nn.Parameter(torch.randn(r, base.in_features, device=base.weight.device) / math.sqrt(base.in_features)) | |
| self.lora_B = nn.Parameter(torch.zeros(base.out_features, r, device=base.weight.device)) | |
| self.scale, self.drop = alpha / r, nn.Dropout(dropout) | |
| def weight(self): | |
| return self.base_layer.weight | |
| def forward(self, x): | |
| y = self.base_layer(x) | |
| out_f, in_f = self.base_layer.weight.shape | |
| lx = F.linear(F.linear(self.drop(x), self.lora_A[:, :in_f].to(x.dtype)), self.lora_B[:out_f].to(x.dtype)) | |
| return y + lx * self.scale | |
| def add_lora(lm, r=16, alpha=32, mlp_from=10): | |
| """Attention q/k/v/o in every layer (KV-shared layers only have q/o) + MLP gate/up/down from `mlp_from`.""" | |
| for li, layer in enumerate(lm.layers): | |
| for name in ("q_proj", "k_proj", "v_proj", "o_proj"): | |
| m = getattr(layer.self_attn, name, None) | |
| if isinstance(m, nn.Linear): | |
| setattr(layer.self_attn, name, LoRALinear(m, r, alpha)) | |
| if li >= mlp_from: | |
| for name in ("gate_proj", "up_proj", "down_proj"): | |
| m = getattr(layer.mlp, name) | |
| if isinstance(m, nn.Linear): | |
| setattr(layer.mlp, name, LoRALinear(m, r, alpha)) | |
| PRESETS = { | |
| "full": (FastConfig(n_latents=8, exit_layer=35, media_exit=None, sibling_from=14), None), | |
| "fast": (FastConfig(n_latents=8, exit_layer=20, media_exit=8, sibling_from=14), 8192), | |
| } | |
| def _from_adapter(cls, base, processor, adapter_path: str, config_path: str | None = None, preset: str = "fast"): | |
| """Build MM-Jev from a Gemma 3n model + the released adapter (LoRA, decision heads, latents, temperatures).""" | |
| import json | |
| from safetensors.torch import load_file | |
| jev = cls(base, processor) | |
| add_lora(jev.lm) | |
| sd = load_file(adapter_path) | |
| params = dict(jev.base.named_parameters()) | |
| with torch.no_grad(): | |
| for k, v in sd.items(): | |
| if k.startswith("lora."): | |
| params[k[5:]].copy_(v.to(params[k[5:]].device)) | |
| jev.heads.load_state_dict({k[6:]: v for k, v in sd.items() if k.startswith("heads.")}) | |
| jev.latents.copy_(sd["latents"].to(jev.latents.device)) | |
| fc, width = PRESETS[preset] | |
| jev.fast = fc | |
| set_ffn_width(jev.lm, width) | |
| if config_path: | |
| cfg = json.load(open(config_path)) | |
| jev.temps_k = cfg.get("temperatures", {}).get(preset, {}) | |
| jev.eval() | |
| return jev | |
| MMJev.from_adapter = classmethod(_from_adapter) | |