twanghcmut/temporal-vla / viz_aam_memory.py
twanghcmut's picture
download
raw
3.59 kB
"""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.