Play with RSL-RL agent.
(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg)
| 74 | |
| 75 | @hydra_task_config(args_cli.task, "rsl_rl_cfg_entry_point") |
| 76 | def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg): |
| 77 | """Play with RSL-RL agent.""" |
| 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 | |
| 81 | # specify directory for logging experiments |
| 82 | log_root_path = os.path.join("logs", "rsl_rl", agent_cfg.experiment_name) |
| 83 | log_root_path = os.path.abspath(log_root_path) |
| 84 | |
| 85 | if args_cli.wandb_path: |
| 86 | raise NotImplementedError("Wandb is not supported") |
| 87 | import wandb |
| 88 | |
| 89 | run_path = args_cli.wandb_path |
| 90 | |
| 91 | api = wandb.Api() |
| 92 | if "model" in args_cli.wandb_path: |
| 93 | run_path = "/".join(args_cli.wandb_path.split("/")[:-1]) |
| 94 | wandb_run = api.run(run_path) |
| 95 | # loop over files in the run |
| 96 | files = [file.name for file in wandb_run.files() if "model" in file.name] |
| 97 | # files are all model_xxx.pt find the largest filename |
| 98 | if "model" in args_cli.wandb_path: |
| 99 | file = args_cli.wandb_path.split("/")[-1] |
| 100 | else: |
| 101 | file = max(files, key=lambda x: int(x.split("_")[1].split(".")[0])) |
| 102 | |
| 103 | wandb_file = wandb_run.file(str(file)) |
| 104 | wandb_file.download("./logs/rsl_rl/temp", replace=True) |
| 105 | |
| 106 | print(f"[INFO]: Loading model checkpoint from: {run_path}/{file}") |
| 107 | resume_path = f"./logs/rsl_rl/temp/{file}" |
| 108 | |
| 109 | if args_cli.motion_file is not None: |
| 110 | print(f"[INFO]: Using motion file from CLI: {args_cli.motion_file}") |
| 111 | env_cfg.commands.motion.motion_file = args_cli.motion_file |
| 112 | |
| 113 | art = next((a for a in wandb_run.used_artifacts() if a.type == "motions"), None) |
| 114 | if art is None: |
| 115 | print("[WARN] No model artifact found in the run.") |
| 116 | else: |
| 117 | env_cfg.commands.motion.motion_file = str(pathlib.Path(art.download()) / "motion.npz") |
| 118 | |
| 119 | elif args_cli.resume_path: |
| 120 | assert args_cli.motion_file is not None, "Motion file is required when resume_path is provided" |
| 121 | resume_path = args_cli.resume_path |
| 122 | # env_cfg.commands.motion.motion_file = str(Path("./artifacts") / Path(args_cli.motion_file) / "motion.npz") |
| 123 | # env_cfg, agent_cfg = load_config(resume_path) |
| 124 | |
| 125 | motion_files = glob.glob(str(Path("./artifacts") / Path(args_cli.motion_file) / "motion.npz")) |
| 126 | if not motion_files: |
| 127 | raise FileNotFoundError(f"No motion.npz found in {Path('./artifacts') / Path(args_cli.motion_file)}") |
| 128 | env_cfg.commands.motion.motion_files = motion_files # List[str] |
| 129 | |
| 130 | print(f"[INFO]: Using motion file from CLI: {args_cli.motion_file}") |
| 131 | print(f"[INFO]: Using resume path from CLI: {args_cli.resume_path}") |
| 132 | else: |
| 133 | print(f"[INFO] Loading experiment from directory: {log_root_path}") |
no test coverage detected