| import argparse |
| import os |
|
|
| import gymnasium as gym |
| import imageio.v2 as imageio |
|
|
| from enjoy import load_agent_from_checkpoint |
|
|
|
|
| def record(weights_path, output_path, seed=42, max_steps=1000): |
| agent = load_agent_from_checkpoint(weights_path, seed=seed) |
| env = gym.make("Acrobot-v1", render_mode="rgb_array") |
| frames = [] |
| obs, _ = env.reset(seed=seed) |
| for _ in range(max_steps): |
| Hr, Hi = agent.encoder.encode(obs) |
| action = agent.actor.argmax(Hr, Hi) |
| obs, reward, term, trunc, _ = env.step(action) |
| frames.append(env.render()) |
| if term or trunc: |
| break |
| env.close() |
| os.makedirs(os.path.dirname(output_path) or ".", exist_ok=True) |
| imageio.mimsave(output_path, frames, fps=30) |
| return len(frames) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Record Acrobot-v1 replay video") |
| parser.add_argument("--weights", default="hdppo-Acrobot-v1/weights.npz") |
| parser.add_argument("--output", default="replay.mp4") |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--max-steps", type=int, default=1000) |
| args = parser.parse_args() |
| n = record(args.weights, args.output, seed=args.seed, max_steps=args.max_steps) |
| print(f"saved {n} frames to {args.output}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|