"""Model-card demos (run in the Colab kernel with the trained `jev`). snake_demo : the model plays Snake from a TEXT state, one `choice` decision per tick; per-move latency overlaid. image/audio/video demos : several typed questions answered in ONE pass over a held-out sample, rendered as cards. """ import json, os, time, random, subprocess import numpy as np, torch from PIL import Image, ImageDraw, ImageFont from mmjev import Seg OUT = "/content/demo_assets" os.makedirs(OUT, exist_ok=True) def font(sz): for p in ("/usr/share/fonts/truetype/dejavu/DejaVuSansMono-Bold.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf"): if os.path.exists(p): return ImageFont.truetype(p, sz) return ImageFont.load_default() # ------------------------------------------------------------------------------------------ snake DIRS = {"up": (0, -1), "down": (0, 1), "left": (-1, 0), "right": (1, 0)} OPP = {"up": "down", "down": "up", "left": "right", "right": "left"} class Snake: def __init__(self, n=12, seed=0): self.n, self.rng = n, random.Random(seed) c = n // 2 self.body = [(c, c), (c - 1, c), (c - 2, c)] self.dir, self.alive, self.score, self.steps = "right", True, 0, 0 self.place_food() def place_food(self): free = [(x, y) for x in range(self.n) for y in range(self.n) if (x, y) not in self.body] self.food = self.rng.choice(free) def blocked(self, d): hx, hy = self.body[0] dx, dy = DIRS[d] nx, ny = hx + dx, hy + dy if not (0 <= nx < self.n and 0 <= ny < self.n): return "wall" if (nx, ny) in self.body[:-1]: return "own body" return None def state_text(self): hx, hy = self.body[0] fx, fy = self.food dx, dy = fx - hx, fy - hy rel = [] if dx: rel.append(f"{abs(dx)} {'right' if dx > 0 else 'left'}") if dy: rel.append(f"{abs(dy)} {'down' if dy > 0 else 'up'}") blk = "; ".join(f"{d}: {self.blocked(d) or 'free'}" for d in DIRS) return (f"Snake game on a {self.n}x{self.n} board (x grows right, y grows down). Snake length {len(self.body)}, " f"head at x={hx}, y={hy}, moving {self.dir}. Food at x={fx}, y={fy}, i.e. {', '.join(rel) or 'here'} " f"from the head. Next cell in each direction -> {blk}. Moving into a wall or the snake's own body " f"ends the game; reversing into the neck ({OPP[self.dir]}) is not allowed.") def question(self): return {"type": "choice", "instructions": "Which way should the snake move next to get closer to the food " "without dying?", "criteria": {"up": "move one cell up (y - 1)", "down": "move one cell down (y + 1)", "left": "move one cell left (x - 1)", "right": "move one cell right (x + 1)"}} def step(self, d): if d == OPP[self.dir]: d = self.dir self.dir = d if self.blocked(d): self.alive = False return hx, hy = self.body[0] nh = (hx + DIRS[d][0], hy + DIRS[d][1]) self.body.insert(0, nh) if nh == self.food: self.score += 1 self.place_food() else: self.body.pop() self.steps += 1 def render(self, lat_ms, probs, cell=40, title="MM-Jev plays Snake", subtitle="text state -> 5 questions, 1 pass", bar_label="P"): W = max(self.n * cell, 400) im = Image.new("RGB", (W + 440, W), (18, 20, 26)) d = ImageDraw.Draw(im) for x in range(self.n): for y in range(self.n): d.rectangle([x * cell, y * cell, x * cell + cell - 1, y * cell + cell - 1], outline=(32, 36, 44)) fx, fy = self.food d.ellipse([fx * cell + 6, fy * cell + 6, fx * cell + cell - 6, fy * cell + cell - 6], fill=(235, 70, 70)) for i, (x, y) in enumerate(self.body): c = (90, 220, 120) if i == 0 else (50, 160, 90) d.rounded_rectangle([x * cell + 3, y * cell + 3, x * cell + cell - 3, y * cell + cell - 3], 8, fill=c) f1, f2 = font(22), font(18) X = W + 20 d.text((X, 20), title, fill=(240, 240, 240), font=f1) d.text((X, 52), subtitle, fill=(150, 160, 175), font=f2) d.text((X, 100), f"apples {self.score}", fill=(235, 70, 70), font=f1) d.text((X, 132), f"moves {self.steps}", fill=(200, 200, 200), font=f1) col = (90, 220, 120) if lat_ms < 80 else (240, 200, 60) d.text((X, 180), f"decision {lat_ms:5.1f} ms", fill=col, font=f1) y0 = 230 if probs: d.text((X, y0), bar_label, fill=(150, 160, 175), font=f2); y0 += 30 for k, v in sorted(probs.items(), key=lambda kv: -kv[1]): d.text((X, y0), f"{k:>5s}", fill=(220, 220, 220), font=f2) d.rectangle([X + 70, y0 + 4, X + 70 + int(220 * v), y0 + 18], fill=(90, 160, 240)) d.text((X + 296, y0), f"{v:.2f}", fill=(180, 180, 180), font=f2) y0 += 30 return im def snake_demo(jev, fc, max_steps=300, seed=3, fps=12, name="snake"): g = Snake(seed=seed) frames, lats = [], [] q = g.question() jev.decide(g.state_text(), [q], fc=fc, graph=True) # warm the graph bucket while g.alive and g.steps < max_steps: st = g.state_text() torch.cuda.synchronize(); t0 = time.perf_counter() r = jev.decide(st, [q], fc=fc, graph=True)[0] torch.cuda.synchronize(); ms = (time.perf_counter() - t0) * 1000 lats.append(ms) frames.append(g.render(ms, r["probabilities"])) g.step(r["choice"]) frames.append(g.render(lats[-1] if lats else 0, {})) d = f"{OUT}/{name}_frames"; os.makedirs(d, exist_ok=True) for i, f in enumerate(frames): f.save(f"{d}/{i:05d}.png") subprocess.run(f"ffmpeg -y -loglevel error -framerate {fps} -i {d}/%05d.png -c:v libx264 -pix_fmt yuv420p " f"-vf 'scale=trunc(iw/2)*2:trunc(ih/2)*2' {OUT}/{name}.mp4", shell=True) return dict(apples=g.score, moves=g.steps, died=not g.alive, p50_ms=float(np.median(lats)), p95_ms=float(np.percentile(lats, 95)), max_ms=float(np.max(lats))) # ------------------------------------------------------------------------------------------ cards def _card(title, media_im, answers, lat_ms, subtitle=""): W = 1280 im = Image.new("RGB", (W, 620), (250, 250, 252)) d = ImageDraw.Draw(im) if media_im is not None: m = media_im.copy(); m.thumbnail((560, 560)) im.paste(m, (30, 30)) X = 620 d.text((X, 30), title, fill=(20, 20, 30), font=font(28)) if subtitle: d.text((X, 70), subtitle, fill=(110, 110, 125), font=font(18)) y = 115 for q, a in answers: d.text((X, y), q[:62], fill=(40, 40, 60), font=font(18)); y += 28 d.text((X + 20, y), a[:60], fill=(20, 110, 60), font=font(20)); y += 40 d.text((X, 570), f"{len(answers)} questions, one forward pass: {lat_ms:.0f} ms (towers included)", fill=(90, 90, 110), font=font(18)) return im def _fmt(q, r): if q["type"] == "noul": return f"noul P(yes) = {r['noul']:.2f}" if q["type"] == "score": top = max(r["probabilities"].items(), key=lambda kv: kv[1]) return f"score level {r['level']} ({q['criteria'][r['level']]}), p={top[1]:.2f}" return f"choice {r['choice']} (p={r['probabilities'][r['choice']]:.2f})" def _timed(jev, st, qs, fc): """decide() on a question list returns {index: answer}; hand back a list in question order.""" jev.decide(st, qs, fc=fc, graph=True) torch.cuda.synchronize(); t0 = time.perf_counter() r = jev.decide(st, qs, fc=fc, graph=True) torch.cuda.synchronize() return [r[i] for i in range(len(qs))], (time.perf_counter() - t0) * 1000 def image_demo(jev, fc, img, qs, name="demo_image", title="Image state"): r, ms = _timed(jev, [Seg("image", img)], qs, fc) _card(title, img, [(q["instructions"], _fmt(q, x)) for q, x in zip(qs, r)], ms, "768x768 input, 64 visual tokens").save(f"{OUT}/{name}.png") return r, ms def audio_demo(jev, fc, wav, qs, name="demo_audio", title="Audio state"): import matplotlib; matplotlib.use("Agg"); import matplotlib.pyplot as plt fig, ax = plt.subplots(figsize=(5.6, 3.0), dpi=100) t = np.arange(len(wav)) / 16000 ax.plot(t, wav, lw=0.4, color="#3b6fd8"); ax.set_xlabel("seconds"); ax.set_yticks([]) fig.tight_layout(); fig.canvas.draw() wim = Image.frombuffer("RGBA", fig.canvas.get_width_height(), fig.canvas.buffer_rgba()).convert("RGB") plt.close(fig) r, ms = _timed(jev, [Seg("audio", wav)], qs, fc) _card(title, wim, [(q["instructions"], _fmt(q, x)) for q, x in zip(qs, r)], ms, "16 kHz waveform, 6.25 tokens / s").save(f"{OUT}/{name}.png") return r, ms def video_demo(jev, fc, frames, qs, audio=None, fps=2.0, name="demo_video", title="Video state"): r, ms = _timed(jev, [Seg("video", frames, audio=audio, fps=fps)], qs, fc) ans = [(q["instructions"], _fmt(q, x)) for q, x in zip(qs, r)] d = f"{OUT}/{name}_frames"; os.makedirs(d, exist_ok=True) for i, fr in enumerate(frames): _card(title, fr, ans, ms, f"{len(frames)} frames @768, 16 tokens/frame" + (" + soundtrack" if audio is not None else "")).save(f"{d}/{i:03d}.png") subprocess.run(f"ffmpeg -y -loglevel error -framerate {fps} -i {d}/%03d.png -c:v libx264 -pix_fmt yuv420p {OUT}/{name}.mp4", shell=True) return r, ms # ------------------------------------------------------------------------------------------ snake, decomposed questions FREE_THR = 0.6 # chosen on held-out states (seeds 60-99), not on the demo games def snake_questions(g): qs = [{"type": "noul", "instructions": f"According to the state, is the next cell {d} of the head free?"} for d in DIRS] qs.append({"type": "choice", "instructions": "Which single move brings the snake's head closest to the food?", "criteria": {"up": "y - 1", "down": "y + 1", "left": "x - 1", "right": "x + 1"}}) return qs def snake_policy(g, ans): """Everything is judged by the model in ONE pass (4 noul + 1 choice); this only combines the answers.""" p_hit = {d: 1.0 - ans[i]["noul"] for i, d in enumerate(DIRS)} # noul = P(cell is free) p_close = ans[4]["probabilities"] cand = [d for d in DIRS if d != OPP[g.dir]] safe = [d for d in cand if p_hit[d] <= 1 - FREE_THR] # P(free) >= FREE_THR pick = max(safe, key=lambda d: p_close[d]) if safe else min(cand, key=lambda d: p_hit[d]) return pick, {d: round(p_close[d] * (1 - p_hit[d]), 3) for d in DIRS} def snake_demo2(jev, fc, max_steps=300, seed=3, fps=12, name="snake2"): g = Snake(seed=seed) frames, lats = [], [] jev.decide(g.state_text(), snake_questions(g), fc=fc, graph=True) while g.alive and g.steps < max_steps: qs = snake_questions(g) torch.cuda.synchronize(); t0 = time.perf_counter() r = jev.decide(g.state_text(), qs, fc=fc, graph=True) torch.cuda.synchronize(); ms = (time.perf_counter() - t0) * 1000 ans = [r[i] for i in range(len(qs))] move, shown = snake_policy(g, ans) lats.append(ms) frames.append(g.render(ms, shown)) g.step(move) frames.append(g.render(lats[-1] if lats else 0, {})) d = f"{OUT}/{name}_frames"; os.makedirs(d, exist_ok=True) for i, f in enumerate(frames): f.save(f"{d}/{i:05d}.png") subprocess.run(f"ffmpeg -y -loglevel error -framerate {fps} -i {d}/%05d.png -c:v libx264 -pix_fmt yuv420p " f"-vf 'scale=trunc(iw/2)*2:trunc(ih/2)*2' {OUT}/{name}.mp4", shell=True) return dict(apples=g.score, moves=g.steps, died=not g.alive, p50_ms=float(np.median(lats)), p95_ms=float(np.percentile(lats, 95))) # ------------------------------------------------------------------------------------------ snake, Open-Jev snake-v1 format # The public Open-Jev snake-v1 records (trained on) use a JSON board state, one `choice` (move) and one collision `noul` # per legal move. The demo asks exactly those 4 questions in ONE pass; the collision answers veto unsafe moves. COLL_Q = "Will moving {d} cause an immediate wall or body collision? The tail vacates unless food is eaten." COLL_THR = 0.5 # re-chosen on held-out seeds 60-71 by tune_snake_native(), never on the demo games def native_state(g): return json.dumps({"coordinates": "x increases right; y increases down; walls are outside the board", "direction": g.dir, "food": list(g.food), "game": "snake", "height": g.n, "snake_head_first": [list(c) for c in g.body], "width": g.n}, sort_keys=True) def native_questions(g): legal = [d for d in ("up", "right", "down", "left") if d != OPP[g.dir]] qs = [{"type": "choice", "instructions": "Choose a direction to keep the snake alive and collect food. The tail vacates " "on a move unless food is eaten; reversing direction is forbidden.", "criteria": {d: "" for d in legal}}] qs += [{"type": "noul", "instructions": COLL_Q.format(d=d)} for d in legal] return legal, qs def native_policy(legal, ans, thr): p_move = ans[0]["probabilities"] p_col = {d: ans[1 + i]["noul"] for i, d in enumerate(legal)} # noul = P(yes, collision) safe = [d for d in legal if p_col[d] < thr] pick = max(safe, key=lambda d: p_move[d]) if safe else min(legal, key=lambda d: p_col[d]) return pick, {d: round(float(p_move[d]), 3) for d in legal} def play_native(jev, fc, n=10, seed=0, max_steps=300, thr=None, frames=None): thr = COLL_THR if thr is None else thr g = Snake(n=n, seed=seed) lats = [] while g.alive and g.steps < max_steps: legal, qs = native_questions(g) torch.cuda.synchronize(); t0 = time.perf_counter() r = jev.decide(native_state(g), qs, fc=fc, graph=True) torch.cuda.synchronize(); ms = (time.perf_counter() - t0) * 1000 move, shown = native_policy(legal, r, thr) lats.append(ms) if frames is not None: frames.append(g.render(ms, shown, title="OmniJev plays Snake", subtitle="JSON state -> 4 questions, 1 pass", bar_label="P(move)")) g.step(move) if frames is not None: frames.append(g.render(lats[-1] if lats else 0, {}, title="OmniJev plays Snake", subtitle="JSON state -> 4 questions, 1 pass")) return g, lats def tune_snake_native(jev, fc, n=10, seeds=range(60, 72), thrs=(0.3, 0.5, 0.7, 0.9), max_steps=150): """Pick the collision-veto threshold on held-out games (apples first, then moves survived).""" global COLL_THR score = {} for t in thrs: gs = [play_native(jev, fc, n=n, seed=s, max_steps=max_steps, thr=t)[0] for s in seeds] score[t] = (sum(g.score for g in gs), sum(g.steps for g in gs), sum(not g.alive for g in gs)) COLL_THR = max(thrs, key=lambda t: score[t][:2]) return COLL_THR, {str(k): v for k, v in score.items()} def snake_demo_native(jev, fc, n=10, seed=3, max_steps=300, fps=18, name="snake_native"): legal, qs = native_questions(Snake(n=n, seed=seed)) jev.decide(native_state(Snake(n=n, seed=seed)), qs, fc=fc, graph=True) # warm the graph bucket frames = [] g, lats = play_native(jev, fc, n=n, seed=seed, max_steps=max_steps, frames=frames) d = f"{OUT}/{name}_frames"; os.makedirs(d, exist_ok=True) for i, f in enumerate(frames): f.save(f"{d}/{i:05d}.png") subprocess.run(f"ffmpeg -y -loglevel error -framerate {fps} -i {d}/%05d.png -c:v libx264 -pix_fmt yuv420p " f"-vf 'scale=trunc(iw/2)*2:trunc(ih/2)*2' {OUT}/{name}.mp4", shell=True) return dict(board=n, apples=g.score, moves=g.steps, died=not g.alive, thr=COLL_THR, p50_ms=float(np.median(lats)), p95_ms=float(np.percentile(lats, 95)))