File size: 4,146 Bytes
2fc96b6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from datetime import datetime
from pathlib import Path

import gradio as gr

from traffic_rl.env.traffic_env import TrafficEnv
from traffic_rl.evaluation.evaluator import (
    compare_policies,
    evaluate_agent,
    evaluate_fixed_controller,
)
from traffic_rl.training.trainer import TrainingConfig, train_dqn
from traffic_rl.visualization.dashboard import plot_comparison, plot_training_history


def _default_env_config(ambulance_prob: float, seed: int) -> dict:
    return {
        "max_steps": 120,
        "arrival_mode": "stochastic",
        "lane_bias": (1.6, 0.8, 1.4, 0.6),
        "peak_rates": (3.8, 2.2, 3.4, 1.5),
        "offpeak_rates": (1.4, 0.9, 1.2, 0.7),
        "peak_duration": 35,
        "cycle_duration": 60,
        "service_rate": 2,
        "ambulance_spawn_prob": ambulance_prob,
        "seed": seed,
    }


def run_experiment(train_episodes: int, eval_episodes: int, ambulance_prob: float, seed: int):
    env_config = _default_env_config(ambulance_prob=ambulance_prob, seed=seed)

    baseline_metrics = evaluate_fixed_controller(
        env_config=env_config,
        episodes=eval_episodes,
        switch_interval=5,
    )

    train_env = TrafficEnv(config=env_config)
    cfg = TrainingConfig(
        episodes=train_episodes,
        max_steps=env_config["max_steps"],
        batch_size=64,
        target_sync_interval=10,
        epsilon_decay=0.97,
        seed=seed,
    )

    agent, history = train_dqn(env=train_env, config=cfg)
    rl_metrics = evaluate_agent(agent=agent, env_config=env_config, episodes=eval_episodes)
    improvement = compare_policies(baseline_metrics, rl_metrics)

    run_dir = Path("outputs") / "space_runs" / datetime.now().strftime("%Y%m%d_%H%M%S")
    run_dir.mkdir(parents=True, exist_ok=True)

    train_plot = plot_training_history(history, output_dir=run_dir)["training_history"]
    compare_plot = plot_comparison(baseline_metrics, rl_metrics, output_dir=run_dir)["policy_comparison"]

    result = {
        "baseline": baseline_metrics,
        "rl": rl_metrics,
        "improvement": improvement,
        "config": {
            "train_episodes": train_episodes,
            "eval_episodes": eval_episodes,
            "ambulance_prob": ambulance_prob,
            "seed": seed,
        },
    }

    summary_md = (
        "## Run Summary\n"
        f"- Waiting time improvement: **{improvement['waiting_time_improvement_pct']:.2f}%**\n"
        f"- Queue length improvement: **{improvement['queue_length_improvement_pct']:.2f}%**\n"
        f"- Throughput gain: **{improvement['throughput_gain_pct']:.2f}%**\n"
        f"- Ambulance clearance gain: **{improvement['ambulance_clearance_gain_pct']:.2f}%**"
    )

    return result, summary_md, train_plot, compare_plot


def build_app() -> gr.Blocks:
    with gr.Blocks(title="RL-Based Adaptive Traffic Intelligence") as demo:
        gr.Markdown(
            """
# ?? RL-Based Adaptive Traffic Intelligence System
Train and compare a DQN traffic controller against a fixed-time baseline.
This demo optimizes waiting time, queue length, throughput, and emergency handling.
            """
        )

        with gr.Row():
            train_episodes = gr.Slider(20, 200, value=70, step=10, label="Training Episodes")
            eval_episodes = gr.Slider(5, 50, value=20, step=5, label="Evaluation Episodes")
        with gr.Row():
            ambulance_prob = gr.Slider(0.0, 0.4, value=0.08, step=0.01, label="Ambulance Spawn Probability")
            seed = gr.Number(value=42, precision=0, label="Random Seed")

        run_btn = gr.Button("Run RL vs Baseline", variant="primary")

        metrics_json = gr.JSON(label="Metrics")
        summary = gr.Markdown()
        train_img = gr.Image(label="Training Trends")
        compare_img = gr.Image(label="Policy Comparison")

        run_btn.click(
            fn=run_experiment,
            inputs=[train_episodes, eval_episodes, ambulance_prob, seed],
            outputs=[metrics_json, summary, train_img, compare_img],
        )

    return demo


if __name__ == "__main__":
    app = build_app()
    app.launch()