Spaces:
Sleeping
Sleeping
| 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, | |
| ) | |