omnijev-work / code /demos.py
fnruha0921's picture
stage3: demos.py
08e3ea8 verified
Raw History Blame Contribute Delete
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)))