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,
    )