File size: 16,116 Bytes
8c867d9
 
 
 
 
08e3ea8
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
08e3ea8
 
 
 
8c867d9
 
 
 
 
 
 
 
 
 
 
08e3ea8
 
8c867d9
 
1eac9b8
8c867d9
 
08e3ea8
 
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1eac9b8
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1eac9b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
08e3ea8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
"""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)))