File size: 1,907 Bytes
a26d6ed | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | 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()
|