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