init commit.
This commit is contained in:
177
scripts/reinforcement_learning/rsl_rl/utils.py
Normal file
177
scripts/reinforcement_learning/rsl_rl/utils.py
Normal file
@@ -0,0 +1,177 @@
|
||||
# base version: IsaacLab/source/isaaclab_rl/isaaclab_rl/rsl_rl/exporter.py
|
||||
|
||||
import copy
|
||||
import os
|
||||
import torch
|
||||
import re
|
||||
import os
|
||||
import sys
|
||||
from typing import NamedTuple
|
||||
|
||||
"Script to log terminal output to a file, stripping ANSI escape codes."
|
||||
class Logger:
|
||||
def __init__(self, filename):
|
||||
self.terminal = sys.stdout
|
||||
os.makedirs(os.path.dirname(filename), exist_ok=True)
|
||||
self.log = open(filename, 'w', encoding='utf-8')
|
||||
|
||||
self.ansi_escape = re.compile(r'\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])')
|
||||
|
||||
def write(self, message):
|
||||
clean_message = self.ansi_escape.sub('', message)
|
||||
|
||||
self.terminal.write(message)
|
||||
self.log.write(clean_message)
|
||||
self.log.flush()
|
||||
|
||||
def flush(self):
|
||||
self.terminal.flush()
|
||||
self.log.flush()
|
||||
|
||||
# Inputs of CTS Policy is a TensorDict with 'policy' and 'single_obs' keys, we simulate this with a NamedTuple.
|
||||
class CTSPolicyInputs(NamedTuple):
|
||||
policy: torch.Tensor
|
||||
single_obs: torch.Tensor
|
||||
|
||||
def export_cts_policy_as_jit(policy: object, actor_obs_normalizer: object | None, single_obs_normalizer: object | None, path: str, filename="policy.pt"):
|
||||
"""Export CTS policy into a Torch JIT file.
|
||||
|
||||
Args:
|
||||
policy: The CTS policy torch module.
|
||||
actor_obs_normalizer: The empirical normalizer module for actor observations. If None, Identity is used.
|
||||
single_obs_normalizer: The empirical normalizer module for single observations. If None, Identity is used.
|
||||
path: The path to the saving directory.
|
||||
filename: The name of exported JIT file. Defaults to "policy.pt".
|
||||
"""
|
||||
policy_exporter = _TorchPolicyExporter(policy, actor_obs_normalizer, single_obs_normalizer)
|
||||
policy_exporter.export(path, filename)
|
||||
|
||||
|
||||
def export_cts_policy_as_onnx(
|
||||
policy: object, path: str, actor_obs_normalizer: object | None = None, single_obs_normalizer: object | None = None, filename="policy.onnx", verbose=False
|
||||
):
|
||||
"""Export CTS policy into a Torch ONNX file.
|
||||
|
||||
Args:
|
||||
policy: The CTS policy torch module.
|
||||
actor_obs_normalizer: The empirical normalizer module for actor observations. If None, Identity is used.
|
||||
single_obs_normalizer: The empirical normalizer module for single observations. If None, Identity is used.
|
||||
path: The path to the saving directory.
|
||||
filename: The name of exported ONNX file. Defaults to "policy.onnx".
|
||||
verbose: Whether to print the model summary. Defaults to False.
|
||||
"""
|
||||
if not os.path.exists(path):
|
||||
os.makedirs(path, exist_ok=True)
|
||||
policy_exporter = _OnnxPolicyExporter(policy, actor_obs_normalizer, single_obs_normalizer, verbose)
|
||||
policy_exporter.export(path, filename)
|
||||
|
||||
|
||||
"""
|
||||
Helper Classes - Private.
|
||||
"""
|
||||
|
||||
|
||||
class _TorchPolicyExporter(torch.nn.Module):
|
||||
"""Exporter of actor-critic into JIT file."""
|
||||
|
||||
def __init__(self, policy, actor_obs_normalizer=None, single_obs_normalizer=None):
|
||||
assert not policy.is_recurrent, "CTS policy should not be recurrent"
|
||||
super().__init__()
|
||||
|
||||
# copy policy parameters
|
||||
if hasattr(policy, "actor"):
|
||||
self.actor = copy.deepcopy(policy.actor)
|
||||
elif hasattr(policy, "student"):
|
||||
self.actor = copy.deepcopy(policy.student)
|
||||
else:
|
||||
raise ValueError("Policy does not have an actor/student module.")
|
||||
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
||||
self.state_dependent_std = policy.state_dependent_std
|
||||
|
||||
# copy normalizer if exists
|
||||
if actor_obs_normalizer:
|
||||
self.actor_obs_normalizer = copy.deepcopy(actor_obs_normalizer)
|
||||
else:
|
||||
self.actor_obs_normalizer = torch.nn.Identity()
|
||||
if single_obs_normalizer:
|
||||
self.single_obs_normalizer = copy.deepcopy(single_obs_normalizer)
|
||||
else:
|
||||
self.single_obs_normalizer = torch.nn.Identity()
|
||||
|
||||
def forward(self, x: CTSPolicyInputs):
|
||||
single_obs = self.single_obs_normalizer(x.single_obs)
|
||||
obs_a = self.actor_obs_normalizer(x.policy)
|
||||
latent, _ = self.student_moe_encoder(obs_a)
|
||||
latent_and_obs = torch.cat([latent, single_obs], dim=-1)
|
||||
if self.state_dependent_std:
|
||||
return self.actor(latent_and_obs)[..., 0, :]
|
||||
else:
|
||||
return self.actor(latent_and_obs)
|
||||
|
||||
@torch.jit.export
|
||||
def reset(self):
|
||||
pass
|
||||
|
||||
def export(self, path, filename):
|
||||
os.makedirs(path, exist_ok=True)
|
||||
path = os.path.join(path, filename)
|
||||
self.to("cpu")
|
||||
traced_script_module = torch.jit.script(self)
|
||||
traced_script_module.save(path)
|
||||
|
||||
|
||||
class _OnnxPolicyExporter(torch.nn.Module):
|
||||
"""Exporter of actor-critic into ONNX file."""
|
||||
|
||||
def __init__(self, policy, actor_obs_normalizer=None, single_obs_normalizer=None, verbose=False):
|
||||
assert not policy.is_recurrent, "CTS policy should not be recurrent"
|
||||
super().__init__()
|
||||
self.verbose = verbose
|
||||
|
||||
# copy policy parameters
|
||||
if hasattr(policy, "actor"):
|
||||
self.actor = copy.deepcopy(policy.actor)
|
||||
elif hasattr(policy, "student"):
|
||||
self.actor = copy.deepcopy(policy.student)
|
||||
else:
|
||||
raise ValueError("Policy does not have an actor/student module.")
|
||||
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
||||
self.num_single_obs = policy.num_single_obs
|
||||
self.num_actor_obs = policy.num_actor_obs
|
||||
self.state_dependent_std = policy.state_dependent_std
|
||||
|
||||
# copy normalizer if exists
|
||||
if actor_obs_normalizer:
|
||||
self.actor_obs_normalizer = copy.deepcopy(actor_obs_normalizer)
|
||||
else:
|
||||
self.actor_obs_normalizer = torch.nn.Identity()
|
||||
if single_obs_normalizer:
|
||||
self.single_obs_normalizer = copy.deepcopy(single_obs_normalizer)
|
||||
else:
|
||||
self.single_obs_normalizer = torch.nn.Identity()
|
||||
|
||||
def forward(self, history, single_obs):
|
||||
single_obs = self.single_obs_normalizer(single_obs)
|
||||
obs_a = self.actor_obs_normalizer(history)
|
||||
latent, _ = self.student_moe_encoder(obs_a)
|
||||
latent_and_obs = torch.cat([latent, single_obs], dim=-1)
|
||||
if self.state_dependent_std:
|
||||
return self.actor(latent_and_obs)[..., 0, :]
|
||||
else:
|
||||
return self.actor(latent_and_obs)
|
||||
|
||||
def export(self, path, filename):
|
||||
self.to("cpu")
|
||||
self.eval()
|
||||
opset_version = 18 # was 11, but it caused problems with linux-aarch, and 18 worked well across all systems.
|
||||
torch.onnx.export(
|
||||
self,
|
||||
(torch.zeros(1, self.num_actor_obs), torch.zeros(1, self.num_single_obs)),
|
||||
os.path.join(path, filename),
|
||||
export_params=True,
|
||||
opset_version=opset_version,
|
||||
verbose=self.verbose,
|
||||
input_names=["obs"],
|
||||
output_names=["actions"],
|
||||
dynamic_axes={},
|
||||
)
|
||||
Reference in New Issue
Block a user