Return list of per-step diffs between two policy trajectories.
(nat, wrp)
| 132 | |
| 133 | |
| 134 | def _diff_frames(nat, wrp): |
| 135 | """Return list of per-step diffs between two policy trajectories.""" |
| 136 | diffs = [] |
| 137 | fn, fw = nat["frames"], wrp["frames"] |
| 138 | if len(fn) != len(fw): |
| 139 | diffs.append(f"length {len(fn)} vs {len(fw)}") |
| 140 | for i, (a, b) in enumerate(zip(fn, fw)): |
| 141 | for k in ("obs", "rew", "term", "trunc", "success", "instr"): |
| 142 | if a[k] != b[k]: |
| 143 | diffs.append(f"step{i}.{k}: {a[k]} != {b[k]}") |
| 144 | # action vector: exact equality (policy is deterministic given identical obs) |
| 145 | if a["act"] != b["act"]: |
| 146 | diffs.append(f"step{i}.act") |
| 147 | return diffs |
| 148 | |
| 149 | |
| 150 | def check_task(task, *, seed, steps, init_rng): |