Buckets:
| import json | |
| import sys | |
| from pathlib import Path | |
| import av | |
| import imageio | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| from matplotlib.patches import Rectangle | |
| import numpy as np | |
| import torch | |
| import run_planner as R | |
| OUT = R.ROOT / "viz/out" | |
| COLORS = {"pick": "#1f77b4", "place": "#d62728", "table": "#8c564b", "press": "#ff7f0e", "move": "#9467bd", | |
| "static": "#7f7f7f", "insert": "#2ca02c", "none": "#d9d9d9"} | |
| KC = [COLORS[k] for k in R.KINDS] | |
| EXAMPLES = [("VideoPlaceOrder", 214), ("PickXtimes", 100), ("VideoUnmaskSwap", 5), ("StopCube", 1206), | |
| ("ButtonUnmaskSwap", 501), ("BinFill", 901), ("VideoRepick", 608), ("InsertPeg", 1305)] | |
| def frames_at(i, wanted, key="image"): | |
| wanted = sorted(set(int(f) for f in wanted)) | |
| out, s = {}, set(wanted) | |
| for t, frame in enumerate(av.open(str(R.video(i, key))).decode(video=0)): | |
| if t in s: | |
| out[t] = frame.to_ndarray(format="rgb24") | |
| if t >= wanted[-1]: | |
| break | |
| return out | |
| def px(point): | |
| return point[..., 1] * R.SPAN + R.MARGIN, point[..., 0] * R.SPAN + R.MARGIN | |
| def runs(seq): | |
| out, start = [], 0 | |
| for j in range(1, len(seq) + 1): | |
| if j == len(seq) or seq[j] != seq[start]: | |
| out.append((start, j, int(seq[start]))) | |
| start = j | |
| return out | |
| def band(ax, seq, y, h=0.8): | |
| for a, b, k in runs(seq): | |
| ax.add_patch(Rectangle((a, y), b - a, h, color=KC[k], lw=0)) | |
| def timeline(task, i, model, device): | |
| item, ent = R.load(i) | |
| instruction, subgoals = R.episode_text(i) | |
| o = R.run(model, item, ent, device) | |
| frames = item["frames"] | |
| t = R.truth(i, frames) | |
| n = len(frames) | |
| demo = int((frames < item["exec_start"]).sum()) | |
| kp, kt = o["kind"].argmax(-1).numpy(), t["kind"].numpy() | |
| op, ot = o["ordinal"].argmax(-1).numpy(), t["ordinal"].numpy() | |
| has = (t["point"] >= 0).all(-1) | |
| err = ((o["point"] - t["point"]).norm(dim=-1) * R.SPAN).numpy() | |
| err[~has.numpy() | (o["point"] < 0).any(-1).numpy()] = np.nan | |
| picks = np.linspace(0, n - 1, 8).round().astype(int) | |
| if demo: | |
| picks = np.unique(np.concatenate([np.linspace(0, demo - 1, 5).round(), np.linspace(demo, n - 1, 3).round()]).astype(int)) | |
| images = frames_at(i, frames[picks].tolist()) | |
| fig = plt.figure(figsize=(16, 10.5)) | |
| gs = fig.add_gridspec(5, len(picks), height_ratios=[2.2, 0.9, 0.9, 1.6, 1.2], hspace=0.45) | |
| for c, j in enumerate(picks): | |
| ax = fig.add_subplot(gs[0, c]) | |
| img = images[int(frames[j])] | |
| ax.imshow(img) | |
| if has[j]: | |
| x, y = px(t["point"][j]) | |
| ax.plot(x, y, "o", ms=9, mfc="none", mec="lime", mew=2) | |
| if (o["point"][j] >= 0).all(): | |
| x, y = px(o["point"][j]) | |
| ax.plot(x, y, "x", ms=9, color="red", mew=2) | |
| ax.set_title(f"node {j} frame {int(frames[j])}\n{'demo' if j < demo else 'execution'}", fontsize=8) | |
| ax.axis("off") | |
| ax = fig.add_subplot(gs[1, :]) | |
| band(ax, kt, 1) | |
| band(ax, kp, 0) | |
| ax.set_xlim(0, n) | |
| ax.set_ylim(0, 1.8) | |
| ax.set_yticks([0.4, 1.4], ["predicted", "label"]) | |
| ax.axvline(demo, color="k", lw=1.5) | |
| ax.set_title("step kind per node (colour) : " + ", ".join(f"{k}" for k in R.KINDS), fontsize=9) | |
| handles = [Rectangle((0, 0), 1, 1, color=COLORS[k]) for k in R.KINDS] | |
| ax.legend(handles, R.KINDS, ncol=8, fontsize=7, loc="upper center", bbox_to_anchor=(0.5, -0.25), frameon=False) | |
| ax = fig.add_subplot(gs[2, :]) | |
| ax.step(range(n), ot, where="post", color="k", lw=2, label="label") | |
| ax.step(range(n), op, where="post", color="tab:red", lw=1, ls="--", label="predicted") | |
| ax.set_yticks(range(6), R.ORDINALS, fontsize=7) | |
| ax.set_xlim(0, n) | |
| ax.axvline(demo, color="k", lw=1.5) | |
| ax.set_title("ordinal ('for the n-th time')", fontsize=9) | |
| ax.legend(fontsize=7, loc="upper left") | |
| ax = fig.add_subplot(gs[3, :]) | |
| steps = np.zeros((n, 7)) | |
| for a, b, k in runs(kt): | |
| if k < 7: | |
| steps[a:, k] += 1 | |
| for k in range(7): | |
| if steps[-1, k] > 0 or o["counts"][-1, k] > 0.5: | |
| ax.plot(o["counts"][:, k], color=KC[k], lw=2, label=f"{R.KINDS[k]} (counter)") | |
| ax.plot(steps[:, k], color=KC[k], lw=1, ls=":") | |
| ax.set_xlim(0, n) | |
| ax.axvline(demo, color="k", lw=1.5) | |
| ax.set_ylabel("count") | |
| ax.set_title("event counter: running count per kind (solid = learned prefix sum, dotted = steps in the labels)", fontsize=9) | |
| ax.legend(fontsize=7, ncol=4, loc="upper left") | |
| ax = fig.add_subplot(gs[4, :]) | |
| ax.plot(err, ".", ms=3, color="tab:red") | |
| ax.axhline(8, color="gray", ls="--", lw=1) | |
| ax.set_xlim(0, n) | |
| ax.set_ylim(0, max(30, np.nanpercentile(err, 98) if np.isfinite(err).any() else 30)) | |
| ax.axvline(demo, color="k", lw=1.5) | |
| ax.set_ylabel("px") | |
| ax.set_xlabel("node (one node every 4 frames; vertical line = end of demo, start of execution)") | |
| ax.set_title("target point error of the planner (pixels in the 256 px image, dashed = 8 px)", fontsize=9) | |
| fig.suptitle(f"{task}, val episode {i}: \"{instruction}\"\nfilmstrip: green circle = label target, red x = planner target", fontsize=11) | |
| path = OUT / "02_planner_timeline" / f"{task}_{i}.png" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(path, dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| return o, item, ent, t, demo | |
| def node_anatomy(i, model, device): | |
| item, ent = R.load(i) | |
| instruction, _ = R.episode_text(i) | |
| frames = item["frames"] | |
| j = len(frames) - 10 | |
| f = int(frames[j]) | |
| front = frames_at(i, [f])[f] | |
| wrist = frames_at(i, [f], "wrist_image")[f] | |
| pos = ent["position"][j, :192].float().numpy() | |
| vis = ent["visible"][j, :192].numpy() | |
| fig, axes = plt.subplots(1, 4, figsize=(18, 5), gridspec_kw={"width_ratios": [1, 1, 1, 0.9]}) | |
| ax = axes[0] | |
| ax.imshow(front) | |
| side = R.SPAN / 9 | |
| for k in range(10): | |
| ax.axhline(R.MARGIN + k * side, color="w", lw=0.6) | |
| ax.axvline(R.MARGIN + k * side, color="w", lw=0.6) | |
| for p in range(81): | |
| ax.text(R.MARGIN + (p % 9 + 0.5) * side, R.MARGIN + (p // 9 + 0.5) * side, str(p), color="yellow", fontsize=6, ha="center", va="center") | |
| ax.set_title("81 front tokens: Eagle patch features on a fixed 9x9 grid\n(slot p is always the same place on the table)", fontsize=9) | |
| ax.axis("off") | |
| axes[1].imshow(wrist) | |
| axes[1].set_title("token 82: mean of the 81 wrist-camera patches", fontsize=9) | |
| axes[1].axis("off") | |
| ax = axes[2] | |
| ax.imshow(front) | |
| ax.scatter(pos[vis, 1], pos[vis, 0], s=10, c="cyan", edgecolors="k", linewidths=0.3) | |
| ax.scatter(pos[~vis, 1], pos[~vis, 0], s=10, c="none", edgecolors="r", linewidths=0.6) | |
| ax.set_title("192 entity tokens: CoTracker points seeded on the foreground\n(cyan visible, red outline occluded)", fontsize=9) | |
| ax.axis("off") | |
| ax = axes[3] | |
| pick = [int(np.argmin(((pos - c) ** 2).sum(-1))) for c in np.array([[90, 90], [120, 160], [160, 120], [100, 130]])] | |
| patches = ent["patch"][j, pick, 0].numpy() | |
| first = ent["patch"][0, pick, 0].numpy() | |
| ax.axis("off") | |
| ax.set_title("each entity token sees a 12x12 patch now\nand at the first node, plus its motion", fontsize=9) | |
| for r, (a, b) in enumerate(zip(patches, first)): | |
| sub = ax.inset_axes([0.05, 0.78 - r * 0.25, 0.4, 0.2]) | |
| sub.imshow(b) | |
| sub.axis("off") | |
| sub.set_title(f"entity {pick[r]}, first node", fontsize=6) | |
| sub = ax.inset_axes([0.55, 0.78 - r * 0.25, 0.4, 0.2]) | |
| sub.imshow(a) | |
| sub.axis("off") | |
| sub.set_title(f"node {j}", fontsize=6) | |
| fig.suptitle(f"What one graph node holds (VideoPlaceOrder val episode {i}, frame {f}). The instruction tokens \"{instruction}\" are read by every token through query edges.", fontsize=10) | |
| path = OUT / "01_graph_node" / "node_anatomy.png" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(path, dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| def edges(i): | |
| item, ent = R.load(i) | |
| frames = item["frames"] | |
| demo = int((frames < item["exec_start"]).sum()) | |
| slot = 4 * 9 + 6 | |
| nodes = np.linspace(0, len(frames) - 1, 10).round().astype(int) | |
| images = frames_at(i, frames[nodes].tolist()) | |
| side = R.SPAN / 9 | |
| r0, c0 = R.MARGIN + (slot // 9) * side, R.MARGIN + (slot % 9) * side | |
| fig, axes = plt.subplots(3, 10, figsize=(18, 6.5)) | |
| for c, j in enumerate(nodes): | |
| img = images[int(frames[j])] | |
| ax = axes[0, c] | |
| ax.imshow(img[int(r0) - 10 : int(r0 + side) + 10, int(c0) - 10 : int(c0 + side) + 10]) | |
| ax.set_title(f"node {j}{' (exec)' if j >= demo else ''}", fontsize=7) | |
| ax.axis("off") | |
| pos = ent["position"][:, :192].float().numpy() | |
| k = int(np.argmax(np.sqrt((np.diff(pos, axis=0) ** 2).sum(-1)).sum(0))) | |
| for c, j in enumerate(nodes): | |
| img = images[int(frames[j])] | |
| y, x = pos[j, k] | |
| ax = axes[1, c] | |
| ax.imshow(img[max(0, int(y) - 24) : int(y) + 24, max(0, int(x) - 24) : int(x) + 24]) | |
| ax.axis("off") | |
| j = nodes[-1] | |
| img = images[int(frames[j])] | |
| for c in range(10): | |
| axes[2, c].axis("off") | |
| ax = axes[2, 0].inset_axes([0, 0, 3.2, 1]) | |
| ax.imshow(img) | |
| d = np.sqrt(((pos[j] - pos[j, k]) ** 2).sum(-1)) / R.SPAN | |
| w = np.exp(-(d ** 2) / (2 * 0.08 ** 2)) | |
| ax.scatter(pos[j, :, 1], pos[j, :, 0], s=6 + 60 * w, c=w, cmap="viridis", edgecolors="k", linewidths=0.2) | |
| ax.plot(pos[j, k, 1], pos[j, k, 0], "*", ms=14, color="red") | |
| ax.axis("off") | |
| axes[0, 0].text(-0.15, 0.5, f"temporal edge\n(frame slot {slot})", transform=axes[0, 0].transAxes, ha="right", va="center", fontsize=9) | |
| axes[1, 0].text(-0.15, 0.5, f"temporal edge\n(entity {k})", transform=axes[1, 0].transAxes, ha="right", va="center", fontsize=9) | |
| ax.set_title("spatial edge between entities: the red entity reads all others in the same node,\nweighted by a learned Gaussian of distance (one width per head; shown: 0.08 of the image)", fontsize=8) | |
| fig.suptitle("Edges of the graph on a real episode. Row 1: one front slot (same place on the table) across nodes; this is what a temporal edge of a frame token reads.\n" | |
| "Row 2: one tracked entity across nodes (the crop follows the track); this is what a temporal edge of an entity token reads.", fontsize=10) | |
| path = OUT / "01_graph_node" / "edges.png" | |
| fig.savefig(path, dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| def query_swap(i, model, device): | |
| item, ent = R.load(i) | |
| instruction, subgoals = R.episode_text(i) | |
| bank = {} | |
| for e in range(200, 300): | |
| bank.setdefault(R.episode_text(e)[0], e) | |
| frames = item["frames"] | |
| demo = int((frames < item["exec_start"]).sum()) | |
| drops = [] | |
| prev = None | |
| for f, s in enumerate(subgoals[: item["exec_start"]]): | |
| if s != prev and s.startswith("drop the cube onto target at"): | |
| drops.append((f, [float(v) for v in s.split("<")[1].rstrip(">").split(",")])) | |
| prev = s | |
| kinds = R.truth(i, frames)["kind"].numpy() | |
| j = int(np.nonzero((np.arange(len(frames)) >= demo) & (kinds == R.KINDS.index("place")))[0][0]) | |
| f = int(frames[j]) | |
| img = frames_at(i, [f])[f] | |
| words = ["first", "second", "third", "fourth"] | |
| match = R.ORDINAL_WORD.search(instruction) | |
| fig, axes = plt.subplots(1, 4, figsize=(18, 5)) | |
| for c, w in enumerate(words): | |
| text = instruction[: match.start()] + w + instruction[match.end():] | |
| ax = axes[c] | |
| ax.imshow(img) | |
| for n, (_, (r, cc)) in enumerate(drops): | |
| ax.text(cc, r, str(n + 1), color="white", fontsize=12, weight="bold", ha="center", va="center", | |
| bbox=dict(boxstyle="circle", fc="black", alpha=0.6)) | |
| if text not in bank: | |
| ax.set_title(f"\"...on the {w} target...\"\n(no such instruction in the data)", fontsize=9) | |
| ax.axis("off") | |
| continue | |
| other = torch.load(R.ROOT / f"cache/online16_s4/ep_{bank[text]:04d}.pt", weights_only=False)["text"] | |
| o = R.run(model, item, ent, device, other) | |
| p = o["pointer"][j, :192].numpy() | |
| pos = o["position"][j].numpy() * R.SPAN + R.MARGIN | |
| order = np.argsort(p) | |
| ax.scatter(pos[order, 1], pos[order, 0], s=8 + 300 * p[order] / p.max(), c=p[order], cmap="magma", edgecolors="w", linewidths=0.3) | |
| x, y = px(o["point"][j]) | |
| ax.plot(x, y, "x", color="cyan", ms=14, mew=3) | |
| ax.set_title(f"\"...on the {w} target...\"\npredicted kind: {R.KINDS[int(o['kind'][j].argmax())]}", fontsize=9) | |
| ax.axis("off") | |
| fig.suptitle(f"Query dependence. Same video and same tracks (VideoPlaceOrder val episode {i}, execution frame {f}); only the ordinal word of the instruction changes.\n" | |
| "Numbers = where the cube was dropped in the demo, in order. Dots = pointer probability over the 192 entities, cyan x = planner target.", fontsize=10) | |
| path = OUT / "03_query_dependence" / f"VideoPlaceOrder_{i}_ordinal_swap.png" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(path, dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| def mark_fovea(cases): | |
| fig, axes = plt.subplots(len(cases), 4, figsize=(15, 3.9 * len(cases))) | |
| centres = (np.stack(np.meshgrid(np.arange(9), np.arange(9), indexing="ij"), -1).reshape(-1, 2) + 0.5) / 9 | |
| for r, (task, i, f) in enumerate(cases): | |
| _, subgoals = R.episode_text(i) | |
| lab = R.labels()[i] | |
| has = (lab["point"] >= 0).all(-1) | |
| if not has[f]: | |
| f = int(has.nonzero()[0]) | |
| point = lab["point"][f].float().numpy() | |
| img = frames_at(i, [f])[f] | |
| weight = np.exp(-((centres - point) ** 2).sum(-1) / (2 * 0.06 ** 2)).reshape(9, 9) | |
| y, x = point[0] * R.SPAN + R.MARGIN, point[1] * R.SPAN + R.MARGIN | |
| ax = axes[r, 0] | |
| ax.imshow(img) | |
| ax.plot(x, y, "o", ms=10, mfc="none", mec="lime", mew=2) | |
| ax.set_title(f"{task}: \"{subgoals[f][:60]}\"", fontsize=8) | |
| ax.axis("off") | |
| ax = axes[r, 1] | |
| ax.imshow(img) | |
| ax.imshow(np.kron(weight, np.ones((27, 27)))[: 243, : 243], extent=(R.MARGIN, R.MARGIN + R.SPAN, R.MARGIN + R.SPAN, R.MARGIN), alpha=0.6, cmap="hot", vmin=0, vmax=1) | |
| ax.set_title("mark weight on the 81 front tokens\nexp(-d^2 / 2*0.06^2)", fontsize=8) | |
| ax.axis("off") | |
| ax = axes[r, 2] | |
| half = 32 | |
| yy, xx = int(round(y)), int(round(x)) | |
| pad = np.pad(img, ((half, half), (half, half), (0, 0))) | |
| ax.imshow(pad[yy : yy + 2 * half, xx : xx + 2 * half]) | |
| ax.set_title("fovea: 64x64 crop at the target\n(read by a small CNN into 81 extra tokens)", fontsize=8) | |
| ax.axis("off") | |
| ax = axes[r, 3] | |
| ax.axis("off") | |
| kind = R.KINDS[int(lab["kind"][f])] | |
| ordinal = R.ORDINALS[int(lab["ordinal"][f])] | |
| ax.text(0, 0.9, "controller memory tokens", fontsize=10, weight="bold", va="top") | |
| ax.text(0, 0.72, f"step token: kind = {kind}, ordinal = {ordinal}\n + words of the subgoal text", fontsize=9, family="monospace", va="top") | |
| ax.text(0, 0.45, f"point token: Fourier((row, col)) = ({point[0]:.3f}, {point[1]:.3f})", fontsize=9, family="monospace", va="top") | |
| ax.text(0, 0.3, "goal token: fovea reader output", fontsize=9, family="monospace", va="top") | |
| ax.text(0, 0.12, "each DiT block: actions cross-attend these\ntokens -> FiLM (scale, shift) on its FFN", fontsize=9, va="top") | |
| fig.suptitle("How a subgoal reaches the controller: the target point is drawn into the image tokens (mark), cropped at full resolution (fovea), and encoded as memory tokens", fontsize=11) | |
| path = OUT / "04_controller_conditioning" / "mark_fovea_tokens.png" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(path, dpi=100, bbox_inches="tight") | |
| plt.close(fig) | |
| def planner_video(task, i, o, item, ent, t, demo): | |
| instruction, subgoals = R.episode_text(i) | |
| frames = item["frames"] | |
| images = frames_at(i, frames.tolist()) | |
| wrist = frames_at(i, frames.tolist(), "wrist_image") | |
| out = [] | |
| for j, f in enumerate(frames.tolist()): | |
| fig = plt.figure(figsize=(9.6, 5.6), dpi=100) | |
| ax = fig.add_axes([0.0, 0.18, 0.5, 0.7]) | |
| ax.imshow(images[f]) | |
| p = o["pointer"][j, :192].numpy() | |
| pos = o["position"][j].numpy() * R.SPAN + R.MARGIN | |
| order = np.argsort(p) | |
| ax.scatter(pos[order, 1], pos[order, 0], s=4 + 160 * p[order] / max(p.max(), 1e-6), c=p[order], cmap="magma", edgecolors="none", alpha=0.85) | |
| if (t["point"][j] >= 0).all(): | |
| x, y = px(t["point"][j]) | |
| ax.plot(x, y, "o", ms=12, mfc="none", mec="lime", mew=2) | |
| if (o["point"][j] >= 0).all(): | |
| x, y = px(o["point"][j]) | |
| ax.plot(x, y, "x", ms=12, color="cyan", mew=3) | |
| ax.axis("off") | |
| ax2 = fig.add_axes([0.5, 0.48, 0.25, 0.4]) | |
| ax2.imshow(wrist[f]) | |
| ax2.axis("off") | |
| ax3 = fig.add_axes([0.77, 0.48, 0.22, 0.4]) | |
| counts = o["counts"][j].numpy() | |
| ax3.barh(range(7), counts, color=KC[:7]) | |
| ax3.set_yticks(range(7), R.KINDS[:7], fontsize=7) | |
| ax3.set_xlim(0, max(6, float(o["counts"][-1].max()) + 0.5)) | |
| ax3.invert_yaxis() | |
| ax3.set_title("event counter", fontsize=8) | |
| ax3.tick_params(labelsize=7) | |
| kp, kt = int(o["kind"][j].argmax()), int(t["kind"][j]) | |
| op, ot = int(o["ordinal"][j].argmax()), int(t["ordinal"][j]) | |
| phase = "DEMO (watching)" if j < demo else "EXECUTION (robot acting)" | |
| fig.text(0.01, 0.97, f"{task} ep {i} node {j}/{len(frames)} frame {f} {phase}", fontsize=10, weight="bold", va="top") | |
| fig.text(0.01, 0.92, f"instruction: {instruction[:120]}", fontsize=8, va="top") | |
| fig.text(0.51, 0.40, f"label : {R.KINDS[kt]:6s} ordinal {R.ORDINALS[ot]}", fontsize=9, family="monospace", color="green") | |
| fig.text(0.51, 0.34, f"planner : {R.KINDS[kp]:6s} ordinal {R.ORDINALS[op]}", fontsize=9, family="monospace", color="red" if (kp, op) != (kt, ot) else "black") | |
| fig.text(0.51, 0.26, f"subgoal text: {subgoals[f][:52]}", fontsize=8, va="top") | |
| fig.text(0.01, 0.12, "dots = planner pointer probability over 192 tracked entities; green circle = label target; cyan x = planner target", fontsize=8) | |
| fig.canvas.draw() | |
| out.append(np.asarray(fig.canvas.buffer_rgba())[..., :3].copy()) | |
| plt.close(fig) | |
| path = OUT / "05_planner_videos" / f"{task}_{i}.mp4" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| imageio.mimwrite(path, out, fps=8, codec="libx264") | |
| def results(): | |
| tasks = ["BinFill", "PickXtimes", "SwingXtimes", "StopCube", "VideoUnmask", "ButtonUnmask", "VideoUnmaskSwap", "ButtonUnmaskSwap", | |
| "PickHighlight", "VideoRepick", "VideoPlaceButton", "VideoPlaceOrder", "MoveCube", "InsertPeg", "PatternLock", "RouteStick"] | |
| rows = { | |
| "ours, non-oracle (50 ep)": [44, 60, 92, 36, 94, 88, 82, 50, 86, 76, 96, 94, 46, 22, 20, 24], | |
| "ours, oracle subgoal (20 ep)": [70, 80, 85, 95, 100, 85, 80, 85, 70, 85, 95, 100, 35, 35, 25, 35], | |
| "SimpleMemVLA": [78, 100, 96, 92, 100, 100, 90, 94, 90, 64, 86, 90, 94, 46, 94, 98], | |
| "FrameSamp+Modul": [39.6, 87.3, 92, 42, 32.7, 25.1, 24.4, 18.2, 22.9, 30.4, 60, 32, 77.8, 7.6, 53.6, 66.7], | |
| "pi0.5": [30, 42.9, 35.6, 6.7, 20.4, 22.2, 18.7, 6.7, 11.3, 0.4, 31.1, 25.8, 26, 1.6, 2.9, 4.7], | |
| } | |
| colors = ["#d62728", "#ff9896", "#1f77b4", "#7f7f7f", "#c7c7c7"] | |
| fig, ax = plt.subplots(figsize=(17, 5.5)) | |
| w = 0.16 | |
| for n, (name, v) in enumerate(rows.items()): | |
| ax.bar(np.arange(16) + (n - 2) * w, v, w, label=f"{name}: {np.mean(v):.1f}", color=colors[n]) | |
| for x in (3.5, 7.5, 11.5): | |
| ax.axvline(x, color="k", lw=0.8, ls=":") | |
| for x, name in zip((1.5, 5.5, 9.5, 13.5), ("Counting", "Permanence", "Reference", "Imitation")): | |
| ax.text(x, 104, name, ha="center", fontsize=10) | |
| ax.set_xticks(range(16), tasks, rotation=35, ha="right", fontsize=9) | |
| ax.set_ylim(0, 110) | |
| ax.set_ylabel("success rate (%)") | |
| ax.legend(fontsize=9, ncol=5, loc="upper center", bbox_to_anchor=(0.5, -0.32), frameon=False) | |
| ax.set_title("RoboMME, 16 tasks, official test split. Baseline numbers from the SimpleMemVLA paper, Table 8", fontsize=11) | |
| path = OUT / "06_results" / "robomme16_per_task.png" | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| fig.savefig(path, dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| fig, ax = plt.subplots(figsize=(9, 3.6)) | |
| names = ["ours (A100)", "HAMLET, GR00T N1.5 (A100)", "SimpleMemVLA (H100)"] | |
| hours = [31, 74, 2560] | |
| ax.barh(names, hours, color=["#d62728", "#2ca02c", "#1f77b4"]) | |
| ax.set_xscale("log") | |
| for y, h in enumerate(hours): | |
| ax.text(h * 1.1, y, f"{h} GPU-h", va="center", fontsize=10) | |
| ax.set_xlim(10, 10000) | |
| ax.set_xlabel("training GPU-hours (log scale)") | |
| ax.set_title("Training compute. HAMLET is on its own benchmarks, not RoboMME", fontsize=10) | |
| fig.savefig(OUT / "06_results" / "training_compute.png", dpi=110, bbox_inches="tight") | |
| plt.close(fig) | |
| def main(): | |
| device = torch.device("cuda", int(sys.argv[1]) if len(sys.argv) > 1 else 1) | |
| model = R.model(device) | |
| results() | |
| node_anatomy(214, model, device) | |
| edges(214) | |
| query_swap(214, model, device) | |
| mark_fovea([("VideoPlaceOrder", 214, 1200), ("PickXtimes", 100, 120), ("VideoUnmaskSwap", 5, 330), ("InsertPeg", 1305, 0), ("StopCube", 1206, 0)]) | |
| for task, i in EXAMPLES: | |
| o, item, ent, t, demo = timeline(task, i, model, device) | |
| planner_video(task, i, o, item, ent, t, demo) | |
| print("done", task, i, flush=True) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 21.7 kB
- Xet hash:
- f12a5c0d96de5cd38ec77d4d94b279095e3181ce4cc9eb40f565850e0be2d407
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.