from __future__ import annotations import argparse import json import os import random import sys from dataclasses import dataclass from pathlib import Path from typing import Dict, List REPO_ROOT = Path(__file__).resolve().parents[1] if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) from baseline.trained_q_agent import NUM_ACTIONS, QTable, TrainedQAgent, action_from_index, observation_key from env.environment import LastMileDeliveryEnvironment from tasks.registry import get_task_config @dataclass class EpisodeStats: total_reward: float completion_rate: float success: bool steps: int def _new_q_row() -> List[float]: return [0.0 for _ in range(NUM_ACTIONS)] def _argmax(values: List[float]) -> int: best_index = 0 best_value = values[0] for index in range(1, len(values)): if values[index] > best_value: best_value = values[index] best_index = index return best_index def _episode_summary(env: LastMileDeliveryEnvironment) -> EpisodeStats: state = env.state() observation = state.observation if observation is None: return EpisodeStats(total_reward=state.total_reward, completion_rate=0.0, success=False, steps=state.step_count) delivered = state.delivered_orders remaining = len(observation.pending_orders) if observation.current_order is not None: remaining += 1 total_orders = max(1, delivered + remaining) completion_rate = delivered / total_orders success = bool(state.done and remaining == 0) return EpisodeStats( total_reward=state.total_reward, completion_rate=completion_rate, success=success, steps=state.step_count, ) def train_q_agent( task: str, episodes: int, seed: int, alpha: float, gamma: float, epsilon_start: float, epsilon_end: float, epsilon_decay: float, ) -> QTable: rng = random.Random(seed) q_table: QTable = {} running_reward = 0.0 running_completion = 0.0 for episode in range(episodes): current_seed = seed + episode config = get_task_config(task_name=task, seed=current_seed) env = LastMileDeliveryEnvironment(config=config) observation = env.reset(seed=current_seed) epsilon = max(epsilon_end, epsilon_start * (epsilon_decay ** episode)) for _ in range(config.max_steps): state_key = observation_key(observation) row = q_table.setdefault(state_key, _new_q_row()) if rng.random() < epsilon: action_index = rng.randrange(NUM_ACTIONS) else: action_index = _argmax(row) action = action_from_index(action_index, observation) next_observation, reward, done, _ = env.step(action) next_key = observation_key(next_observation) next_row = q_table.setdefault(next_key, _new_q_row()) td_target = reward if done else reward + gamma * max(next_row) row[action_index] += alpha * (td_target - row[action_index]) observation = next_observation if done: break summary = _episode_summary(env) running_reward += summary.total_reward running_completion += summary.completion_rate if (episode + 1) % max(1, episodes // 10) == 0: avg_reward = running_reward / (episode + 1) avg_completion = running_completion / (episode + 1) print( f"[TRAIN] episode={episode + 1} " f"epsilon={epsilon:.4f} " f"avg_reward={avg_reward:.2f} " f"avg_completion={avg_completion:.3f}", flush=True, ) return q_table def evaluate_q_agent(task: str, q_table: QTable, episodes: int, seed: int) -> Dict[str, float]: agent = TrainedQAgent(q_table=q_table) total_reward = 0.0 total_completion = 0.0 success_count = 0 total_steps = 0 for episode in range(episodes): current_seed = seed + 10_000 + episode config = get_task_config(task_name=task, seed=current_seed) env = LastMileDeliveryEnvironment(config=config) observation = env.reset(seed=current_seed) for _ in range(config.max_steps): action = agent.act(observation) observation, _, done, _ = env.step(action) if done: break summary = _episode_summary(env) total_reward += summary.total_reward total_completion += summary.completion_rate success_count += 1 if summary.success else 0 total_steps += summary.steps count = max(1, episodes) return { "avg_reward": total_reward / count, "avg_completion": total_completion / count, "success_rate": success_count / count, "avg_steps": total_steps / count, } def save_model(model_path: str, task: str, seed: int, episodes: int, q_table: QTable) -> None: output_dir = os.path.dirname(model_path) if output_dir: os.makedirs(output_dir, exist_ok=True) payload = { "model_type": "tabular_q_learning", "task": task, "seed": seed, "episodes": episodes, "num_actions": NUM_ACTIONS, "q_values": q_table, } with open(model_path, "w", encoding="utf-8") as handle: json.dump(payload, handle, separators=(",", ":"), sort_keys=True) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Train a tabular Q-learning agent for delivery-openenv") parser.add_argument("--task", type=str, default="easy", choices=["easy", "medium", "hard"]) parser.add_argument("--episodes", type=int, default=800) parser.add_argument("--eval-episodes", type=int, default=100) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--alpha", type=float, default=0.20) parser.add_argument("--gamma", type=float, default=0.98) parser.add_argument("--epsilon-start", type=float, default=1.0) parser.add_argument("--epsilon-end", type=float, default=0.05) parser.add_argument("--epsilon-decay", type=float, default=0.995) parser.add_argument("--output", type=str, default="models/q_agent_easy.json") return parser.parse_args() def main() -> None: args = parse_args() q_table = train_q_agent( task=args.task, episodes=args.episodes, seed=args.seed, alpha=args.alpha, gamma=args.gamma, epsilon_start=args.epsilon_start, epsilon_end=args.epsilon_end, epsilon_decay=args.epsilon_decay, ) metrics = evaluate_q_agent( task=args.task, q_table=q_table, episodes=args.eval_episodes, seed=args.seed, ) save_model(model_path=args.output, task=args.task, seed=args.seed, episodes=args.episodes, q_table=q_table) print( f"[EVAL] task={args.task} episodes={args.eval_episodes} " f"avg_reward={metrics['avg_reward']:.2f} " f"avg_completion={metrics['avg_completion']:.3f} " f"success_rate={metrics['success_rate']:.3f} " f"avg_steps={metrics['avg_steps']:.2f}", flush=True, ) print(f"[MODEL] saved={args.output} states={len(q_table)}", flush=True) if __name__ == "__main__": main()