Spaces:
Sleeping
Sleeping
File size: 1,519 Bytes
1b72fa2 | 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 | from traffic_rl.env.traffic_env import TrafficEnv
from traffic_rl.training.trainer import TrainingConfig, train_dqn
from traffic_rl.evaluation.evaluator import evaluate_agent, evaluate_fixed_controller
def test_training_loop_returns_history():
env = TrafficEnv(
config={
"max_steps": 20,
"arrival_mode": "deterministic",
"arrival_sequence": [[1, 1, 1, 1]],
"ambulance_spawn_prob": 0.0,
"seed": 11,
}
)
cfg = TrainingConfig(episodes=4, max_steps=20, batch_size=8, target_sync_interval=2)
agent, history = train_dqn(env=env, config=cfg)
assert len(history["episode_reward"]) == 4
assert len(history["avg_queue"]) == 4
assert agent.action_dim == 3
def test_evaluation_returns_metrics():
env_config = {
"max_steps": 20,
"arrival_mode": "deterministic",
"arrival_sequence": [[2, 1, 2, 1]],
"ambulance_spawn_prob": 0.0,
"seed": 12,
}
env = TrafficEnv(config=env_config)
cfg = TrainingConfig(episodes=3, max_steps=20, batch_size=8, target_sync_interval=2)
agent, _ = train_dqn(env=env, config=cfg)
fixed_metrics = evaluate_fixed_controller(env_config=env_config, episodes=2, switch_interval=2)
rl_metrics = evaluate_agent(agent=agent, env_config=env_config, episodes=2)
assert fixed_metrics["avg_waiting_time"] >= 0
assert rl_metrics["avg_queue_length"] >= 0
assert "throughput" in fixed_metrics
assert "throughput" in rl_metrics
|