Buckets:
| """Visualize AAM memory selection: per-episode gate-weight heatmaps aligned to rollout videos. | |
| Reads a memdump JSONL (one record per policy call: reset/t_hist/steps_back/weight), | |
| splits into episodes by `reset` markers, and renders one heatmap PNG per episode. | |
| Rows = the 8 log-window slots (deepest history -> current frame); columns = policy | |
| calls over the episode; color = gate weight (how strongly that slot is kept). | |
| Usage: | |
| python scripts/viz_aam_memory.py <task_dir> <memdump.jsonl> <out_dir> | |
| where <task_dir> holds the rollout mp4s named ..._episode_N-{success|failure}.mp4. | |
| """ | |
| import json | |
| import re | |
| import sys | |
| from pathlib import Path | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| LOGK = [-1024, -512, -256, -128, -64, -32, -16, 0] # nominal env-step offsets per slot | |
| def split_episodes(records): | |
| eps, cur = [], [] | |
| for r in records: | |
| if r["reset"] and cur: | |
| eps.append(cur) | |
| cur = [] | |
| cur.append(r) | |
| if cur: | |
| eps.append(cur) | |
| return eps | |
| def episode_result(task_dir: Path, ep_idx: int): | |
| for mp4 in task_dir.glob(f"*episode_{ep_idx}-*.mp4"): | |
| return "success" if "success" in mp4.name else "failure", mp4 | |
| return "unknown", None | |
| def render(task: str, ep_idx: int, records, result: str, out_png: Path): | |
| W = np.array([r["weight"] for r in records]).T # (8 slots, n_calls) | |
| n_calls = W.shape[1] | |
| # demo/exec boundary: prime calls are those before the first non-reset stretch stops | |
| # growing t_hist past the demo; simplest cue = mark where reset happened (call 0). | |
| fig, ax = plt.subplots(figsize=(max(6, n_calls * 0.14), 4.2)) | |
| im = ax.imshow(W, aspect="auto", cmap="magma", vmin=0, vmax=max(0.35, W.max()), | |
| interpolation="nearest") | |
| ax.set_yticks(range(8)) | |
| ax.set_yticklabels([f"{o:+d} steps" if o != 0 else "current" for o in LOGK]) | |
| ax.set_xlabel("policy call (time →)") | |
| ax.set_ylabel("memory slot (log-window offset)") | |
| color = "#2ca02c" if result == "success" else "#d62728" | |
| ax.set_title(f"{task} · episode {ep_idx} · {result.upper()}\n" | |
| f"gate weight per memory slot over the rollout", color=color, fontsize=11) | |
| cb = fig.colorbar(im, ax=ax, fraction=0.03, pad=0.01) | |
| cb.set_label("gate weight (0 = discarded, high = kept)") | |
| fig.tight_layout() | |
| fig.savefig(out_png, dpi=110) | |
| plt.close(fig) | |
| def main(): | |
| task_dir = Path(sys.argv[1]) | |
| memdump = Path(sys.argv[2]) | |
| out_dir = Path(sys.argv[3]) | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| task = task_dir.name | |
| records = [json.loads(l) for l in memdump.read_text().splitlines() if l.strip()] | |
| episodes = split_episodes(records) | |
| print(f"{task}: {len(records)} calls -> {len(episodes)} episodes") | |
| manifest = [] | |
| for i, ep in enumerate(episodes): | |
| result, mp4 = episode_result(task_dir, i) | |
| png = out_dir / f"{task}_episode_{i}_{result}_memheatmap.png" | |
| render(task, i, ep, result, png) | |
| manifest.append({ | |
| "task": task, "episode": i, "result": result, | |
| "n_calls": len(ep), | |
| "heatmap": png.name, | |
| "video": mp4.name if mp4 else None, | |
| "mean_gate_current": round(float(np.mean([r["weight"][-1] for r in ep])), 3), | |
| "mean_gate_deep": round(float(np.mean([r["weight"][0] for r in ep])), 3), | |
| }) | |
| print(f" ep{i} [{result}] -> {png.name}") | |
| (out_dir / f"{task}_manifest.json").write_text(json.dumps(manifest, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 3.59 kB
- Xet hash:
- f8308369bac648ae5299ebeb5e9bf21dcb7cc33f7fc497c34a1edf3873acf073
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.