diff --git a/requirements.txt b/requirements.txt index 1f1849d..99ca341 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,6 +4,7 @@ git+https://github.com/ZhengyiLuo/SMPLSim.git@dd65a86 easydict warp-lang dataclass-wizard +onnxscript -e neural_wbc/core -e neural_wbc/data diff --git a/scripts/rsl_rl/players.py b/scripts/rsl_rl/players.py index b006b30..10f20ef 100644 --- a/scripts/rsl_rl/players.py +++ b/scripts/rsl_rl/players.py @@ -70,6 +70,13 @@ def __init__(self, args_cli: argparse.Namespace, randomize: bool, custom_config: student_cfg = StudentPolicyTrainerCfg(**config_dict) student_trainer = StudentPolicyTrainer(env=self.wrapped_env, cfg=student_cfg) self.policy = student_trainer.get_inference_policy(device=self.env.device) + + # export student policy to onnx + onnx_program = torch.onnx.export(self.policy, self.wrapped_env.get_observations(), dynamo=True) + export_model_dir = os.path.join(student_path, "exported") + os.makedirs(export_model_dir, exist_ok=True) + onnx_program.save(os.path.join(export_model_dir, "policy.onnx")) + else: raise ValueError("student_policy.resume_path is needed for play or eval. Please specify a value.") else: