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