openenv-claw1 / scripts /train_q_agent.py
vishal harkal
Deploy full OpenEnv API instead of starter app
f104717
Raw History Blame
7.36 kB
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()