Spaces:
Sleeping
Sleeping
File size: 5,340 Bytes
0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 0fca41e f88d446 | 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 | import argparse
import glob
import json
import os
from pathlib import Path
from typing import Dict, List
def _load_metric_files(log_dir: str, latest_run_only: bool = True) -> List[Dict]:
metric_files = sorted(glob.glob(os.path.join(log_dir, "*_metrics.json")))
training_logs = sorted(glob.glob(os.path.join(log_dir, "training_log_*.json")))
payloads = []
if latest_run_only:
if metric_files:
metric_files = [metric_files[-1]]
if training_logs:
training_logs = [training_logs[-1]]
for path in metric_files:
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
data["_path"] = path
payloads.append(data)
# Backward-compatible support for training logs produced by train_rl.py.
for path in training_logs:
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
episode_data = data.get("episode_data", [])
if not episode_data:
continue
rewards = [float(ep.get("reward", 0.0)) for ep in episode_data]
r_min = min(rewards)
r_max = max(rewards)
denom = (r_max - r_min) if r_max > r_min else 1.0
episodes = []
for i, ep in enumerate(episode_data):
reward = float(ep.get("reward", 0.0))
# Training logs do not store per-episode task score, so derive a 0-1 proxy from reward.
score_proxy = (reward - r_min) / denom if r_max > r_min else 1.0
episodes.append(
{
"episode": int(ep.get("episode", i + 1)),
"score": float(score_proxy),
"total_reward": reward,
"steps": float(ep.get("steps", 0.0)),
"final_error_rate": float(ep.get("final_error_rate", 0.0)),
}
)
payloads.append(
{
"agent_name": Path(path).stem,
"episodes": episodes,
"statistics": {
"avg_reward": float(sum(rewards) / len(rewards)),
},
"_path": path,
}
)
return payloads
def visualize_training(log_dir: str, save_path: str = "training_comparison.png", latest_run_only: bool = True):
try:
import matplotlib.pyplot as plt
except ImportError:
print("matplotlib not installed. Install with: pip install matplotlib")
return
payloads = _load_metric_files(log_dir, latest_run_only=latest_run_only)
if not payloads:
print(f"No metrics files found in {log_dir}")
return
fig, axes = plt.subplots(2, 2, figsize=(12, 8))
for payload in payloads:
agent = payload.get("agent_name", "unknown")
episodes = payload.get("episodes", [])
stats = payload.get("statistics", {})
if episodes:
x = [ep.get("episode", i + 1) for i, ep in enumerate(episodes)]
scores = [float(ep.get("score", 0.0)) for ep in episodes]
rewards = [float(ep.get("total_reward", 0.0)) for ep in episodes]
steps = [float(ep.get("steps", 0.0)) for ep in episodes]
errors = [float(ep.get("final_error_rate", 0.0)) * 100.0 for ep in episodes]
axes[0, 0].plot(x, scores, marker="o", linewidth=1.8, label=agent)
axes[0, 1].plot(x, rewards, marker="o", linewidth=1.8, label=agent)
axes[1, 0].plot(x, steps, marker="o", linewidth=1.8, label=agent)
axes[1, 1].plot(x, errors, marker="o", linewidth=1.8, label=agent)
else:
print(f"Warning: {agent} has no per-episode data in {payload['_path']}")
# Fall back to a single-point visualization from statistics.
axes[0, 0].scatter([1], [float(stats.get("avg_score", 0.0))], label=agent)
axes[0, 1].scatter([1], [float(stats.get("avg_reward", 0.0))], label=agent)
axes[1, 0].scatter([1], [float(stats.get("avg_steps", 0.0))], label=agent)
axes[1, 1].scatter([1], [float(stats.get("avg_final_error", 0.0)) * 100.0], label=agent)
axes[0, 0].set_title("Score per Episode")
axes[0, 1].set_title("Reward per Episode")
axes[1, 0].set_title("Steps per Episode")
axes[1, 1].set_title("Final Error Rate per Episode")
axes[0, 0].set_ylabel("Score")
axes[0, 1].set_ylabel("Reward")
axes[1, 0].set_ylabel("Steps")
axes[1, 1].set_ylabel("Error %")
for ax in axes.flatten():
ax.set_xlabel("Episode")
ax.grid(True, alpha=0.3)
ax.legend()
plt.tight_layout()
plt.savefig(save_path, dpi=300)
print(f"Saved training visualization to {save_path}")
plt.show()
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Visualize training logs from metrics JSON files")
parser.add_argument("--log-dir", type=str, default="logs/training")
parser.add_argument("--save-path", type=str, default="training_comparison.png")
parser.add_argument(
"--all-runs",
action="store_true",
help="Overlay all runs found in log-dir instead of only the latest run",
)
args = parser.parse_args()
visualize_training(
log_dir=args.log_dir,
save_path=args.save_path,
latest_run_only=not args.all_runs,
)
|