Skip to content

Commit be9ed48

Browse files
committed
feat: adds changes for rsl_rl 3.0.1
1 parent 68b147d commit be9ed48

29 files changed

Lines changed: 181 additions & 104 deletions

File tree

scripts/reinforcement_learning/rsl_rl/cli_args.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
from typing import TYPE_CHECKING
1414

1515
if TYPE_CHECKING:
16-
from isaaclab_rl.rsl_rl import RslRlOnPolicyRunnerCfg
16+
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg
1717

1818

1919
def add_rsl_rl_args(parser: argparse.ArgumentParser):
@@ -42,7 +42,7 @@ def add_rsl_rl_args(parser: argparse.ArgumentParser):
4242
)
4343

4444

45-
def parse_rsl_rl_cfg(task_name: str, args_cli: argparse.Namespace) -> RslRlOnPolicyRunnerCfg:
45+
def parse_rsl_rl_cfg(task_name: str, args_cli: argparse.Namespace) -> RslRlBaseRunnerCfg:
4646
"""Parse configuration for RSL-RL agent based on inputs.
4747
4848
Args:
@@ -55,12 +55,12 @@ def parse_rsl_rl_cfg(task_name: str, args_cli: argparse.Namespace) -> RslRlOnPol
5555
from isaaclab_tasks.utils.parse_cfg import load_cfg_from_registry
5656

5757
# load the default configuration
58-
rslrl_cfg: RslRlOnPolicyRunnerCfg = load_cfg_from_registry(task_name, "rsl_rl_cfg_entry_point")
58+
rslrl_cfg: RslRlBaseRunnerCfg = load_cfg_from_registry(task_name, "rsl_rl_cfg_entry_point")
5959
rslrl_cfg = update_rsl_rl_cfg(rslrl_cfg, args_cli)
6060
return rslrl_cfg
6161

6262

63-
def update_rsl_rl_cfg(agent_cfg: RslRlOnPolicyRunnerCfg, args_cli: argparse.Namespace):
63+
def update_rsl_rl_cfg(agent_cfg: RslRlBaseRunnerCfg, args_cli: argparse.Namespace):
6464
"""Update configuration for RSL-RL agent based on inputs.
6565
6666
Args:

scripts/reinforcement_learning/rsl_rl/play.py

Lines changed: 31 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818

1919
# local imports
2020
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
21-
import cli_args
21+
import cli_args # isort: skip
2222
from rl_utils import camera_follow
2323

2424
# add argparse arguments
@@ -64,7 +64,7 @@
6464
import time
6565
import torch
6666

67-
from rsl_rl.runners import OnPolicyRunner
67+
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
6868

6969
from isaaclab.devices import Se2Keyboard, Se2KeyboardCfg
7070
from isaaclab.envs import (
@@ -78,20 +78,24 @@
7878
from isaaclab.utils.assets import retrieve_file_path
7979
from isaaclab.utils.dict import print_dict
8080
from isaaclab.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
81-
from isaaclab_rl.rsl_rl import RslRlOnPolicyRunnerCfg, RslRlVecEnvWrapper, export_policy_as_jit, export_policy_as_onnx
81+
82+
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg, RslRlVecEnvWrapper, export_policy_as_jit, export_policy_as_onnx
83+
8284
from isaaclab_tasks.utils import get_checkpoint_path
8385
from isaaclab_tasks.utils.hydra import hydra_task_config
8486

8587
import robot_lab.tasks # noqa: F401
8688

8789

8890
@hydra_task_config(args_cli.task, args_cli.agent)
89-
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg):
91+
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
9092
"""Play with RSL-RL agent."""
93+
# grab task name for checkpoint path
9194
task_name = args_cli.task.split(":")[-1]
95+
9296
# override configurations with non-hydra CLI arguments
93-
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
94-
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else 50
97+
agent_cfg: RslRlBaseRunnerCfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
98+
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else 64
9599

96100
# set the environment seed
97101
# note: certain randomizations occur in the environment initialization so we set the seed here
@@ -167,40 +171,43 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
167171

168172
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
169173
# load previously trained model
170-
ppo_runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
171-
ppo_runner.load(resume_path)
174+
if agent_cfg.class_name == "OnPolicyRunner":
175+
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
176+
elif agent_cfg.class_name == "DistillationRunner":
177+
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
178+
else:
179+
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
180+
runner.load(resume_path)
172181

