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

Function main

TextOpTracker/scripts/rsl_rl/train.py:75–146  ·  view source on GitHub ↗

Train with RSL-RL agent.

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

Source from the content-addressed store, hash-verified

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

Callers 1

train.pyFile · 0.70

Calls 4

add_git_repo_to_logMethod · 0.95
loadMethod · 0.95
OnPolicyRunnerClass · 0.85
closeMethod · 0.45

Tested by

no test coverage detected