Bot3dData/DEMO/scripts/rsl_rl/play.py

389 lines
16 KiB
Python
Raw Permalink Normal View History

2026-09-06 04:05:13 +00:00
"""Script to play a trained RSL-RL policy."""
import argparse
import os
import sys
from isaaclab.app import AppLauncher
parser = argparse.ArgumentParser(description="Play a trained RSL-RL policy.")
parser.add_argument("--num_envs", type=int, default=1, 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="Agent configuration entry point.")
parser.add_argument("--seed", type=int, default=None, help="Seed used for the environment.")
parser.add_argument("--load_run", type=str, default=".*", help="Run folder name or regex to load from.")
parser.add_argument(
"--checkpoint",
type=str,
default="model_.*.pt",
help="Checkpoint filename regex under logs, or a direct path to a .pt file.",
)
parser.add_argument("--num_steps", type=int, default=500, help="Number of inference steps to simulate.")
parser.add_argument(
"--disable_resets",
action="store_true",
default=False,
help="Disable environment reset during visual play so short or unstable policies do not instantly jump back to the start pose.",
)
parser.add_argument("--export_io_descriptors", action="store_true", default=False, help="Export IO descriptors.")
parser.add_argument("--export_onnx", action="store_true", default=False, help="Export policy.onnx from the loaded checkpoint.")
AppLauncher.add_app_launcher_args(parser)
args_cli, hydra_args = parser.parse_known_args()
sys.argv = [sys.argv[0]] + hydra_args
app_launcher = AppLauncher(args_cli)
simulation_app = app_launcher.app
import gymnasium as gym
import inspect
import os
import torch
from rsl_rl.runners import OnPolicyRunner
try:
from rsl_rl.algorithms import PPO
except ImportError:
PPO = None
try:
from tensordict import TensorDict
except ImportError:
TensorDict = None
from isaaclab.envs import DirectMARLEnv, DirectMARLEnvCfg, DirectRLEnvCfg, ManagerBasedRLEnvCfg, multi_agent_to_single_agent
from isaaclab_rl.rsl_rl import RslRlBaseRunnerCfg
from isaaclab_tasks.utils import get_checkpoint_path
from isaaclab_tasks.utils.hydra import hydra_task_config
import stackforce_simready_2foot_lab.tasks # noqa: F401
import os
import torch
class _StackForceOnnxPolicy(torch.nn.Module):
def __init__(self, policy):
super().__init__()
self.policy = policy
def forward(self, obs):
actions = self.policy(obs)
if isinstance(actions, (tuple, list)):
return actions[0]
if isinstance(actions, dict):
if "actions" in actions:
return actions["actions"]
if "action" in actions:
return actions["action"]
return next(iter(actions.values()))
return actions
def _stackforce_policy_obs_tensor(obs):
if isinstance(obs, dict):
obs = obs["policy"] if "policy" in obs else next(iter(obs.values()))
elif not isinstance(obs, torch.Tensor) and hasattr(obs, "get"):
try:
candidate = obs.get("policy")
except Exception:
candidate = None
if candidate is not None:
obs = candidate
if not isinstance(obs, torch.Tensor):
raise TypeError(f"ONNX export requires a tensor policy observation, got {type(obs)!r}")
if hasattr(obs, "detach"):
obs = obs.detach()
if obs.dim() == 1:
obs = obs.unsqueeze(0)
elif obs.shape[0] > 1:
obs = obs[:1]
return obs.contiguous()
def stackforce_export_policy_as_onnx(policy, obs, output_dir, file_name="policy.onnx", opset=17):
os.makedirs(output_dir, exist_ok=True)
output_path = os.path.join(output_dir, file_name)
sample_obs = _stackforce_policy_obs_tensor(obs)
module = _StackForceOnnxPolicy(policy).to(sample_obs.device).eval()
with torch.no_grad():
torch.onnx.export(
module,
sample_obs,
output_path,
input_names=["obs"],
output_names=["actions"],
dynamic_axes={"obs": {0: "batch"}, "actions": {0: "batch"}},
opset_version=opset,
)
print(f"Exported ONNX policy to: {output_path}")
return output_path
def _runner_uses_obs_groups():
try:
source = inspect.getsource(OnPolicyRunner)
if PPO is not None and hasattr(PPO, "construct_algorithm"):
source += "\n" + inspect.getsource(PPO.construct_algorithm)
except OSError:
return False
return "resolve_obs_groups" in source or '"obs_groups"' in source or "'obs_groups'" in source
def _runner_uses_split_actor_critic():
try:
source = inspect.getsource(OnPolicyRunner)
if PPO is not None and hasattr(PPO, "construct_algorithm"):
source += "\n" + inspect.getsource(PPO.construct_algorithm)
except OSError:
return False
return 'cfg["actor"]' in source or "cfg['actor']" in source or 'cfg["critic"]' in source or "cfg['critic']" in source
def _runner_expects_privileged_step():
try:
source = inspect.getsource(OnPolicyRunner.learn)
except OSError:
return True
return "privileged_obs" in source or "critic_obs" in source
def _format_rsl_rl_obs(obs_dict, use_obs_groups):
if not use_obs_groups:
return obs_dict["policy"]
if TensorDict is not None and not isinstance(obs_dict, TensorDict):
first_obs = next(iter(obs_dict.values()))
return TensorDict(dict(obs_dict), batch_size=[first_obs.shape[0]], device=first_obs.device)
return obs_dict
class LegacyRslRlVecEnvWrapper:
def __init__(self, env, clip_actions=None):
self.env = env
self.clip_actions = clip_actions
self.use_obs_groups = _runner_uses_obs_groups()
self.return_privileged_obs = _runner_expects_privileged_step()
self.num_envs = env.unwrapped.num_envs
self.device = env.unwrapped.device
self.max_episode_length = env.unwrapped.max_episode_length
self.cfg = env.unwrapped.cfg
self.num_actions = gym.spaces.flatdim(env.unwrapped.single_action_space)
obs_dict, extras = self.env.reset()
self.obs_buf = _format_rsl_rl_obs(obs_dict, self.use_obs_groups)
self.privileged_obs_buf = obs_dict.get("critic")
self.num_obs = obs_dict["policy"].shape[-1]
self.num_privileged_obs = self.privileged_obs_buf.shape[-1] if self.privileged_obs_buf is not None else None
self.rew_buf = torch.zeros(self.num_envs, device=self.device)
self.reset_buf = torch.zeros(self.num_envs, dtype=torch.long, device=self.device)
self.episode_length_buf = env.unwrapped.episode_length_buf
self.extras = extras
def get_observations(self):
return self.obs_buf
def get_privileged_observations(self):
return self.privileged_obs_buf
def reset(self, env_ids=None):
del env_ids
obs_dict, extras = self.env.reset()
self.obs_buf = _format_rsl_rl_obs(obs_dict, self.use_obs_groups)
self.privileged_obs_buf = obs_dict.get("critic")
self.extras = extras
if not self.return_privileged_obs:
return self.obs_buf
return self.obs_buf, self.privileged_obs_buf
def step(self, actions):
if self.clip_actions is not None:
actions = torch.clamp(actions, -self.clip_actions, self.clip_actions)
obs_dict, rewards, terminated, truncated, extras = self.env.step(actions)
dones = (terminated | truncated).to(dtype=torch.long)
if not self.env.unwrapped.cfg.is_finite_horizon:
extras["time_outs"] = truncated
episode_log = dict(extras.get("episode", {}))
if torch.any(dones.bool()) and "log" in extras:
episode_log.update(extras["log"])
episode_log["Step_Reward/mean"] = torch.mean(rewards.detach())
episode_log["Step_Reward/abs_mean"] = torch.mean(torch.abs(rewards.detach()))
extras["episode"] = episode_log
self.obs_buf = _format_rsl_rl_obs(obs_dict, self.use_obs_groups)
self.privileged_obs_buf = obs_dict.get("critic")
self.rew_buf = rewards
self.reset_buf = dones
self.extras = extras
if not self.return_privileged_obs:
return self.obs_buf, rewards, dones, extras
return self.obs_buf, self.privileged_obs_buf, rewards, dones, extras
def close(self):
return self.env.close()
def _runner_expects_nested_runner():
try:
source = inspect.getsource(OnPolicyRunner)
except OSError:
return True
return 'train_cfg["runner"]' in source or "train_cfg['runner']" in source
def _runner_uses_nested_class_name():
try:
source = inspect.getsource(OnPolicyRunner)
except OSError:
return False
return (
'algorithm"]["class_name' in source
or "algorithm']['class_name" in source
or 'policy_cfg.pop("class_name")' in source
or "policy_cfg.pop('class_name')" in source
or 'self.policy_cfg.pop("class_name")' in source
or "self.policy_cfg.pop('class_name')" in source
or "resolve_callable" in source
)
def _runner_uses_split_actor_critic():
try:
source = inspect.getsource(OnPolicyRunner)
if PPO is not None and hasattr(PPO, "construct_algorithm"):
source += "\n" + inspect.getsource(PPO.construct_algorithm)
except OSError:
return False
return 'cfg["actor"]' in source or "cfg['actor']" in source or 'cfg["critic"]' in source or "cfg['critic']" in source
def to_compatible_rsl_rl_cfg(agent_cfg):
data = agent_cfg.to_dict() if hasattr(agent_cfg, "to_dict") else dict(agent_cfg)
allowed_policy_keys = {
"actor_hidden_dims",
"critic_hidden_dims",
"activation",
"init_noise_std",
"clip_actions",
"actor_obs_normalization",
"critic_obs_normalization",
}
allowed_algorithm_keys = {
"num_learning_epochs",
"num_mini_batches",
"clip_param",
"gamma",
"lam",
"value_loss_coef",
"entropy_coef",
"learning_rate",
"max_grad_norm",
"use_clipped_value_loss",
"schedule",
"desired_kl",
"use_spo",
}
if "runner" in data and "policy" in data and "algorithm" in data:
runner_cfg = dict(data["runner"])
policy_cfg = {key: value for key, value in dict(data["policy"]).items() if key in allowed_policy_keys or key == "class_name"}
algorithm_cfg = {key: value for key, value in dict(data["algorithm"]).items() if key in allowed_algorithm_keys or key == "class_name"}
else:
policy_cfg = {key: value for key, value in dict(data["policy"]).items() if key in allowed_policy_keys}
algorithm_cfg = {key: value for key, value in dict(data["algorithm"]).items() if key in allowed_algorithm_keys}
runner_cfg = {key: value for key, value in data.items() if key not in {"policy", "algorithm", "class_name"}}
runner_cfg.setdefault("num_steps_per_env", getattr(agent_cfg, "num_steps_per_env", 24))
runner_cfg.setdefault("max_iterations", getattr(agent_cfg, "max_iterations", 1500))
runner_cfg.setdefault("save_interval", getattr(agent_cfg, "save_interval", 50))
runner_cfg.setdefault("obs_groups", {"policy": ["policy"], "critic": ["policy"]})
runner_cfg.setdefault("experiment_name", getattr(agent_cfg, "experiment_name", "stackforce"))
runner_cfg.setdefault("run_name", getattr(agent_cfg, "run_name", ""))
runner_cfg.setdefault("resume", getattr(agent_cfg, "resume", False))
runner_cfg.setdefault("load_run", getattr(agent_cfg, "load_run", ".*"))
runner_cfg.setdefault("checkpoint", getattr(agent_cfg, "load_checkpoint", "model_.*.pt"))
if _runner_uses_nested_class_name():
policy_cfg.setdefault("class_name", "ActorCritic")
algorithm_cfg.setdefault("class_name", "PPO")
else:
runner_cfg.setdefault("policy_class_name", "ActorCritic")
runner_cfg.setdefault("algorithm_class_name", "PPO")
if _runner_expects_nested_runner():
return {"runner": runner_cfg, "policy": policy_cfg, "algorithm": algorithm_cfg}
if _runner_uses_split_actor_critic():
algorithm_cfg.setdefault("class_name", "PPO")
algorithm_cfg.pop("use_spo", None)
actor_cfg = {
"class_name": "MLPModel",
"hidden_dims": policy_cfg.get("actor_hidden_dims", [256, 256, 128]),
"activation": policy_cfg.get("activation", "elu"),
"obs_normalization": policy_cfg.get("actor_obs_normalization", False),
"distribution_cfg": {
"class_name": "GaussianDistribution",
"init_std": policy_cfg.get("init_noise_std", 1.0),
},
}
critic_cfg = {
"class_name": "MLPModel",
"hidden_dims": policy_cfg.get("critic_hidden_dims", [256, 256, 128]),
"activation": policy_cfg.get("activation", "elu"),
"obs_normalization": policy_cfg.get("critic_obs_normalization", False),
}
runner_cfg.pop("policy_class_name", None)
runner_cfg.pop("algorithm_class_name", None)
runner_cfg["obs_groups"] = {"actor": ["policy"], "critic": ["policy"], "policy": ["policy"]}
runner_cfg.setdefault("multi_gpu", None)
return {**runner_cfg, "actor": actor_cfg, "critic": critic_cfg, "algorithm": algorithm_cfg}
return {**runner_cfg, "policy": policy_cfg, "algorithm": algorithm_cfg}
@hydra_task_config(args_cli.task, args_cli.agent)
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
env_cfg.scene.num_envs = args_cli.num_envs if args_cli.num_envs is not None else env_cfg.scene.num_envs
env_cfg.seed = args_cli.seed if args_cli.seed is not None else agent_cfg.seed
env_cfg.sim.device = args_cli.device if args_cli.device is not None else env_cfg.sim.device
if isinstance(env_cfg, ManagerBasedRLEnvCfg):
env_cfg.export_io_descriptors = args_cli.export_io_descriptors
elif args_cli.disable_resets and hasattr(env_cfg, "visual_disable_resets"):
env_cfg.visual_disable_resets = True
log_root_path = os.path.abspath(os.path.join("logs", "rsl_rl", agent_cfg.experiment_name))
checkpoint_arg = args_cli.checkpoint
if os.path.isfile(checkpoint_arg):
resume_path = os.path.abspath(checkpoint_arg)
else:
resume_path = get_checkpoint_path(log_root_path, args_cli.load_run, checkpoint_arg)
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
env = gym.make(args_cli.task, cfg=env_cfg)
if isinstance(env.unwrapped, DirectMARLEnv):
env = multi_agent_to_single_agent(env)
legacy_agent_cfg = to_compatible_rsl_rl_cfg(agent_cfg)
wrapped_env = LegacyRslRlVecEnvWrapper(env, clip_actions=getattr(agent_cfg, "clip_actions", None))
runner = OnPolicyRunner(wrapped_env, legacy_agent_cfg, log_dir=None, device=env.unwrapped.device)
runner.load(resume_path, load_optimizer=False)
policy = runner.get_inference_policy(device=env.unwrapped.device)
obs = wrapped_env.get_observations()
if args_cli.export_onnx:
stackforce_export_policy_as_onnx(
policy,
obs,
os.path.join(os.path.dirname(resume_path), "exported", "policies"),
)
steps = 0
with torch.inference_mode():
while simulation_app.is_running():
actions = policy(obs)
step_result = wrapped_env.step(actions)
obs = step_result[0]
steps += 1
if args_cli.num_steps > 0 and steps >= args_cli.num_steps:
break
wrapped_env.close()
if __name__ == "__main__":
main()
simulation_app.close()