Download code/demos.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 16.1 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/demos.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/demos.py
-
curl -L -o demos.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/demos.py
16.1 kB
| """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))) | |