173182
# obtain the trained policy for inference
174-
policy = ppo_runner.get_inference_policy(device=env.unwrapped.device)
183+
policy = runner.get_inference_policy(device=env.unwrapped.device)
175184

176185
# extract the neural network module
177186
# we do this in a try-except to maintain backwards compatibility.
178187
try:
179188
# version 2.3 onwards
180-
policy_nn = ppo_runner.alg.policy
189+
policy_nn = runner.alg.policy
181190
except AttributeError:
182191
# version 2.2 and below
183-
policy_nn = ppo_runner.alg.actor_critic
192+
policy_nn = runner.alg.actor_critic
193+
194+
# extract the normalizer
195+
if hasattr(policy_nn, "actor_obs_normalizer"):
196+
normalizer = policy_nn.actor_obs_normalizer
197+
elif hasattr(policy_nn, "student_obs_normalizer"):
198+
normalizer = policy_nn.student_obs_normalizer
199+
else:
200+
normalizer = None
184201

185202
# export policy to onnx/jit
186203
export_model_dir = os.path.join(os.path.dirname(resume_path), "exported")
187-
export_policy_as_onnx(
188-
policy=policy_nn,
189-
normalizer=ppo_runner.obs_normalizer,
190-
path=export_model_dir,
191-
filename="policy.onnx",
192-
)
193-
export_policy_as_jit(
194-
policy=policy_nn,
195-
normalizer=ppo_runner.obs_normalizer,
196-
path=export_model_dir,
197-
filename="policy.pt",
198-
)
204+
export_policy_as_jit(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.pt")
205+
export_policy_as_onnx(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.onnx")
199206

200207
dt = env.unwrapped.step_dt
201208

202209
# reset environment
203-
obs, _ = env.get_observations()
210+
obs = env.get_observations()
204211
timestep = 0
205212
# simulate environment
206213
while simulation_app.is_running():

scripts/reinforcement_learning/rsl_rl/play_cs.py

Lines changed: 33 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,13 @@
1111
"""Launch Isaac Sim Simulator first."""
1212

1313
import argparse
14-
import os
1514
import sys
1615

1716
from isaaclab.app import AppLauncher
1817

1918
# local imports
2019
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
21-
import cli_args
20+
import cli_args # isort: skip
2221
from rl_utils import camera_follow
2322

2423
# add argparse arguments
@@ -62,10 +61,11 @@
6261
"""Rest everything follows."""
6362

6463
import gymnasium as gym
64+
import os
6565
import time
6666
import torch
6767

68-
from rsl_rl.runners import OnPolicyRunner
68+
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
6969

7070
from isaaclab.devices import Se2Keyboard, Se2KeyboardCfg
7171
from isaaclab.envs import (
@@ -80,25 +80,29 @@
8080
from isaaclab.utils.assets import retrieve_file_path
8181
from isaaclab.utils.dict import print_dict
8282
from isaaclab.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
83-
from isaaclab_rl.rsl_rl import RslRlOnPolicyRunnerCfg, RslRlVecEnvWrapper, export_policy_as_jit, export_policy_as_onnx
83+
84+
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg, RslRlVecEnvWrapper, export_policy_as_jit, export_policy_as_onnx
85+
8486
from isaaclab_tasks.utils import get_checkpoint_path
8587
from isaaclab_tasks.utils.hydra import hydra_task_config
8688

8789
import robot_lab.tasks # noqa: F401
8890

8991

9092
@hydra_task_config(args_cli.task, args_cli.agent)
91-
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg):
93+
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
9294
"""Play with RSL-RL agent."""
95+
# grab task name for checkpoint path
9396
task_name = args_cli.task.split(":")[-1]
97+
9498
# override configurations with non-hydra CLI arguments
95-
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
96-
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs
99+
agent_cfg: RslRlBaseRunnerCfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
100+
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else 64
97101

98102
# set the environment seed
99103
# note: certain randomizations occur in the environment initialization so we set the seed here
100104
env_cfg.seed = agent_cfg.seed
101-
env_cfg.sim.device = args_cli.device if args_cli.device is not None else 50
105+
env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device
102106

103107
# cs map config
104108
env_cfg.scene.terrain = TerrainImporterCfg(
@@ -206,40 +210,43 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
206210

207211
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
208212
# load previously trained model
209-
ppo_runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
210-
ppo_runner.load(resume_path)
213+
if agent_cfg.class_name == "OnPolicyRunner":
214+
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
215+
elif agent_cfg.class_name == "DistillationRunner":
216+
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=None, device=agent_cfg.device)
217+
else:
218+
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
219+
runner.load(resume_path)
211220

212221
# obtain the trained policy for inference
213-
policy = ppo_runner.get_inference_policy(device=env.unwrapped.device)
222+
policy = runner.get_inference_policy(device=env.unwrapped.device)
214223

215224
# extract the neural network module
216225
# we do this in a try-except to maintain backwards compatibility.
217226
try:
218227
# version 2.3 onwards
219-
policy_nn = ppo_runner.alg.policy
228+
policy_nn = runner.alg.policy
220229
except AttributeError:
221230
# version 2.2 and below
222-
policy_nn = ppo_runner.alg.actor_critic
231+
policy_nn = runner.alg.actor_critic
232+
233+
# extract the normalizer
234+
if hasattr(policy_nn, "actor_obs_normalizer"):
235+
normalizer = policy_nn.actor_obs_normalizer
236+
elif hasattr(policy_nn, "student_obs_normalizer"):
237+
normalizer = policy_nn.student_obs_normalizer
238+
else:
239+
normalizer = None
223240

224241
# export policy to onnx/jit
225242
export_model_dir = os.path.join(os.path.dirname(resume_path), "exported")
226-
export_policy_as_onnx(
227-
policy=policy_nn,
228-
normalizer=ppo_runner.obs_normalizer,
229-
path=export_model_dir,
230-
filename="policy.onnx",
231-
)
232-
export_policy_as_jit(
233-
policy=policy_nn,
234-
normalizer=ppo_runner.obs_normalizer,
235-
path=export_model_dir,
236-
filename="policy.pt",
237-
)
243+
export_policy_as_jit(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.pt")
244+
export_policy_as_onnx(policy_nn, normalizer=normalizer, path=export_model_dir, filename="policy.onnx")
238245

239246
dt = env.unwrapped.step_dt
240247

241248
# reset environment
242-
obs, _ = env.get_observations()
249+
obs = env.get_observations()
243250
timestep = 0
244251
# simulate environment
245252
while simulation_app.is_running():

scripts/reinforcement_learning/rsl_rl/train.py

Lines changed: 42 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -17,8 +17,7 @@
1717
from isaaclab.app import AppLauncher
1818

1919
# local imports
20-
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
21-
import cli_args
20+
import cli_args # isort: skip
2221

2322
# add argparse arguments
2423
parser = argparse.ArgumentParser(description="Train an RL agent with RSL-RL.")
@@ -35,6 +34,7 @@
3534
parser.add_argument(
3635
"--distributed", action="store_true", default=False, help="Run training with multiple GPUs or nodes."
3736
)
37+
parser.add_argument("--export_io_descriptors", action="store_true", default=False, help="Export IO descriptors.")
3838
# append RSL-RL cli arguments
3939
cli_args.add_rsl_rl_args(parser)
4040
# append AppLauncher cli args
@@ -52,14 +52,34 @@
5252
app_launcher = AppLauncher(args_cli)
5353
simulation_app = app_launcher.app
5454

55+
"""Check for minimum supported RSL-RL version."""
56+
57+
import importlib.metadata as metadata
58+
import platform
59+
60+
from packaging import version
61+
62+
# check minimum supported rsl-rl version
63+
RSL_RL_VERSION = "3.0.1"
64+
installed_version = metadata.version("rsl-rl-lib")
65+
if version.parse(installed_version) < version.parse(RSL_RL_VERSION):
66+
cmd = [r"python", "-m", "pip", "install", f"rsl-rl-lib=={RSL_RL_VERSION}"]
67+
print(
68+
f"Please install the correct version of RSL-RL.\nExisting version is: '{installed_version}'"
69+
f" and required version is: '{RSL_RL_VERSION}'.\nTo install the correct version, run:"
70+
f"\n\n\t{' '.join(cmd)}\n"
71+
)
72+
exit(1)
73+
5574
"""Rest everything follows."""
5675

5776
import gymnasium as gym
5877
import os
5978
import torch
6079
from datetime import datetime
6180

62-
from rsl_rl.runners import OnPolicyRunner
81+
import omni
82+
from rsl_rl.runners import DistillationRunner, OnPolicyRunner
6383

6484
from isaaclab.envs import (
6585
DirectMARLEnv,
@@ -70,7 +90,9 @@
7090
)
7191
from isaaclab.utils.dict import print_dict
7292
from isaaclab.utils.io import dump_pickle, dump_yaml
73-
from isaaclab_rl.rsl_rl import RslRlOnPolicyRunnerCfg, RslRlVecEnvWrapper
93+
94+
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg, RslRlVecEnvWrapper
95+
7496
from isaaclab_tasks.utils import get_checkpoint_path
7597
from isaaclab_tasks.utils.hydra import hydra_task_config
7698

@@ -83,7 +105,7 @@
83105

84106

85107
@hydra_task_config(args_cli.task, args_cli.agent)
86-
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlOnPolicyRunnerCfg):
108+
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
87109
"""Train with RSL-RL agent."""
88110
# override configurations with non-hydra CLI arguments
89111
agent_cfg = cli_args.update_rsl_rl_cfg(agent_cfg, args_cli)
@@ -119,6 +141,15 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
119141
log_dir += f"_{agent_cfg.run_name}"
120142
log_dir = os.path.join(log_root_path, log_dir)
121143

144+
# set the IO descriptors output directory if requested
145+
if isinstance(env_cfg, ManagerBasedRLEnvCfg):
146+
env_cfg.export_io_descriptors = args_cli.export_io_descriptors
147+
env_cfg.io_descriptors_output_dir = log_dir
148+
else:
149+
omni.log.warn(
150+
"IO descriptors are only supported for manager based RL environments. No IO descriptors will be exported."
151+
)
152+
122153
# create isaac environment
123154
env = gym.make(args_cli.task, cfg=env_cfg, render_mode="rgb_array" if args_cli.video else None)
124155

@@ -146,7 +177,12 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
146177
env = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
147178

148179
# create runner from rsl-rl
149-
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
180+
if agent_cfg.class_name == "OnPolicyRunner":
181+
runner = OnPolicyRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
182+
elif agent_cfg.class_name == "DistillationRunner":
183+
runner = DistillationRunner(env, agent_cfg.to_dict(), log_dir=log_dir, device=agent_cfg.device)
184+
else:
185+
raise ValueError(f"Unsupported runner class: {agent_cfg.class_name}")
150186
# write git state to logs
151187
runner.add_git_repo_to_log(__file__)
152188
# load the checkpoint

source/robot_lab/robot_lab/tasks/manager_based/locomotion/velocity/config/humanoid/booster_t1/agents/rsl_rl_ppo_cfg.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,9 +11,10 @@ class BoosterT1RoughPPORunnerCfg(RslRlOnPolicyRunnerCfg):
1111
max_iterations = 3000
1212
save_interval = 50
1313
experiment_name = "booster_t1_rough"
14-
empirical_normalization = False
1514
policy = RslRlPpoActorCriticCfg(
1615
init_noise_std=1.0,
16+
actor_obs_normalization=False,
17+
critic_obs_normalization=False,
1718
actor_hidden_dims=[512, 256, 128],
1819
critic_hidden_dims=[512, 256, 128],
1920
activation="elu",

source/robot_lab/robot_lab/tasks/manager_based/locomotion/velocity/config/humanoid/fftai_gr1t1/agents/rsl_rl_ppo_cfg.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,9 +11,10 @@ class FFTAIGR1T1RoughPPORunnerCfg(RslRlOnPolicyRunnerCfg):
1111
max_iterations = 3000
1212
save_interval = 100
1313
experiment_name = "fftai_gr1t1_rough"
14-
empirical_normalization = False
1514
policy = RslRlPpoActorCriticCfg(
1615
init_noise_std=1.0,
16+
actor_obs_normalization=False,
17+
critic_obs_normalization=False,
1718
actor_hidden_dims=[512, 256, 128],
1819
critic_hidden_dims=[512, 256, 128],
1920
activation="elu",

0 commit comments

Comments
 (0)