Update ONNX export
This commit is contained in:
@@ -191,8 +191,21 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
raise ValueError("Policy does not have an actor/student module.")
|
raise ValueError("Policy does not have an actor/student module.")
|
||||||
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
||||||
self.num_single_obs = policy.num_single_obs
|
self.num_actions = int(policy.num_actions)
|
||||||
self.num_actor_obs = policy.num_actor_obs
|
self.num_single_obs = int(policy.num_single_obs)
|
||||||
|
self.num_actor_obs = int(policy.num_actor_obs)
|
||||||
|
if self.num_actor_obs % self.num_single_obs != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"num_actor_obs ({self.num_actor_obs}) must be divisible by num_single_obs ({self.num_single_obs})."
|
||||||
|
)
|
||||||
|
self.history_len = self.num_actor_obs // self.num_single_obs
|
||||||
|
# Keep the same per-term history layout as deploy-side push_obs_history:
|
||||||
|
# [ang_vel(3), gravity(3), cmd(3), joint_pos(A), joint_vel(A), last_action(A)].
|
||||||
|
self.feature_dims = [3, 3, 3, self.num_actions, self.num_actions, self.num_actions]
|
||||||
|
if sum(self.feature_dims) != self.num_single_obs:
|
||||||
|
raise ValueError(
|
||||||
|
"Unsupported single_obs layout: expected 3+3+3+3*num_actions to match num_single_obs."
|
||||||
|
)
|
||||||
self.state_dependent_std = policy.state_dependent_std
|
self.state_dependent_std = policy.state_dependent_std
|
||||||
|
|
||||||
# copy normalizer if exists
|
# copy normalizer if exists
|
||||||
@@ -205,7 +218,26 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
self.single_obs_normalizer = torch.nn.Identity()
|
self.single_obs_normalizer = torch.nn.Identity()
|
||||||
|
|
||||||
def forward(self, history, single_obs):
|
def _extract_single_obs_from_history(self, history: torch.Tensor) -> torch.Tensor:
|
||||||
|
if history.dim() == 1:
|
||||||
|
history = history.unsqueeze(0)
|
||||||
|
if history.shape[-1] != self.num_actor_obs:
|
||||||
|
raise ValueError(
|
||||||
|
f"Expected history last dimension {self.num_actor_obs}, got {history.shape[-1]}."
|
||||||
|
)
|
||||||
|
|
||||||
|
single_obs_terms = []
|
||||||
|
offset = 0
|
||||||
|
for dim in self.feature_dims:
|
||||||
|
end = offset + dim * self.history_len
|
||||||
|
single_obs_terms.append(history[:, end - dim:end])
|
||||||
|
offset = end
|
||||||
|
return torch.cat(single_obs_terms, dim=-1)
|
||||||
|
|
||||||
|
def forward(self, history):
|
||||||
|
if history.dim() == 1:
|
||||||
|
history = history.unsqueeze(0)
|
||||||
|
single_obs = self._extract_single_obs_from_history(history)
|
||||||
single_obs = self.single_obs_normalizer(single_obs)
|
single_obs = self.single_obs_normalizer(single_obs)
|
||||||
obs_a = self.actor_obs_normalizer(history)
|
obs_a = self.actor_obs_normalizer(history)
|
||||||
latent, _ = self.student_moe_encoder(obs_a)
|
latent, _ = self.student_moe_encoder(obs_a)
|
||||||
@@ -221,7 +253,7 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
opset_version = 18 # was 11, but it caused problems with linux-aarch, and 18 worked well across all systems.
|
opset_version = 18 # was 11, but it caused problems with linux-aarch, and 18 worked well across all systems.
|
||||||
torch.onnx.export(
|
torch.onnx.export(
|
||||||
self,
|
self,
|
||||||
(torch.zeros(1, self.num_actor_obs), torch.zeros(1, self.num_single_obs)),
|
torch.zeros(1, self.num_actor_obs),
|
||||||
os.path.join(path, filename),
|
os.path.join(path, filename),
|
||||||
export_params=True,
|
export_params=True,
|
||||||
opset_version=opset_version,
|
opset_version=opset_version,
|
||||||
|
|||||||
Reference in New Issue
Block a user