Initial commit
This commit is contained in:
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,93 @@
|
||||
# Copyright (c) 2022-2026, The Isaac Lab Project Developers (https://github.com/isaac-sim/IsaacLab/blob/main/CONTRIBUTORS.md).
|
||||
# All rights reserved.
|
||||
#
|
||||
# SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import random
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg
|
||||
|
||||
|
||||
def add_rsl_rl_args(parser: argparse.ArgumentParser):
|
||||
"""Add RSL-RL arguments to the parser.
|
||||
|
||||
Args:
|
||||
parser: The parser to add the arguments to.
|
||||
"""
|
||||
# create a new argument group
|
||||
arg_group = parser.add_argument_group("rsl_rl", description="Arguments for RSL-RL agent.")
|
||||
# -- experiment arguments
|
||||
arg_group.add_argument(
|
||||
"--experiment_name", type=str, default=None, help="Name of the experiment folder where logs will be stored."
|
||||
)
|
||||
arg_group.add_argument("--run_name", type=str, default=None, help="Run name suffix to the log directory.")
|
||||
# -- load arguments
|
||||
arg_group.add_argument("--resume", action="store_true", default=False, help="Whether to resume from a checkpoint.")
|
||||
arg_group.add_argument("--load_run", type=str, default=None, help="Name of the run folder to resume from.")
|
||||
arg_group.add_argument("--checkpoint", type=str, default=None, help="Checkpoint file to resume from.")
|
||||
# -- logger arguments
|
||||
arg_group.add_argument(
|
||||
"--logger", type=str, default=None, choices={"wandb", "tensorboard", "neptune"}, help="Logger module to use."
|
||||
)
|
||||
arg_group.add_argument(
|
||||
"--log_project_name", type=str, default=None, help="Name of the logging project when using wandb or neptune."
|
||||
)
|
||||
|
||||
|
||||
def parse_rsl_rl_cfg(task_name: str, args_cli: argparse.Namespace) -> RslRlBaseRunnerCfg:
|
||||
"""Parse configuration for RSL-RL agent based on inputs.
|
||||
|
||||
Args:
|
||||
task_name: The name of the environment.
|
||||
args_cli: The command line arguments.
|
||||
|
||||
Returns:
|
||||
The parsed configuration for RSL-RL agent based on inputs.
|
||||
"""
|
||||
from isaaclab_tasks.utils.parse_cfg import load_cfg_from_registry
|
||||
|
||||
# load the default configuration
|
||||
rslrl_cfg: RslRlBaseRunnerCfg = load_cfg_from_registry(task_name, "rsl_rl_cfg_entry_point")
|
||||
rslrl_cfg = update_rsl_rl_cfg(rslrl_cfg, args_cli)
|
||||
return rslrl_cfg
|
||||
|
||||
|
||||
def update_rsl_rl_cfg(agent_cfg: RslRlBaseRunnerCfg, args_cli: argparse.Namespace):
|
||||
"""Update configuration for RSL-RL agent based on inputs.
|
||||
|
||||
Args:
|
||||
agent_cfg: The configuration for RSL-RL agent.
|
||||
args_cli: The command line arguments.
|
||||
|
||||
Returns:
|
||||
The updated configuration for RSL-RL agent based on inputs.
|
||||
"""
|
||||
# override the default configuration with CLI arguments
|
||||
if hasattr(args_cli, "seed") and args_cli.seed is not None:
|
||||
# randomly sample a seed if seed = -1
|
||||
if args_cli.seed == -1:
|
||||
args_cli.seed = random.randint(0, 10000)
|
||||
agent_cfg.seed = args_cli.seed
|
||||
if args_cli.resume is not None:
|
||||
agent_cfg.resume = args_cli.resume
|
||||
if args_cli.load_run is not None:
|
||||
agent_cfg.load_run = args_cli.load_run
|
||||
if args_cli.checkpoint is not None:
|
||||
agent_cfg.load_checkpoint = args_cli.checkpoint
|
||||
if args_cli.experiment_name is not None:
|
||||
agent_cfg.experiment_name = args_cli.experiment_name
|
||||
if args_cli.run_name is not None:
|
||||
agent_cfg.run_name = args_cli.run_name
|
||||
if args_cli.logger is not None:
|
||||
agent_cfg.logger = args_cli.logger
|
||||
# set the project name for wandb and neptune
|
||||
if agent_cfg.logger in {"wandb", "neptune"} and args_cli.log_project_name:
|
||||
agent_cfg.wandb_project = args_cli.log_project_name
|
||||
agent_cfg.neptune_project = args_cli.log_project_name
|
||||
|
||||
return agent_cfg
|
||||
@@ -0,0 +1,251 @@
|
||||
# Copyright (c) 2022-2026, The Isaac Lab Project Developers (https://github.com/isaac-sim/IsaacLab/blob/main/CONTRIBUTORS.md).
|
||||
# All rights reserved.
|
||||
#
|
||||
# SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
"""Script to play a checkpoint if an RL agent from RSL-RL."""
|
||||
|
||||
import warnings
|
||||
|
||||
warnings.warn(
|
||||
"scripts/reinforcement_learning/rsl_rl/play.py is deprecated. Use "
|
||||
"`./isaaclab.sh play --rl_library rsl_rl --task <TASK>` instead. "
|
||||
"Example: `./isaaclab.sh play --rl_library rsl_rl --task Isaac-Cartpole-v0`.",
|
||||
DeprecationWarning,
|
||||
stacklevel=1,
|
||||
)
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import importlib.metadata as metadata
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
import gymnasium as gym
|
||||
import torch
|
||||
from packaging import version
|
||||
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
|
||||
|
||||
from isaaclab.envs import DirectMARLEnvCfg, DirectRLEnvCfg, ManagerBasedRLEnvCfg
|
||||
from isaaclab.utils.assets import retrieve_file_path
|
||||
from isaaclab.utils.dict import print_dict
|
||||
from isaaclab.utils.seed import configure_seed
|
||||
from isaaclab.utils.string import list_intersection, string_to_callable
|
||||
|
||||
from isaaclab_rl.rsl_rl import (
|
||||
RslRlBaseRunnerCfg,
|
||||
RslRlVecEnvWrapper,
|
||||
export_policy_as_jit,
|
||||
export_policy_as_onnx,
|
||||
handle_deprecated_rsl_rl_cfg,
|
||||
)
|
||||
from isaaclab_rl.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
|
||||
|
||||
import isaaclab_tasks # noqa: F401
|
||||
from isaaclab_tasks.utils import (
|
||||
add_launcher_args,
|
||||
get_checkpoint_path,
|
||||
launch_simulation,
|
||||
setup_preset_cli,
|
||||
)
|
||||
from isaaclab_tasks.utils.hydra import hydra_task_config
|
||||
|
||||
# local imports
|
||||
import cli_args # isort: skip
|
||||
|
||||
import dex_workbench.tasks # noqa: F401
|
||||
with contextlib.suppress(ImportError):
|
||||
import isaaclab_tasks_experimental # noqa: F401
|
||||
|
||||
# -- argparse ----------------------------------------------------------------
|
||||
parser = argparse.ArgumentParser(description="Train an RL agent with RSL-RL.")
|
||||
parser.add_argument("--video", action="store_true", default=False, help="Record videos during training.")
|
||||
parser.add_argument("--video_length", type=int, default=200, help="Length of the recorded video (in steps).")
|
||||
parser.add_argument(
|
||||
"--disable_fabric", action="store_true", default=False, help="Disable fabric and use USD I/O operations."
|
||||
)
|
||||
parser.add_argument("--num_envs", type=int, default=None, help="Number of environments to simulate.")
|
||||
parser.add_argument("--task", type=str, default=None, help="Name of the task.")
|
||||
parser.add_argument(
|
||||
"--agent", type=str, default="rsl_rl_cfg_entry_point", help="Name of the RL agent configuration entry point."
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed used for the environment")
|
||||
parser.add_argument(
|
||||
"--use_pretrained_checkpoint",
|
||||
action="store_true",
|
||||
help="Use the pre-trained checkpoint from Nucleus.",
|
||||
)
|
||||
parser.add_argument("--real-time", action="store_true", default=False, help="Run in real-time, if possible.")
|
||||
parser.add_argument("--external_callback", default=None, help="Fully qualified path to an externally defined callback.")
|
||||
cli_args.add_rsl_rl_args(parser)
|
||||
add_launcher_args(parser)
|
||||
args_cli, remaining_args = setup_preset_cli(parser)
|
||||
|
||||
if args_cli.video:
|
||||
args_cli.enable_cameras = True
|
||||
|
||||
|
||||
# Call an external callback if requested. This gives opportunity to external code to register the environments
|
||||
# The function is expected to return a list of arguments that were not consumed by the callback.
|
||||
remaining_args_env_registration = None
|
||||
if args_cli.external_callback:
|
||||
external_callback_function = string_to_callable(args_cli.external_callback, separator=".")
|
||||
remaining_args_env_registration = external_callback_function()
|
||||
|
||||
# clear out sys.argv for Hydra
|
||||
# The remaining arguments are the arguments that were not consumed by both this scripts
|
||||
# argparser and (optionally) the external callback function. Both sides of this
|
||||
# intersection are pre-fold (the callback reads the user's original sys.argv), so
|
||||
# preset tokens like ``physics=NAME`` compare correctly here. Fold runs after.
|
||||
remaining_args = list_intersection(remaining_args, remaining_args_env_registration)
|
||||
sys.argv = [sys.argv[0]] + remaining_args
|
||||
|
||||
# Check for installed RSL-RL version
|
||||
installed_version = metadata.version("rsl-rl-lib")
|
||||
|
||||
|
||||
@hydra_task_config(args_cli.task, args_cli.agent)
|
||||
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
|
||||
"""Play with RSL-RL agent."""
|
||||
with launch_simulation(env_cfg, args_cli):
|
||||
# grab task name for checkpoint path
|
||||
task_name = args_cli.task.split(":")[-1]
|
||||
train_task_name = task_name.replace("-Play", "")
|
||||
|
||||
# override configurations with non-hydra CLI arguments
|
||||
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
|
||||
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs
|
||||
|
||||
# handle deprecated configurations
|
||||
agent_cfg = handle_deprecated_rsl_rl_cfg(agent_cfg, installed_version)
|
||||
|
||||
# set the environment seed
|
||||
# note: certain randomizations occur in the environment initialization so we set the seed here
|
||||
env_cfg.seed = agent_cfg.seed
|
||||
env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device
|
||||
|
||||
# specify directory for logging experiments
|
||||
log_root_path = os.path.join("logs", "rsl_rl", agent_cfg.experiment_name)
|
||||
log_root_path = os.path.abspath(log_root_path)
|
||||
print(f"[INFO] Loading experiment from directory: {log_root_path}")
|
||||
if args_cli.use_pretrained_checkpoint:
|
||||
resume_path = get_published_pretrained_checkpoint("rsl_rl", train_task_name)
|
||||
if not resume_path:
|
||||
print("[INFO] Unfortunately a pre-trained checkpoint is currently unavailable for this task.")
|
||||
return
|
||||
elif args_cli.checkpoint:
|
||||
resume_path = retrieve_file_path(args_cli.checkpoint)
|
||||
else:
|
||||
resume_path = get_checkpoint_path(log_root_path, agent_cfg.load_run, agent_cfg.load_checkpoint)
|
||||
|
||||
log_dir = os.path.dirname(resume_path)
|
||||
|
||||
# set the log directory for the environment
|
||||
env_cfg.log_dir = log_dir
|
||||
|
||||
# create isaac environment
|
||||
env = gym.make(args_cli.task, cfg=env_cfg, render_mode="rgb_array" if args_cli.video else None)
|
||||
|
||||
# convert to single-agent instance if required by the RL algorithm
|
||||
if isinstance(env.unwrapped.cfg, DirectMARLEnvCfg):
|
||||
from isaaclab.envs import multi_agent_to_single_agent
|
||||
|
||||
env = multi_agent_to_single_agent(env)
|
||||
|
||||
# wrap for video recording
|
||||
if args_cli.video:
|
||||
video_kwargs = {
|
||||
"video_folder": os.path.join(log_dir, "videos", "play"),
|
||||
"step_trigger": lambda step: step == 0,
|
||||
"video_length": args_cli.video_length,
|
||||
"disable_logger": True,
|
||||
}
|
||||
print("[INFO] Recording videos during training.")
|
||||
print_dict(video_kwargs, nesting=4)
|
||||
env = gym.wrappers.RecordVideo(env, **video_kwargs)
|
||||
|
||||
# wrap around environment for rsl-rl
|
||||
env = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
# load previously trained model
|
||||
if agent_cfg.class_name == "OnPolicyRunner":
|
||||
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
|
||||
elif agent_cfg.class_name == "DistillationRunner":
|
||||
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
|
||||
# configure_seed must be called after runner construction so that PyTorch deterministic settings
|
||||
# do not interfere with the runner's internal initialization.
|
||||
if args_cli.deterministic:
|
||||
configure_seed(env_cfg.seed, True)
|
||||
runner.load(resume_path)
|
||||
|
||||
# obtain the trained policy for inference
|
||||
policy = runner.get_inference_policy(device=env.unwrapped.device)
|
||||
|
||||
# export the trained policy to JIT and ONNX formats
|
||||
export_model_dir = os.path.join(os.path.dirname(resume_path), "exported")
|
||||
|
||||
if version.parse(installed_version) >= version.parse("4.0.0"):
|
||||
# use the new export functions for rsl-rl >= 4.0.0
|
||||
runner.export_policy_to_jit(path=export_model_dir, filename="policy.pt")
|
||||
runner.export_policy_to_onnx(path=export_model_dir, filename="policy.onnx")
|
||||
policy_nn = None # Not needed for rsl-rl >= 4.0.0
|
||||
else:
|
||||
# extract the neural network for rsl-rl < 4.0.0
|
||||
if version.parse(installed_version) >= version.parse("2.3.0"):
|
||||
policy_nn = runner.alg.policy
|
||||
else:
|
||||
policy_nn = runner.alg.actor_critic
|
||||
|
||||
# extract the normalizer
|
||||
if hasattr(policy_nn, "actor_obs_normalizer"):
|
||||
normalizer = policy_nn.actor_obs_normalizer
|
||||
elif hasattr(policy_nn, "student_obs_normalizer"):
|
||||
normalizer = policy_nn.student_obs_normalizer
|
||||
else:
|
||||
normalizer = None
|
||||
|
||||
# export to JIT and ONNX
|
||||
export_policy_as_jit(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.pt")
|
||||
export_policy_as_onnx(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.onnx")
|
||||
|
||||
dt = env.unwrapped.step_dt
|
||||
|
||||
# reset environment
|
||||
obs = env.get_observations()
|
||||
timestep = 0
|
||||
# simulate environment
|
||||
try:
|
||||
while True:
|
||||
start_time = time.time()
|
||||
# run everything in inference mode
|
||||
with torch.inference_mode():
|
||||
# agent stepping
|
||||
actions = policy(obs)
|
||||
# env stepping
|
||||
obs, _, dones, _ = env.step(actions)
|
||||
# reset recurrent states for episodes that have terminated
|
||||
if version.parse(installed_version) >= version.parse("4.0.0"):
|
||||
policy.reset(dones)
|
||||
else:
|
||||
policy_nn.reset(dones)
|
||||
if args_cli.video:
|
||||
timestep += 1
|
||||
if timestep == args_cli.video_length:
|
||||
break
|
||||
|
||||
sleep_time = dt - (time.time() - start_time)
|
||||
if args_cli.real_time and sleep_time > 0:
|
||||
time.sleep(sleep_time)
|
||||
|
||||
# close the simulator
|
||||
env.close()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,234 @@
|
||||
# Copyright (c) 2022-2026, The Isaac Lab Project Developers (https://github.com/isaac-sim/IsaacLab/blob/main/CONTRIBUTORS.md).
|
||||
# All rights reserved.
|
||||
#
|
||||
# SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
"""Script to play a checkpoint of an RL agent from RSL-RL."""
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import importlib.metadata as metadata
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
import gymnasium as gym
|
||||
import torch
|
||||
from packaging import version
|
||||
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
|
||||
|
||||
from isaaclab.envs import DirectMARLEnvCfg, DirectRLEnvCfg, ManagerBasedRLEnvCfg
|
||||
from isaaclab.utils.assets import retrieve_file_path
|
||||
from isaaclab.utils.dict import print_dict
|
||||
from isaaclab.utils.string import list_intersection, string_to_callable
|
||||
|
||||
from isaaclab_rl.rsl_rl import (
|
||||
RslRlBaseRunnerCfg,
|
||||
RslRlVecEnvWrapper,
|
||||
export_policy_as_jit,
|
||||
export_policy_as_onnx,
|
||||
handle_deprecated_rsl_rl_cfg,
|
||||
)
|
||||
from isaaclab_rl.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
|
||||
|
||||
import isaaclab_tasks # noqa: F401
|
||||
from isaaclab_tasks.utils import (
|
||||
add_launcher_args,
|
||||
get_checkpoint_path,
|
||||
launch_simulation,
|
||||
setup_preset_cli,
|
||||
)
|
||||
from isaaclab_tasks.utils.hydra import hydra_task_config
|
||||
|
||||
# local imports
|
||||
import cli_args # isort: skip
|
||||
|
||||
import dex_workbench.tasks # noqa: F401
|
||||
with contextlib.suppress(ImportError):
|
||||
import isaaclab_tasks_experimental # noqa: F401
|
||||
|
||||
# -- argparse ----------------------------------------------------------------
|
||||
parser = argparse.ArgumentParser(description="Play a checkpoint of an RL agent from RSL-RL.")
|
||||
parser.add_argument("--video", action="store_true", default=False, help="Record videos during play.")
|
||||
parser.add_argument("--video_length", type=int, default=200, help="Length of the recorded video (in steps).")
|
||||
parser.add_argument(
|
||||
"--disable_fabric", action="store_true", default=False, help="Disable fabric and use USD I/O operations."
|
||||
)
|
||||
parser.add_argument("--num_envs", type=int, default=None, help="Number of environments to simulate.")
|
||||
parser.add_argument("--task", type=str, default=None, help="Name of the task.")
|
||||
parser.add_argument(
|
||||
"--agent", type=str, default="rsl_rl_cfg_entry_point", help="Name of the RL agent configuration entry point."
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed used for the environment")
|
||||
parser.add_argument(
|
||||
"--use_pretrained_checkpoint",
|
||||
action="store_true",
|
||||
help="Use the pre-trained checkpoint from Nucleus.",
|
||||
)
|
||||
parser.add_argument("--real-time", action="store_true", default=False, help="Run in real-time, if possible.")
|
||||
parser.add_argument("--external_callback", default=None, help="Fully qualified path to an externally defined callback.")
|
||||
cli_args.add_rsl_rl_args(parser)
|
||||
add_launcher_args(parser)
|
||||
args_cli, remaining_args = setup_preset_cli(parser)
|
||||
|
||||
if args_cli.video:
|
||||
args_cli.enable_cameras = True
|
||||
|
||||
|
||||
# Call an external callback if requested. This gives opportunity to external code to register the environments
|
||||
# The function is expected to return a list of arguments that were not consumed by the callback.
|
||||
remaining_args_env_registration = None
|
||||
if args_cli.external_callback:
|
||||
external_callback_function = string_to_callable(args_cli.external_callback, separator=".")
|
||||
remaining_args_env_registration = external_callback_function()
|
||||
|
||||
# clear out sys.argv for Hydra
|
||||
# The remaining arguments are the arguments that were not consumed by both this scripts
|
||||
# argparser and (optionally) the external callback function.
|
||||
remaining_args = list_intersection(remaining_args, remaining_args_env_registration)
|
||||
sys.argv = [sys.argv[0]] + remaining_args
|
||||
|
||||
# Check for installed RSL-RL version
|
||||
installed_version = metadata.version("rsl-rl-lib")
|
||||
|
||||
|
||||
@hydra_task_config(args_cli.task, args_cli.agent)
|
||||
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
|
||||
"""Play with RSL-RL agent."""
|
||||
with launch_simulation(env_cfg, args_cli):
|
||||
# grab task name for checkpoint path
|
||||
task_name = args_cli.task.split(":")[-1]
|
||||
train_task_name = task_name.replace("-Play", "")
|
||||
|
||||
# override configurations with non-hydra CLI arguments
|
||||
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
|
||||
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs
|
||||
|
||||
# handle deprecated configurations
|
||||
agent_cfg = handle_deprecated_rsl_rl_cfg(agent_cfg, installed_version)
|
||||
|
||||
# set the environment seed
|
||||
# note: certain randomizations occur in the environment initialization so we set the seed here
|
||||
env_cfg.seed = agent_cfg.seed
|
||||
env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device
|
||||
|
||||
# specify directory for logging experiments
|
||||
log_root_path = os.path.join("logs", "rsl_rl", agent_cfg.experiment_name)
|
||||
log_root_path = os.path.abspath(log_root_path)
|
||||
print(f"[INFO] Loading experiment from directory: {log_root_path}")
|
||||
if args_cli.use_pretrained_checkpoint:
|
||||
resume_path = get_published_pretrained_checkpoint("rsl_rl", train_task_name)
|
||||
if not resume_path:
|
||||
print("[INFO] Unfortunately a pre-trained checkpoint is currently unavailable for this task.")
|
||||
return
|
||||
elif args_cli.checkpoint:
|
||||
resume_path = retrieve_file_path(args_cli.checkpoint)
|
||||
else:
|
||||
resume_path = get_checkpoint_path(log_root_path, agent_cfg.load_run, agent_cfg.load_checkpoint)
|
||||
|
||||
log_dir = os.path.dirname(resume_path)
|
||||
|
||||
# set the log directory for the environment
|
||||
env_cfg.log_dir = log_dir
|
||||
|
||||
# create isaac environment
|
||||
env = gym.make(args_cli.task, cfg=env_cfg, render_mode="rgb_array" if args_cli.video else None)
|
||||
|
||||
# convert to single-agent instance if required by the RL algorithm
|
||||
if isinstance(env.unwrapped.cfg, DirectMARLEnvCfg):
|
||||
from isaaclab.envs import multi_agent_to_single_agent
|
||||
|
||||
env = multi_agent_to_single_agent(env)
|
||||
|
||||
# wrap for video recording
|
||||
if args_cli.video:
|
||||
video_kwargs = {
|
||||
"video_folder": os.path.join(log_dir, "videos", "play"),
|
||||
"step_trigger": lambda step: step == 0,
|
||||
"video_length": args_cli.video_length,
|
||||
"disable_logger": True,
|
||||
}
|
||||
print("[INFO] Recording videos during play.")
|
||||
print_dict(video_kwargs, nesting=4)
|
||||
env = gym.wrappers.RecordVideo(env, **video_kwargs)
|
||||
|
||||
# wrap around environment for rsl-rl
|
||||
env = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
# load previously trained model
|
||||
if agent_cfg.class_name == "OnPolicyRunner":
|
||||
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
|
||||
elif agent_cfg.class_name == "DistillationRunner":
|
||||
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
|
||||
runner.load(resume_path)
|
||||
|
||||
# obtain the trained policy for inference
|
||||
policy = runner.get_inference_policy(device=env.unwrapped.device)
|
||||
|
||||
# export the trained policy to JIT and ONNX formats
|
||||
export_model_dir = os.path.join(os.path.dirname(resume_path), "exported")
|
||||
|
||||
if version.parse(installed_version) >= version.parse("4.0.0"):
|
||||
# use the new export functions for rsl-rl >= 4.0.0
|
||||
runner.export_policy_to_jit(path=export_model_dir, filename="policy.pt")
|
||||
runner.export_policy_to_onnx(path=export_model_dir, filename="policy.onnx")
|
||||
policy_nn = None # Not needed for rsl-rl >= 4.0.0
|
||||
else:
|
||||
# extract the neural network for rsl-rl < 4.0.0
|
||||
if version.parse(installed_version) >= version.parse("2.3.0"):
|
||||
policy_nn = runner.alg.policy
|
||||
else:
|
||||
policy_nn = runner.alg.actor_critic
|
||||
|
||||
# extract the normalizer
|
||||
if hasattr(policy_nn, "actor_obs_normalizer"):
|
||||
normalizer = policy_nn.actor_obs_normalizer
|
||||
elif hasattr(policy_nn, "student_obs_normalizer"):
|
||||
normalizer = policy_nn.student_obs_normalizer
|
||||
else:
|
||||
normalizer = None
|
||||
|
||||
# export to JIT and ONNX
|
||||
export_policy_as_jit(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.pt")
|
||||
export_policy_as_onnx(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.onnx")
|
||||
|
||||
dt = env.unwrapped.step_dt
|
||||
|
||||
# reset environment
|
||||
obs = env.get_observations()
|
||||
timestep = 0
|
||||
# simulate environment
|
||||
try:
|
||||
while True:
|
||||
start_time = time.time()
|
||||
# run everything in inference mode
|
||||
with torch.inference_mode():
|
||||
# agent stepping
|
||||
actions = policy(obs)
|
||||
# env stepping
|
||||
obs, _, dones, _ = env.step(actions)
|
||||
# reset recurrent states for episodes that have terminated
|
||||
if version.parse(installed_version) >= version.parse("4.0.0"):
|
||||
policy.reset(dones)
|
||||
else:
|
||||
policy_nn.reset(dones)
|
||||
if args_cli.video:
|
||||
timestep += 1
|
||||
if timestep == args_cli.video_length:
|
||||
break
|
||||
|
||||
sleep_time = dt - (time.time() - start_time)
|
||||
if args_cli.real_time and sleep_time > 0:
|
||||
time.sleep(sleep_time)
|
||||
|
||||
# close the simulator
|
||||
env.close()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,255 @@
|
||||
# Copyright (c) 2022-2026, The Isaac Lab Project Developers (https://github.com/isaac-sim/IsaacLab/blob/main/CONTRIBUTORS.md).
|
||||
# All rights reserved.
|
||||
#
|
||||
# SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
"""Script to train RL agent with RSL-RL."""
|
||||
|
||||
import warnings
|
||||
|
||||
warnings.warn(
|
||||
"scripts/reinforcement_learning/rsl_rl/train.py is deprecated. Use "
|
||||
"`./isaaclab.sh train --rl_library rsl_rl --task <TASK>` instead. "
|
||||
"Example: `./isaaclab.sh train --rl_library rsl_rl --task Isaac-Cartpole-v0`.",
|
||||
DeprecationWarning,
|
||||
stacklevel=1,
|
||||
)
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import importlib.metadata as metadata
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime
|
||||
|
||||
import gymnasium as gym
|
||||
import torch
|
||||
from packaging import version
|
||||
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
|
||||
|
||||
from isaaclab.envs import DirectMARLEnvCfg, DirectRLEnvCfg, ManagerBasedRLEnvCfg
|
||||
from isaaclab.utils.dict import print_dict
|
||||
from isaaclab.utils.io import dump_yaml
|
||||
from isaaclab.utils.seed import configure_seed
|
||||
from isaaclab.utils.string import list_intersection, string_to_callable
|
||||
|
||||
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg, RslRlVecEnvWrapper, handle_deprecated_rsl_rl_cfg
|
||||
|
||||
import isaaclab_tasks # noqa: F401
|
||||
from isaaclab_tasks.utils import (
|
||||
add_launcher_args,
|
||||
get_checkpoint_path,
|
||||
launch_simulation,
|
||||
setup_preset_cli,
|
||||
)
|
||||
from isaaclab_tasks.utils.hydra import hydra_task_config
|
||||
|
||||
# local imports
|
||||
import cli_args # isort: skip
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
import dex_workbench.tasks # noqa: F401
|
||||
with contextlib.suppress(ImportError):
|
||||
import isaaclab_tasks_experimental # noqa: F401
|
||||
|
||||
RSL_RL_VERSION = "5.0.1"
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
torch.backends.cudnn.deterministic = False
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
# -- argparse ----------------------------------------------------------------
|
||||
parser = argparse.ArgumentParser(description="Train an RL agent with RSL-RL.")
|
||||
parser.add_argument("--video", action="store_true", default=False, help="Record videos during training.")
|
||||
parser.add_argument("--video_length", type=int, default=200, help="Length of the recorded video (in steps).")
|
||||
parser.add_argument("--video_interval", type=int, default=2000, help="Interval between video recordings (in steps).")
|
||||
parser.add_argument("--num_envs", type=int, default=None, help="Number of environments to simulate.")
|
||||
parser.add_argument("--task", type=str, default=None, help="Name of the task.")
|
||||
parser.add_argument(
|
||||
"--agent", type=str, default="rsl_rl_cfg_entry_point", help="Name of the RL agent configuration entry point."
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=None, help="Seed used for the environment")
|
||||
parser.add_argument("--max_iterations", type=int, default=None, help="RL Policy training iterations.")
|
||||
parser.add_argument(
|
||||
"--distributed", action="store_true", default=False, help="Run training with multiple GPUs or nodes."
|
||||
)
|
||||
parser.add_argument("--export_io_descriptors", action="store_true", default=False, help="Export IO descriptors.")
|
||||
parser.add_argument(
|
||||
"--ray-proc-id", "-rid", type=int, default=None, help="Automatically configured by Ray integration, otherwise None."
|
||||
)
|
||||
parser.add_argument("--external_callback", default=None, help="Fully qualified path to an externally defined callback.")
|
||||
cli_args.add_rsl_rl_args(parser)
|
||||
add_launcher_args(parser)
|
||||
args_cli, remaining_args = setup_preset_cli(parser)
|
||||
|
||||
if args_cli.video:
|
||||
args_cli.enable_cameras = True
|
||||
|
||||
|
||||
# Call an external callback if requested. This gives opportunity to external code to register the environments
|
||||
# The function is expected to return a list of arguments that were not consumed by the callback.
|
||||
remaining_args_env_registration = None
|
||||
if args_cli.external_callback:
|
||||
external_callback_function = string_to_callable(args_cli.external_callback, separator=".")
|
||||
remaining_args_env_registration = external_callback_function()
|
||||
|
||||
# clear out sys.argv for Hydra
|
||||
# The remaining arguments are the arguments that were not consumed by both this scripts
|
||||
# argparser and (optionally) the external callback function. Both sides of this
|
||||
# intersection share the same token vocabulary (the callback reads the user's
|
||||
# original sys.argv), so preset tokens like ``physics=NAME`` compare correctly.
|
||||
remaining_args = list_intersection(remaining_args, remaining_args_env_registration)
|
||||
sys.argv = [sys.argv[0]] + remaining_args
|
||||
|
||||
# -- check RSL-RL version ----------------------------------------------------
|
||||
installed_version = metadata.version("rsl-rl-lib")
|
||||
if version.parse(installed_version) < version.parse(RSL_RL_VERSION):
|
||||
if platform.system() == "Windows":
|
||||
cmd = [r".\isaaclab.bat", "-p", "-m", "pip", "install", f"rsl-rl-lib=={RSL_RL_VERSION}"]
|
||||
else:
|
||||
cmd = ["./isaaclab.sh", "-p", "-m", "pip", "install", f"rsl-rl-lib=={RSL_RL_VERSION}"]
|
||||
print(
|
||||
f"Please install the correct version of RSL-RL.\nExisting version is: '{installed_version}'"
|
||||
f" and required version is: '{RSL_RL_VERSION}'.\nTo install the correct version, run:"
|
||||
f"\n\n\t{' '.join(cmd)}\n"
|
||||
)
|
||||
exit(1)
|
||||
|
||||
|
||||
@hydra_task_config(args_cli.task, args_cli.agent)
|
||||
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
|
||||
"""Train with RSL-RL agent."""
|
||||
with launch_simulation(env_cfg, args_cli):
|
||||
# override configurations with non-hydra CLI arguments
|
||||
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
|
||||
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs
|
||||
agent_cfg.max_iterations = (
|
||||
args_cli.max_iterations if args_cli.max_iterations is not None else agent_cfg.max_iterations
|
||||
)
|
||||
|
||||
# handle deprecated configurations
|
||||
agent_cfg = handle_deprecated_rsl_rl_cfg(agent_cfg, installed_version)
|
||||
|
||||
# set the environment seed
|
||||
# note: certain randomizations occur in the environment initialization so we set the seed here
|
||||
env_cfg.seed = agent_cfg.seed
|
||||
# For distributed training, launch_simulation() already resolved the
|
||||
# correct per-rank device; only apply a CLI --device override for
|
||||
# non-distributed runs (the default "cuda:0" would clobber the
|
||||
# per-rank device otherwise).
|
||||
if not args_cli.distributed:
|
||||
env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device
|
||||
# check for invalid combination of CPU device with distributed training
|
||||
if args_cli.distributed and args_cli.device is not None and "cpu" in args_cli.device:
|
||||
raise ValueError(
|
||||
"Distributed training is not supported when using CPU device. "
|
||||
"Please use GPU device (e.g., --device cuda) for distributed training."
|
||||
)
|
||||
|
||||
# multi-gpu training configuration
|
||||
if args_cli.distributed:
|
||||
global_rank = int(os.getenv("RANK", "0"))
|
||||
# env_cfg.sim.device is resolved by launch_simulation() which
|
||||
# accounts for CUDA_VISIBLE_DEVICES restrictions.
|
||||
agent_cfg.device = env_cfg.sim.device
|
||||
|
||||
# use global rank for seed diversity across all nodes
|
||||
seed = agent_cfg.seed + global_rank
|
||||
env_cfg.seed = seed
|
||||
agent_cfg.seed = seed
|
||||
|
||||
# specify directory for logging experiments
|
||||
log_root_path = os.path.join("logs", "rsl_rl", agent_cfg.experiment_name)
|
||||
log_root_path = os.path.abspath(log_root_path)
|
||||
print(f"[INFO] Logging experiment in directory: {log_root_path}")
|
||||
# specify directory for logging runs: {time-stamp}_{run_name}
|
||||
log_dir = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
# The Ray Tune workflow extracts experiment name using the logging line below, hence, do not
|
||||
# change it (see PR #2346, comment-2819298849)
|
||||
print(f"Exact experiment name requested from command line: {log_dir}")
|
||||
if agent_cfg.run_name:
|
||||
log_dir += f"_{agent_cfg.run_name}"
|
||||
log_dir = os.path.join(log_root_path, log_dir)
|
||||
|
||||
# set the IO descriptors export flag if requested
|
||||
if isinstance(env_cfg, ManagerBasedRLEnvCfg):
|
||||
env_cfg.export_io_descriptors = args_cli.export_io_descriptors
|
||||
else:
|
||||
logger.warning(
|
||||
"IO descriptors are only supported for manager based RL environments."
|
||||
" No IO descriptors will be exported."
|
||||
)
|
||||
|
||||
# set the log directory for the environment (works for all environment types)
|
||||
env_cfg.log_dir = log_dir
|
||||
|
||||
# create isaac environment
|
||||
env = gym.make(args_cli.task, cfg=env_cfg, render_mode="rgb_array" if args_cli.video else None)
|
||||
|
||||
# convert to single-agent instance if required by the RL algorithm
|
||||
if isinstance(env.unwrapped.cfg, DirectMARLEnvCfg):
|
||||
from isaaclab.envs import multi_agent_to_single_agent
|
||||
|
||||
env = multi_agent_to_single_agent(env)
|
||||
|
||||
# save resume path before creating a new log_dir
|
||||
if agent_cfg.resume or agent_cfg.algorithm.class_name == "Distillation":
|
||||
resume_path = get_checkpoint_path(log_root_path, agent_cfg.load_run, agent_cfg.load_checkpoint)
|
||||
|
||||
# wrap for video recording
|
||||
if args_cli.video:
|
||||
video_kwargs = {
|
||||
"video_folder": os.path.join(log_dir, "videos", "train"),
|
||||
"step_trigger": lambda step: step % args_cli.video_interval == 0,
|
||||
"video_length": args_cli.video_length,
|
||||
"disable_logger": True,
|
||||
}
|
||||
print("[INFO] Recording videos during training.")
|
||||
print_dict(video_kwargs, nesting=4)
|
||||
env = gym.wrappers.RecordVideo(env, **video_kwargs)
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# wrap around environment for rsl-rl
|
||||
env = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
|
||||
# create runner from rsl-rl
|
||||
if agent_cfg.class_name == "OnPolicyRunner":
|
||||
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
|
||||
elif agent_cfg.class_name == "DistillationRunner":
|
||||
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
|
||||
# configure_seed must be called after runner construction so that PyTorch deterministic settings
|
||||
# do not interfere with the runner's internal initialization.
|
||||
if args_cli.deterministic:
|
||||
configure_seed(env_cfg.seed, True)
|
||||
# write git state to logs
|
||||
runner.add_git_repo_to_log(__file__)
|
||||
# load the checkpoint
|
||||
if agent_cfg.resume or agent_cfg.algorithm.class_name == "Distillation":
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
# load previously trained model
|
||||
runner.load(resume_path)
|
||||
|
||||
# dump the configuration into log-directory
|
||||
dump_yaml(os.path.join(log_dir, "params", "env.yaml"), env_cfg)
|
||||
dump_yaml(os.path.join(log_dir, "params", "agent.yaml"), agent_cfg)
|
||||
|
||||
# run training
|
||||
try:
|
||||
runner.learn(num_learning_iterations=agent_cfg.max_iterations, init_at_random_ep_len=True)
|
||||
print(f"Training time: {round(time.time() - start_time, 2)} seconds")
|
||||
# close the simulator
|
||||
env.close()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,183 @@
|
||||
# Copyright (c) 2022-2026, The Isaac Lab Project Developers (https://github.com/isaac-sim/IsaacLab/blob/main/CONTRIBUTORS.md).
|
||||
# All rights reserved.
|
||||
#
|
||||
# SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
"""RSL-RL training logic for the unified reinforcement learning entrypoint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import importlib.metadata as metadata
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import time
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from common import (
|
||||
add_common_train_args,
|
||||
add_isaaclab_launcher_args,
|
||||
apply_env_overrides,
|
||||
configure_io_descriptors,
|
||||
create_isaaclab_env,
|
||||
dump_train_configs,
|
||||
enable_cameras_for_video,
|
||||
import_local_module,
|
||||
set_hydra_args,
|
||||
validate_distributed_device,
|
||||
wrap_record_video,
|
||||
)
|
||||
from packaging import version
|
||||
|
||||
import isaaclab_tasks # noqa: F401
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
RSL_RL_VERSION = "5.0.1"
|
||||
RL_ROOT = Path(__file__).resolve().parents[1]
|
||||
CLI_ARGS = import_local_module("isaaclab_rsl_rl_cli_args", RL_ROOT / "rsl_rl" / "cli_args.py")
|
||||
|
||||
import dex_workbench.tasks # noqa: F401
|
||||
with contextlib.suppress(ImportError):
|
||||
import isaaclab_tasks_experimental # noqa: F401
|
||||
|
||||
|
||||
def _check_rsl_rl_version() -> str:
|
||||
"""Check that the installed RSL-RL version is supported."""
|
||||
installed_version = metadata.version("rsl-rl-lib")
|
||||
if version.parse(installed_version) < version.parse(RSL_RL_VERSION):
|
||||
if platform.system() == "Windows":
|
||||
cmd = [r".\isaaclab.bat", "-p", "-m", "pip", "install", f"rsl-rl-lib=={RSL_RL_VERSION}"]
|
||||
else:
|
||||
cmd = ["./isaaclab.sh", "-p", "-m", "pip", "install", f"rsl-rl-lib=={RSL_RL_VERSION}"]
|
||||
print(
|
||||
f"Please install the correct version of RSL-RL.\nExisting version is: '{installed_version}'"
|
||||
f" and required version is: '{RSL_RL_VERSION}'.\nTo install the correct version, run:"
|
||||
f"\n\n\t{' '.join(cmd)}\n"
|
||||
)
|
||||
raise SystemExit(1)
|
||||
return installed_version
|
||||
|
||||
|
||||
def _parse_args(argv: list[str]) -> argparse.Namespace:
|
||||
"""Parse RSL-RL training arguments."""
|
||||
from isaaclab.utils.string import list_intersection, string_to_callable
|
||||
|
||||
from isaaclab_tasks.utils import setup_preset_cli
|
||||
|
||||
parser = argparse.ArgumentParser(description="Train an RL agent with RSL-RL.")
|
||||
add_common_train_args(
|
||||
parser,
|
||||
agent_default="rsl_rl_cfg_entry_point",
|
||||
agent_help="Name of the RL agent configuration entry point.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--external_callback",
|
||||
default=None,
|
||||
help="Fully qualified path to an externally defined callback.",
|
||||
)
|
||||
CLI_ARGS.add_rsl_rl_args(parser)
|
||||
add_isaaclab_launcher_args(parser)
|
||||
# setup_preset_cli registers preset-selection help text + runs parse_known_args
|
||||
args_cli, remaining_args = setup_preset_cli(parser, argv)
|
||||
enable_cameras_for_video(args_cli)
|
||||
|
||||
remaining_args_env_registration = None
|
||||
if args_cli.external_callback:
|
||||
external_callback_function = string_to_callable(args_cli.external_callback, separator=".")
|
||||
remaining_args_env_registration = external_callback_function()
|
||||
|
||||
# physics=/renderer=/presets= tokens pass through the remainder for hydra to parse later
|
||||
set_hydra_args(list_intersection(remaining_args, remaining_args_env_registration))
|
||||
return args_cli
|
||||
|
||||
|
||||
def run(argv: list[str]) -> None:
|
||||
"""Train an RSL-RL agent."""
|
||||
import torch
|
||||
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
|
||||
|
||||
from isaaclab.envs import DirectMARLEnvCfg
|
||||
|
||||
from isaaclab_rl.rsl_rl import RslRlVecEnvWrapper, handle_deprecated_rsl_rl_cfg
|
||||
|
||||
from isaaclab_tasks.utils import get_checkpoint_path, launch_simulation, resolve_task_config
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
torch.backends.cudnn.allow_tf32 = True
|
||||
torch.backends.cudnn.deterministic = False
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
args_cli = _parse_args(argv)
|
||||
installed_version = _check_rsl_rl_version()
|
||||
env_cfg, agent_cfg = resolve_task_config(args_cli.task, args_cli.agent)
|
||||
|
||||
with launch_simulation(env_cfg, args_cli):
|
||||
agent_cfg = CLI_ARGS.update_rsl_rl_cfg(agent_cfg, args_cli)
|
||||
apply_env_overrides(args_cli, env_cfg)
|
||||
agent_cfg.max_iterations = (
|
||||
args_cli.max_iterations if args_cli.max_iterations is not None else agent_cfg.max_iterations
|
||||
)
|
||||
|
||||
agent_cfg = handle_deprecated_rsl_rl_cfg(agent_cfg, installed_version)
|
||||
|
||||
env_cfg.seed = agent_cfg.seed
|
||||
validate_distributed_device(args_cli)
|
||||
|
||||
if args_cli.distributed:
|
||||
global_rank = int(os.getenv("RANK", "0"))
|
||||
agent_cfg.device = env_cfg.sim.device
|
||||
|
||||
seed = agent_cfg.seed + global_rank
|
||||
env_cfg.seed = seed
|
||||
agent_cfg.seed = seed
|
||||
|
||||
log_root_path = os.path.abspath(os.path.join("logs", "rsl_rl", agent_cfg.experiment_name))
|
||||
print(f"[INFO] Logging experiment in directory: {log_root_path}")
|
||||
log_dir = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
print(f"Exact experiment name requested from command line: {log_dir}")
|
||||
if agent_cfg.run_name:
|
||||
log_dir += f"_{agent_cfg.run_name}"
|
||||
log_dir = os.path.join(log_root_path, log_dir)
|
||||
|
||||
configure_io_descriptors(env_cfg, args_cli, logger)
|
||||
env_cfg.log_dir = log_dir
|
||||
|
||||
env = create_isaaclab_env(
|
||||
args_cli.task,
|
||||
env_cfg,
|
||||
args_cli,
|
||||
convert_marl_to_single_agent=isinstance(env_cfg, DirectMARLEnvCfg),
|
||||
)
|
||||
|
||||
if agent_cfg.resume or agent_cfg.algorithm.class_name == "Distillation":
|
||||
resume_path = get_checkpoint_path(log_root_path, agent_cfg.load_run, agent_cfg.load_checkpoint)
|
||||
|
||||
env = wrap_record_video(env, log_dir, args_cli)
|
||||
|
||||
start_time = time.time()
|
||||
env = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
|
||||
if agent_cfg.class_name == "OnPolicyRunner":
|
||||
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
|
||||
elif agent_cfg.class_name == "DistillationRunner":
|
||||
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
|
||||
else:
|
||||
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
|
||||
|
||||
runner.add_git_repo_to_log(__file__)
|
||||
if agent_cfg.resume or agent_cfg.algorithm.class_name == "Distillation":
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
runner.load(resume_path)
|
||||
|
||||
dump_train_configs(log_dir, env_cfg, agent_cfg)
|
||||
|
||||
try:
|
||||
runner.learn(num_learning_iterations=agent_cfg.max_iterations, init_at_random_ep_len=True)
|
||||
print(f"Training time: {round(time.time() - start_time, 2)} seconds")
|
||||
env.close()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
Reference in New Issue
Block a user