from pathlib import Path import numpy as np import torch import torch.nn as nn def resolve_policy_path(policy_path: Path | str | None, root: Path | None = None) -> Path: """Prefer ONNX for deployment while keeping the PT file as source/fallback.""" root = root or Path(__file__).resolve().parents[1] if policy_path is None: onnx_path = root / "policies" / "model_rough.onnx" pt_path = root / "policies" / "model_rough.pt" return onnx_path if onnx_path.exists() else pt_path path = Path(policy_path) if not path.is_absolute(): path = root / path if path.suffix.lower() == ".pt": onnx_path = path.with_suffix(".onnx") if onnx_path.exists(): return onnx_path return path class PolicyMLP(nn.Module): def __init__(self, obs_dim: int, action_dim: int): super().__init__() self.register_buffer("obs_mean", torch.zeros(obs_dim)) self.register_buffer("obs_std", torch.ones(obs_dim)) self.net = nn.Sequential( nn.Linear(obs_dim, 512), nn.ELU(), nn.Linear(512, 256), nn.ELU(), nn.Linear(256, 128), nn.ELU(), nn.Linear(128, action_dim), ) def forward(self, x: torch.Tensor) -> torch.Tensor: x = (x - self.obs_mean) / torch.clamp(self.obs_std, min=1e-6) return self.net(x) class OnnxPolicy: def __init__(self, model_path: Path): try: import onnxruntime as ort except ImportError as exc: raise ImportError( "onnxruntime is required for ONNX policy inference. " "Install it on Orin with `python -m pip install onnxruntime`." ) from exc opts = ort.SessionOptions() opts.intra_op_num_threads = 1 opts.inter_op_num_threads = 1 opts.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL self.session = ort.InferenceSession( str(model_path), sess_options=opts, providers=["CPUExecutionProvider"], ) self.input_name = self.session.get_inputs()[0].name self.output_name = self.session.get_outputs()[0].name input_shape = self.session.get_inputs()[0].shape output_shape = self.session.get_outputs()[0].shape self.expected_obs_dim = int(input_shape[1]) if len(input_shape) >= 2 and isinstance(input_shape[1], int) else 53 self.expected_action_dim = int(output_shape[1]) if len(output_shape) >= 2 and isinstance(output_shape[1], int) else 16 self.backend = "onnxruntime" self.obs_mean = torch.zeros(self.expected_obs_dim) self.obs_std = torch.ones(self.expected_obs_dim) def __call__(self, x: torch.Tensor) -> torch.Tensor: obs = x.detach().cpu().numpy().astype(np.float32, copy=False) action = self.session.run([self.output_name], {self.input_name: obs})[0] return torch.from_numpy(np.asarray(action, dtype=np.float32)).to(x.device) def load_policy(model_path: Path, device: torch.device): if model_path.suffix.lower() == ".onnx": return OnnxPolicy(model_path) checkpoint = torch.load(model_path, map_location=device, weights_only=False) state_dict = checkpoint["actor_state_dict"] input_key = "mlp.0.weight" if "mlp.0.weight" in state_dict else "net.0.weight" output_key = "mlp.6.weight" if "mlp.6.weight" in state_dict else "net.6.weight" obs_dim = int(state_dict[input_key].shape[1]) action_dim = int(state_dict[output_key].shape[0]) model = PolicyMLP(obs_dim=obs_dim, action_dim=action_dim) remapped_state_dict: dict[str, torch.Tensor] = {} for key, value in state_dict.items(): if key.startswith("mlp."): remapped_state_dict[key.replace("mlp.", "net.")] = value elif key.startswith("net."): remapped_state_dict[key] = value elif key == "obs_normalizer._mean": remapped_state_dict["obs_mean"] = value.squeeze() elif key == "obs_normalizer._var": remapped_state_dict["obs_std"] = torch.sqrt(value.squeeze() + 1e-5) model.load_state_dict(remapped_state_dict, strict=False) model.eval() model.to(device) model.expected_obs_dim = obs_dim model.expected_action_dim = action_dim model.backend = "torch" return model class PolicyRunner: BASE_OBS_DIM = 53 DEFAULT_STAND_POSE = np.array( [ 0.0, 0.9, -1.8, 0.0, 0.9, -1.8, 0.0, 0.9, -1.8, 0.0, 0.9, -1.8, 0.0, 0.0, 0.0, 0.0, ], dtype=np.float32, ) def __init__( self, policy_path: Path, device: torch.device | None = None, enable_zero_cmd_suppression: bool = True, hold_zero_command_pose: bool = True, command_release_s: float = 0.35, action_scale: np.ndarray | None = None, zero_cmd_use_yaw_rate: bool = True, clip_obs: float = 100.0, ): self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") self.policy_path = Path(policy_path) self.enable_zero_cmd_suppression = bool(enable_zero_cmd_suppression) self.hold_zero_command_pose = bool(hold_zero_command_pose) self.command_release_s = max(float(command_release_s), 1e-3) self.clip_obs = max(float(clip_obs), 0.0) self.policy_path = resolve_policy_path(self.policy_path) if not self.policy_path.exists(): raise FileNotFoundError(f"policy file not found: {self.policy_path}") print(f"[PolicyRunner] device={self.device}, policy={self.policy_path}") self.policy = load_policy(self.policy_path, self.device) if self.policy.expected_obs_dim != self.BASE_OBS_DIM: raise ValueError( f"Unsupported policy obs dim {self.policy.expected_obs_dim}. " f"Current sim2real only supports {self.BASE_OBS_DIM}-D actor observations." ) self.default_dof_pos = self.DEFAULT_STAND_POSE.copy() self.last_actions = np.zeros(16, dtype=np.float32) self.action_scale = np.asarray( action_scale if action_scale is not None else [ 0.125, 0.25, 0.25, 0.125, 0.25, 0.25, 0.125, 0.25, 0.25, 0.125, 0.25, 0.25, 5.0, 5.0, 5.0, 5.0, ], dtype=np.float32, ) if self.action_scale.shape != (16,): raise ValueError(f"action_scale must be shape (16,), got {self.action_scale.shape}") self.zero_cmd_lin_thresh = 0.05 self.zero_cmd_yaw_thresh = 0.05 self.zero_yaw_rate_thresh = 0.10 self.zero_cmd_use_yaw_rate = bool(zero_cmd_use_yaw_rate) self._command_release_alpha = 0.0 print( f"[PolicyRunner] obs_dim={self.policy.expected_obs_dim}, " f"base_obs_dim={self.BASE_OBS_DIM}, history=1, " f"action_dim={self.policy.expected_action_dim}, " f"backend={getattr(self.policy, 'backend', 'unknown')}, " f"clip_obs={self.clip_obs:.1f}, " f"zero_cmd_suppression={self.enable_zero_cmd_suppression}, " f"hold_zero_command_pose={self.hold_zero_command_pose}" ) def reset(self, prime_obs: np.ndarray | None = None) -> None: self.last_actions = np.zeros(16, dtype=np.float32) self._command_release_alpha = 0.0 def _is_zero_command(self, command: np.ndarray, base_ang_vel: np.ndarray) -> bool: cmd_is_zero = ( np.linalg.norm(command[:2]) < self.zero_cmd_lin_thresh and abs(command[2]) < self.zero_cmd_yaw_thresh ) if not self.zero_cmd_use_yaw_rate: return cmd_is_zero return cmd_is_zero and abs(base_ang_vel[2]) < self.zero_yaw_rate_thresh def command_activation_metrics(self, command: np.ndarray) -> tuple[float, float]: command = np.asarray(command, dtype=np.float32) planar = float(np.linalg.norm(command[:2])) yaw = float(abs(command[2])) return planar, yaw def is_command_active(self, command: np.ndarray) -> bool: planar, yaw = self.command_activation_metrics(command) return planar >= self.zero_cmd_lin_thresh or yaw >= self.zero_cmd_yaw_thresh def step(self, obs: np.ndarray, dt: float = 0.02) -> tuple[np.ndarray, np.ndarray]: obs = np.asarray(obs, dtype=np.float32) expected_obs_dim = int(self.policy.expected_obs_dim) if obs.shape[0] != expected_obs_dim: raise ValueError( f"Observation dim mismatch: got {obs.shape[0]}, expected {expected_obs_dim}." ) if self.clip_obs > 0.0: obs = np.clip(obs, -self.clip_obs, self.clip_obs).astype(np.float32, copy=False) obs_tensor = torch.tensor(obs, dtype=torch.float32, device=self.device).unsqueeze(0) with torch.no_grad(): raw_actions = self.policy(obs_tensor).squeeze(0).cpu().numpy() raw_actions = np.clip(raw_actions, -10.0, 10.0).astype(np.float32) command = obs[6:9] base_ang_vel = obs[0:3] / 0.25 zero_command = self._is_zero_command(command, base_ang_vel) if zero_command: self._command_release_alpha = 0.0 if self.hold_zero_command_pose: raw_actions[:] = 0.0 elif self.enable_zero_cmd_suppression: raw_actions[12:16] = 0.0 raw_actions[:12] *= 0.5 else: self._command_release_alpha = min(1.0, self._command_release_alpha + dt / self.command_release_s) raw_actions *= self._command_release_alpha self.last_actions = raw_actions.copy() scaled_actions = raw_actions * self.action_scale return scaled_actions, raw_actions