| from __future__ import annotations |
|
|
| import json |
| from pathlib import Path |
|
|
| import torch |
| from castle_env import CastleEnv |
| from model import DuelingQNetwork |
| from safetensors.torch import load_file |
|
|
| PROJECT_DIR = Path(__file__).resolve().parent |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "castle-nav-geodesic-dqfd" |
|
|
|
|
| def main() -> None: |
| model = DuelingQNetwork() |
| model.load_state_dict(load_file(ARTIFACT_DIR / "policy.safetensors")) |
| model.eval() |
| env = CastleEnv(seed=2026) |
| successes = 0 |
| returns = [] |
| lengths = [] |
| rollouts = [] |
| with torch.no_grad(): |
| for episode in range(500): |
| state = env.reset() |
| total_reward = 0.0 |
| path = [env.agent.tolist()] |
| info = {"success": False} |
| for _ in range(env.max_steps): |
| action = int(model(torch.tensor(state)[None]).argmax(dim=1)) |
| state, reward, done, info = env.step(action) |
| total_reward += reward |
| path.append(env.agent.tolist()) |
| if done: |
| break |
| successes += int(info["success"]) |
| returns.append(total_reward) |
| lengths.append(env.steps) |
| if episode < 5: |
| rollouts.append( |
| { |
| "success": info["success"], |
| "steps": env.steps, |
| "path": path, |
| } |
| ) |
| results = { |
| "evaluation_episodes": 500, |
| "success_rate": successes / 500, |
| "average_return": sum(returns) / len(returns), |
| "average_steps": sum(lengths) / len(lengths), |
| "example_rollouts": rollouts, |
| } |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(results, indent=2), |
| encoding="utf-8", |
| ) |
| print(json.dumps(results, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|