ARotting's picture
Publish 1.7K parameter distilled navigation policy
a26d6ed verified
Raw
History Blame Contribute Delete
10.5 kB
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()