from __future__ import annotations import argparse import json import math import random from collections import deque from pathlib import Path import numpy as np import torch import trackio from castle_env import CastleEnv, Transition from model import DuelingQNetwork from safetensors.torch import save_file from torch import nn PROJECT_DIR = Path(__file__).resolve().parent ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "castle-nav-geodesic-dqfd" def epsilon(step: int, decay_steps: int = 45_000) -> float: progress = min(step / decay_steps, 1.0) return 0.03 + (0.35 - 0.03) * math.exp(-4.5 * progress) def collect_demonstrations( env: CastleEnv, episodes: int, ) -> list[Transition]: demonstrations = [] for _ in range(episodes): state = env.reset() for _ in range(env.max_steps): action = env.optimal_action() next_state, reward, done, _ = env.step(action) demonstrations.append( Transition(state, action, reward, next_state, done, is_demo=True) ) state = next_state if done: break return demonstrations def pretrain_from_demonstrations( model: DuelingQNetwork, optimizer: torch.optim.Optimizer, demonstrations: list[Transition], updates: int, ) -> tuple[float, float]: losses = [] for _ in range(updates): batch = random.sample(demonstrations, min(256, len(demonstrations))) states = torch.tensor(np.stack([item.state for item in batch])) actions = torch.tensor([item.action for item in batch], dtype=torch.long) loss = nn.functional.cross_entropy(model(states), actions) optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() losses.append(float(loss.detach())) validation = random.sample(demonstrations, min(2000, len(demonstrations))) states = torch.tensor(np.stack([item.state for item in validation])) actions = torch.tensor([item.action for item in validation], dtype=torch.long) with torch.no_grad(): accuracy = float((model(states).argmax(dim=1) == actions).float().mean()) return sum(losses) / len(losses), accuracy def evaluate_success( model: DuelingQNetwork, episodes: int = 150, seed: int = 10_000, ) -> float: environment = CastleEnv(seed=seed) successes = 0 with torch.no_grad(): for _ in range(episodes): state = environment.reset() info = {"success": False} for _ in range(environment.max_steps): action = int(model(torch.tensor(state)[None]).argmax(dim=1)) state, _, done, info = environment.step(action) if done: break successes += int(info["success"]) return successes / episodes def optimize( online: DuelingQNetwork, target: DuelingQNetwork, optimizer: torch.optim.Optimizer, replay: deque[Transition], batch_size: int, gamma: float, ) -> float: batch = random.sample(replay, batch_size) states = torch.tensor(np.stack([item.state for item in batch])) actions = torch.tensor([item.action for item in batch], dtype=torch.long) rewards = torch.tensor([item.reward for item in batch], dtype=torch.float32) next_states = torch.tensor(np.stack([item.next_state for item in batch])) dones = torch.tensor([item.done for item in batch], dtype=torch.float32) demo_mask = torch.tensor([item.is_demo for item in batch], dtype=torch.bool) predicted = online(states).gather(1, actions[:, None]).squeeze(1) with torch.no_grad(): next_actions = online(next_states).argmax(dim=1) next_values = target(next_states).gather(1, next_actions[:, None]).squeeze(1) expected = rewards + gamma * (1 - dones) * next_values temporal_difference_loss = nn.functional.smooth_l1_loss(predicted, expected) if demo_mask.any(): demonstration_loss = nn.functional.cross_entropy( online(states[demo_mask]), actions[demo_mask], ) else: demonstration_loss = torch.tensor(0.0) loss = temporal_difference_loss + 0.65 * demonstration_loss optimizer.zero_grad(set_to_none=True) loss.backward() nn.utils.clip_grad_norm_(online.parameters(), 5.0) optimizer.step() return float(loss.detach()) def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--episodes", type=int, default=700) parser.add_argument("--demo-episodes", type=int, default=1200) parser.add_argument("--pretrain-updates", type=int, default=800) parser.add_argument("--seed", type=int, default=7) args = parser.parse_args() random.seed(args.seed) np.random.seed(args.seed) torch.manual_seed(args.seed) env = CastleEnv(seed=args.seed) online = DuelingQNetwork() target = DuelingQNetwork() target.load_state_dict(online.state_dict()) optimizer = torch.optim.AdamW(online.parameters(), lr=8e-4, weight_decay=1e-4) replay: deque[Transition] = deque(maxlen=50_000) returns: list[float] = [] successes: list[float] = [] losses: list[float] = [] global_step = 0 demonstrations = collect_demonstrations(env, args.demo_episodes) imitation_loss, imitation_accuracy = pretrain_from_demonstrations( online, optimizer, demonstrations, args.pretrain_updates, ) target.load_state_dict(online.state_dict()) replay.extend(demonstrations) pretrain_success = evaluate_success(online) best_success = pretrain_success best_episode = 0 best_state = { name: tensor.detach().clone() for name, tensor in online.state_dict().items() } trackio.init( project="castle-nav-rl", name="dueling-double-dqn-v4-geodesic-features", config={ "episodes": args.episodes, "demonstration_episodes": args.demo_episodes, "pretrain_updates": args.pretrain_updates, "gamma": 0.985, "replay_size": replay.maxlen, "seed": args.seed, }, ) trackio.log( { "imitation_loss": imitation_loss, "imitation_accuracy": imitation_accuracy, "pretrain_success": pretrain_success, "demonstration_transitions": len(demonstrations), } ) for episode in range(args.episodes): state = env.reset() episode_return = 0.0 info = {"success": False} for _ in range(env.max_steps): explore = random.random() < epsilon(global_step) if explore: action = random.randrange(4) else: with torch.no_grad(): action = int(online(torch.tensor(state)[None]).argmax(dim=1)) next_state, reward, done, info = env.step(action) replay.append(Transition(state, action, reward, next_state, done)) state = next_state episode_return += reward global_step += 1 if len(replay) >= 512 and global_step % 8 == 0: losses.append(optimize(online, target, optimizer, replay, 64, gamma=0.985)) if global_step % 400 == 0: target.load_state_dict(online.state_dict()) if done: break returns.append(episode_return) successes.append(float(info["success"])) if (episode + 1) % 25 == 0: window = min(100, len(returns)) trackio.log( { "episode": episode + 1, "return_100": sum(returns[-window:]) / window, "success_rate_100": sum(successes[-window:]) / window, "loss_100": sum(losses[-100:]) / max(1, len(losses[-100:])), "epsilon": epsilon(global_step), "global_step": global_step, } ) if (episode + 1) % 100 == 0: checkpoint_success = evaluate_success( online, seed=10_000 + episode, ) trackio.log( { "episode": episode + 1, "deterministic_success": checkpoint_success, } ) if checkpoint_success > best_success: best_success = checkpoint_success best_episode = episode + 1 best_state = { name: tensor.detach().clone() for name, tensor in online.state_dict().items() } if episode == 499 and sum(successes[-100:]) < 70: trackio.alert( title="Low navigation success", text="Success remained below 70% after 500 episodes.", level=trackio.AlertLevel.WARN, ) trackio.finish() ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) final_training_success = sum(successes[-100:]) / 100 online.load_state_dict(best_state) save_file(online.state_dict(), ARTIFACT_DIR / "policy.safetensors") config = { "architecture": "DuelingQNetwork", "algorithm": "Dueling Double DQN with demonstration pretraining", "observation_size": 17, "actions": 4, "hidden_size": 128, "parameters": sum(parameter.numel() for parameter in online.parameters()), "episodes": args.episodes, "demonstration_episodes": args.demo_episodes, "demonstration_transitions": len(demonstrations), "imitation_accuracy": imitation_accuracy, "pretrain_success": pretrain_success, "best_deterministic_success": best_success, "best_episode": best_episode, "global_steps": global_step, "success_rate_last_100": final_training_success, "return_last_100": sum(returns[-100:]) / 100, "seed": args.seed, } (ARTIFACT_DIR / "config.json").write_text( json.dumps(config, indent=2), encoding="utf-8", ) (ARTIFACT_DIR / "training_curve.json").write_text( json.dumps( { "returns": returns, "successes": successes, "loss_tail": losses[-1000:], } ), encoding="utf-8", ) print(json.dumps(config, indent=2)) if __name__ == "__main__": main()