MCPcopy Create free account
hub / github.com/TeleHuman/TextOp / main

Function main

TextOpTracker/scripts/rsl_rl/play.py:76–200  ·  view source on GitHub ↗

Play with RSL-RL agent.

(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg)

Source from the content-addressed store, hash-verified

74
75@hydra_task_config(args_cli.task, "rsl_rl_cfg_entry_point")
76def 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}")

Callers 1

play.pyFile · 0.70

Calls 9

loadMethod · 0.95
get_inference_policyMethod · 0.95
OnPolicyRunnerClass · 0.90
attach_onnx_metadataFunction · 0.90
runMethod · 0.80
get_observationsMethod · 0.80
stepMethod · 0.80
closeMethod · 0.45

Tested by

no test coverage detected