Train with RSL-RL agent.
(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg)
| 73 | |
| 74 | @hydra_task_config(args_cli.task, "rsl_rl_cfg_entry_point") |
| 75 | def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg): |
| 76 | """Train with RSL-RL agent.""" |
| 77 | # override configurations with non-hydra CLI arguments |
| 78 | agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli) |
| 79 | env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs |
| 80 | agent_cfg.max_iterations = args_cli.max_iterations if args_cli.max_iterations is not None else agent_cfg.max_iterations |
| 81 | |
| 82 | # set the environment seed |
| 83 | # note: certain randomizations occur in the environment initialization so we set the seed here |
| 84 | env_cfg.seed = agent_cfg.seed |
| 85 | env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device |
| 86 | |
| 87 | motion_files = glob.glob(str(Path("./artifacts") / Path(args_cli.motion_file) / "motion.npz")) |
| 88 | if not motion_files: |
| 89 | raise FileNotFoundError(f"No motion.npz found in {Path('./artifacts') / Path(args_cli.motion_file)}") |
| 90 | env_cfg.commands.motion.motion_files = motion_files # List[str] |
| 91 | |
| 92 | # specify directory for logging experiments |
| 93 | log_root_path = os.path.join("logs", "rsl_rl", agent_cfg.experiment_name) |
| 94 | log_root_path = os.path.abspath(log_root_path) |
| 95 | print(f"[INFO] Logging experiment in directory: {log_root_path}") |
| 96 | # specify directory for logging runs: {time-stamp}_{run_name} |
| 97 | log_dir = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") |
| 98 | if agent_cfg.run_name: |
| 99 | log_dir += f"_{agent_cfg.run_name}" |
| 100 | log_dir = os.path.join(log_root_path, log_dir) |
| 101 | |
| 102 | # create isaac environment |
| 103 | env = gym.make(args_cli.task, cfg=env_cfg, render_mode="rgb_array" if args_cli.video else None) |
| 104 | # wrap for video recording |
| 105 | if args_cli.video: |
| 106 | video_kwargs = { |
| 107 | "video_folder": os.path.join(log_dir, "videos", "train"), |
| 108 | "step_trigger": lambda step: step % args_cli.video_interval == 0, |
| 109 | "video_length": args_cli.video_length, |
| 110 | "disable_logger": True, |
| 111 | } |
| 112 | print("[INFO] Recording videos during training.") |
| 113 | print_dict(video_kwargs, nesting=4) |
| 114 | env = gym.wrappers.RecordVideo(env, **video_kwargs) |
| 115 | |
| 116 | # convert to single-agent instance if required by the RL algorithm |
| 117 | if isinstance(env.unwrapped, DirectMARLEnv): |
| 118 | env = multi_agent_to_single_agent(env) |
| 119 | |
| 120 | # wrap around environment for rsl-rl |
| 121 | env = RslRlVecEnvWrapper(env) |
| 122 | |
| 123 | # create runner from rsl-rl |
| 124 | runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device, registry_name=None) |
| 125 | # write git state to logs |
| 126 | runner.add_git_repo_to_log(__file__) |
| 127 | # save resume path before creating a new log_dir |
| 128 | if agent_cfg.resume: |
| 129 | # get path to previous checkpoint |
| 130 | resume_path = get_checkpoint_path(log_root_path, agent_cfg.load_run, agent_cfg.load_checkpoint) |
| 131 | print(f"[INFO]: Loading model checkpoint from: {resume_path}") |
| 132 | # load previously trained model |
no test coverage detected