Spaces:
Sleeping
Sleeping
File size: 6,713 Bytes
a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 662728a a3143c1 662728a a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 a3143c1 bf46f62 | 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 | """Generate the 'drift-recovery' plot.
For one chosen task, plot per-turn cumulative reward for N policies on
the same axes, averaged across multiple seeds with a ±1σ band. A
vertical red line marks the turn drift fires (modal across seeds).
Run this AFTER Track A delivers a trained policy adapter — on
baselines alone, cumulative reward gives random too much credit
through valid-frame bonuses, so rule_based / keyword / random do not
separate cleanly enough for a deck slide. With a trained policy in
the mix, the trained line dominates and the plot tells a strong story.
Usage (with trained policy):
python scripts/make_money_plot.py \\
--task research_photo_rename --seeds 0 1 2 \\
--policies rule_based keyword random trained:my_module:my_policy_fn \\
--out reports/drift_recovery.png
The format `module:fn` lets you point at any callable returning an
action dict from an obs dict.
"""
from __future__ import annotations
import argparse
import importlib
import os
import sys
from typing import Any, Callable, Dict, List, Tuple
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from arena.server.arena_env import ProtocolArenaEnvironment
from arena.models import OrchestratorAction
from arena.eval.baselines import rule_based_policy, keyword_policy, random_policy
PolicyFn = Callable[[Dict[str, Any]], Dict[str, Any]]
BUILTINS: Dict[str, PolicyFn] = {
"rule_based": rule_based_policy,
"keyword": keyword_policy,
"random": random_policy,
}
def _resolve(spec: str) -> Tuple[str, PolicyFn]:
if spec in BUILTINS:
return spec, BUILTINS[spec]
# format: label:module:fn
parts = spec.split(":")
if len(parts) != 3:
raise ValueError(f"bad policy spec {spec!r}; want 'label:module:fn' "
f"or one of {list(BUILTINS)}")
label, mod, fn = parts
return label, getattr(importlib.import_module(mod), fn)
def _rollout_curve(env: ProtocolArenaEnvironment, policy: PolicyFn,
task_id: str, seed: int,
max_turns: int = 12) -> Dict[str, Any]:
obs = env.reset(task_id=task_id, seed=seed)
cum = [0.0]
drift_turn = None
prior_drift_fired = bool(env.state.drift_fired)
for _ in range(max_turns):
if obs.done:
break
d = policy(obs.model_dump())
d.setdefault("rationale", "money-plot rollout action.")
if len(d["rationale"]) < 20:
d["rationale"] = (d["rationale"] + " " * 25)[:40]
d = {k: v for k, v in d.items()
if k in {"kind", "rationale", "mcp_call", "a2a_call",
"dag_delta", "kg_op", "final", "rewind_n", "confidence"}}
obs = env.step(OrchestratorAction(**d))
cum.append(cum[-1] + float(obs.reward))
try:
now_fired = bool(env.state.drift_fired)
if drift_turn is None and now_fired and not prior_drift_fired:
drift_turn = obs.turn
prior_drift_fired = now_fired
except Exception:
pass
return {"cum": cum, "drift_turn": drift_turn, "final": float(obs.reward)}
def _pad_to(curve: List[float], length: int) -> List[float]:
"""Right-pad a per-turn cumulative reward curve to a common length by
repeating the final value (cumulative reward stays flat once the
episode ends, by definition)."""
if len(curve) >= length:
return curve[:length]
return curve + [curve[-1]] * (length - len(curve))
def main(argv=None) -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--task", default="research_photo_rename")
ap.add_argument("--seeds", nargs="+", type=int, default=[0, 1, 2],
help="Seeds to average over (default 3 seeds).")
ap.add_argument("--policies", nargs="+",
default=["rule_based", "keyword", "random"])
ap.add_argument("--out", default="reports/drift_recovery.png")
ap.add_argument("--max-turns", type=int, default=12)
args = ap.parse_args(argv)
try:
import matplotlib.pyplot as plt
import numpy as np
except Exception as e:
print(f"[plot] matplotlib/numpy unavailable: {e}", file=sys.stderr)
print("[plot] install with: pip install matplotlib numpy", file=sys.stderr)
return 1
env = ProtocolArenaEnvironment()
# policy_label → list of per-seed curves; per-seed-drift-turns
per_policy: Dict[str, List[List[float]]] = {}
drift_turns: List[int] = []
for spec in args.policies:
label, fn = _resolve(spec)
per_policy.setdefault(label, [])
for seed in args.seeds:
c = _rollout_curve(env, fn, args.task, seed, args.max_turns)
per_policy[label].append(c["cum"])
if c["drift_turn"] is not None:
drift_turns.append(c["drift_turn"])
print(f"[plot] {label:>12} seed={seed} "
f"cum_terminal={c['cum'][-1]:.3f} drift@turn={c['drift_turn']}")
common_len = max(max(len(c) for c in cs) for cs in per_policy.values())
plt.figure(figsize=(8.0, 4.8))
colors = {"rule_based": "#1f77b4", "keyword": "#ff7f0e",
"random": "#7f7f7f", "trained": "#2ca02c"}
styles = {"rule_based": "-", "keyword": "--",
"random": ":", "trained": "-"}
for label, curves in per_policy.items():
padded = np.array([_pad_to(c, common_len) for c in curves])
mean = padded.mean(axis=0)
std = padded.std(axis=0)
x = np.arange(common_len)
color = colors.get(label, "#d62728")
plt.plot(x, mean, marker="o", color=color,
linestyle=styles.get(label, "-"),
label=label, linewidth=2.2)
if len(curves) > 1:
plt.fill_between(x, mean - std, mean + std, color=color, alpha=0.15)
if drift_turns:
# Modal drift turn across seeds — most common firing turn.
from collections import Counter
dt = Counter(drift_turns).most_common(1)[0][0]
plt.axvline(dt, linestyle="--", color="#d62728", linewidth=1.5,
label=f"drift fires at turn {dt}")
n_seeds = len(args.seeds)
plt.xlabel("turn")
plt.ylabel("cumulative reward")
plt.title(f"Drift-recovery — {args.task} "
f"(mean ±1σ over {n_seeds} seed{'s' if n_seeds != 1 else ''})")
plt.grid(True, alpha=0.3)
plt.legend(loc="best", fontsize=9)
plt.tight_layout()
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
plt.savefig(args.out, dpi=150)
plt.close()
print(f"[plot] wrote {args.out}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
|