| 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() |
|
|