#!/usr/bin/env python3 """Export JAX/Flax-trained SKRL Go1 policy to ONNX for sim2sim deployment. Converts Flax weights → PyTorch → ONNX, baking in the RunningStandardScaler normalization so the ONNX model accepts raw (scaled) observations directly. Usage: uv run scripts/export_go1_onnx.py uv run scripts/export_go1_onnx.py --output ./my_exports Output files (in output_dir): policy.onnx - ONNX model with normalization baked in normalizer.npz - Normalizer stats (for reference/debugging) metadata.txt - Policy metadata (obs dim, joint order, scales, etc.) """ import argparse import os import pickle from pathlib import Path import msgpack import numpy as np # ── Flax weight decoder ────────────────────────────────────────────────── def _decode_flax_array(ext) -> np.ndarray | None: """Decode a flax-serialized msgpack ExtType to a numpy array.""" if not hasattr(ext, "code"): return None # The ExtType data is a msgpack array: [shape_list, dtype_str, raw_bytes] parts = msgpack.unpackb(ext.data, raw=False) if not isinstance(parts, list) or len(parts) < 3: return None # parts[0]: nested shape list, e.g. [[12]] or [[256, 45]] # parts[1]: dtype string, e.g. "float32" # parts[2]: raw bytes of array data def _flatten(s): if isinstance(s, list): out = [] for item in s: out.extend(_flatten(item)) return out return [s] shape = tuple(_flatten(parts[0])) dtype_str = parts[1] if isinstance(dtype_str, bytes): dtype_str = dtype_str.decode("utf-8") raw = parts[2] return np.frombuffer(raw, dtype=np.dtype(dtype_str)).reshape(shape) def load_jax_checkpoint(ckpt_path: str) -> dict: """Load a SKRL JAX checkpoint and extract all arrays. Returns dict with keys: flax_params: {layer_name: {kernel, bias} | array} – Flax-format weights running_mean: np.ndarray running_var: np.ndarray count: int """ with open(ckpt_path, "rb") as f: data = pickle.load(f) # Decode policy params policy_raw = msgpack.unpackb(data["policy"]) flax_params = {} for name, val in policy_raw["params"].items(): if isinstance(val, dict): flax_params[name] = { k: _decode_flax_array(v) for k, v in val.items() } else: flax_params[name] = _decode_flax_array(val) # Decode state preprocessor prep = msgpack.unpackb(data["state_preprocessor"]) running_mean = _decode_flax_array(prep["running_mean"]) running_var = _decode_flax_array(prep["running_variance"]) count_arr = _decode_flax_array(prep["current_count"]) count = int(count_arr.flat[0]) if count_arr is not None else 0 return { "flax_params": flax_params, "running_mean": running_mean, "running_var": running_var, "count": count, } # ── PyTorch model (for ONNX export) ───────────────────────────────────── import torch import torch.nn as nn class PolicyTorch(nn.Module): """PyTorch MLP matching the SKRL Flax policy architecture.""" def __init__(self, obs_dim: int, action_dim: int, hidden_dims: list[int]): super().__init__() self.obs_dim = obs_dim self.action_dim = action_dim self.hidden_dims = hidden_dims layers = [] in_dim = obs_dim for h in hidden_dims: layers.extend([nn.Linear(in_dim, h), nn.ELU()]) in_dim = h self.net = nn.Sequential(*layers) self.mean_layer = nn.Linear(in_dim, action_dim) def forward(self, x): return self.mean_layer(self.net(x)) class ONNXExporter(nn.Module): """Wraps policy with RunningStandardScaler normalization baked in.""" def __init__(self, policy: PolicyTorch, mean: np.ndarray, std: np.ndarray): super().__init__() self.policy = policy self.register_buffer("mean", torch.from_numpy(mean).float()) self.register_buffer("std", torch.from_numpy(std).float()) self.clip_threshold = 5.0 def forward(self, x): x = (x - self.mean) / (self.std + 1e-8) x = torch.clamp(x, min=-self.clip_threshold, max=self.clip_threshold) return self.policy(x) # ── Flax → PyTorch weight conversion ──────────────────────────────────── def flax_to_torch_weights(flax_params: dict, obs_dim: int, hidden_dims: list[int], action_dim: int) -> dict: """Convert Flax-format params to PyTorch state_dict. Flax Dense kernel: shape [in_dim, out_dim] PyTorch Linear weight: shape [out_dim, in_dim] → needs transpose Architecture: Dense_0..Dense_{N-1} → net hidden layers (Linear+ELU pairs) Dense_N → mean_layer (Linear, no activation) """ state_dict = {} layer_names = sorted([k for k in flax_params if k.startswith("Dense_")]) # Hidden layers: all Dense except the last hidden_dense = layer_names[:-1] layer_idx = 0 for name in hidden_dense: layer_params = flax_params[name] kernel = layer_params["kernel"] # Flax: [in_dim, out_dim] bias = layer_params["bias"] # [out_dim] state_dict[f"net.{layer_idx}.weight"] = torch.from_numpy(kernel.T.copy()).float() state_dict[f"net.{layer_idx}.bias"] = torch.from_numpy(bias.copy()).float() layer_idx += 2 # skip ELU activation (no params) # Output layer (mean_layer) last_name = layer_names[-1] last_params = flax_params[last_name] state_dict["mean_layer.weight"] = torch.from_numpy(last_params["kernel"].T.copy()).float() state_dict["mean_layer.bias"] = torch.from_numpy(last_params["bias"].copy()).float() return state_dict # ── Main export ───────────────────────────────────────────────────────── GO1_JOINT_NAMES = [ "FR_hip", "FR_thigh", "FR_calf", "FL_hip", "FL_thigh", "FL_calf", "RR_hip", "RR_thigh", "RR_calf", "RL_hip", "RL_thigh", "RL_calf", ] GO1_DEFAULT_ANGLES = 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, ], dtype=np.float32) # Observation layout for our 45-dim policy (NO linear velocity): # [0:3] gyro (scaled *0.25) # [3:6] gravity vector (body frame, no scale) # [6:18] joint angle deviation from default (scaled *1.0) # [18:30] joint velocity (scaled *0.05) # [30:42] last action (raw) # [42:45] command [vx, vy, wz] (scaled *[2.0, 2.0, 0.25]) OBS_SCALES = { "lin_vel": 2.0, # NOT used in 45-dim obs (kept for reference) "ang_vel": 0.25, "dof_pos": 1.0, "dof_vel": 0.05, } ACTION_SCALE = 0.05 KP, KD = 80.0, 1.0 CLIP_ACTIONS = 23.7 CLIP_OBS = 100.0 def export(checkpoint_path: str, output_dir: str): """Main export pipeline.""" ckpt = load_jax_checkpoint(checkpoint_path) flax_params = ckpt["flax_params"] running_mean = ckpt["running_mean"] running_var = ckpt["running_var"] running_std = np.sqrt(running_var) # Infer architecture from Flax params dense_keys = sorted([k for k in flax_params if k.startswith("Dense_")]) hidden_dims = [flax_params[k]["bias"].shape[0] for k in dense_keys[:-1]] obs_dim = flax_params[dense_keys[0]]["kernel"].shape[0] action_dim = flax_params[dense_keys[-1]]["bias"].shape[0] print(f"Architecture: obs={obs_dim}, hidden={hidden_dims}, action={action_dim}") print(f"Normalizer mean range: [{running_mean.min():.4f}, {running_mean.max():.4f}]") print(f"Normalizer std range: [{running_std.min():.6f}, {running_std.max():.6f}]") # Build PyTorch model and load weights policy = PolicyTorch(obs_dim, action_dim, hidden_dims) torch_weights = flax_to_torch_weights(flax_params, obs_dim, hidden_dims, action_dim) policy.load_state_dict(torch_weights, strict=True) policy.eval() # Verify conversion with a random input rng = np.random.RandomState(42) test_obs = rng.randn(1, obs_dim).astype(np.float32) with torch.no_grad(): torch_out = policy(torch.from_numpy(test_obs)).numpy() print(f"Test forward pass: input shape={test_obs.shape}, output shape={torch_out.shape}") print(f" output sample: {np.array2string(torch_out[0, :4], precision=4, suppress_small=True)} ...") # Export ONNX os.makedirs(output_dir, exist_ok=True) onnx_path = os.path.join(output_dir, "policy.onnx") exporter = ONNXExporter(policy, running_mean, running_std) exporter.eval() dummy = torch.zeros(1, obs_dim, dtype=torch.float32) torch.onnx.export( exporter, dummy, onnx_path, export_params=True, opset_version=11, input_names=["observations"], output_names=["actions"], dynamic_axes={}, ) print(f"✓ ONNX exported to: {onnx_path}") # Save normalizer stats for reference npz_path = os.path.join(output_dir, "normalizer.npz") np.savez(npz_path, mean=running_mean, std=running_std) print(f"✓ Normalizer saved to: {npz_path}") # Save metadata meta_path = os.path.join(output_dir, "metadata.txt") with open(meta_path, "w") as f: f.write(f"# Go1 Flat Terrain Walk - ONNX Policy Metadata\n") f.write(f"obs_dim: {obs_dim}\n") f.write(f"action_dim: {action_dim}\n") f.write(f"hidden_dims: {hidden_dims}\n") f.write(f"observation_layout: gyro(3) + gravity(3) + joint_angle(12) + joint_vel(12) + last_action(12) + command(3)\n") f.write(f" - NO linear velocity in observation\n") f.write(f"\n# Joint order: {GO1_JOINT_NAMES}\n") f.write(f"default_angles: {GO1_DEFAULT_ANGLES.tolist()}\n") f.write(f"action_scale: {ACTION_SCALE}\n") f.write(f"kp: {KP}\n") f.write(f"kd: {KD}\n") f.write(f"clip_actions: {CLIP_ACTIONS}\n") f.write(f"clip_observations: {CLIP_OBS}\n") f.write(f"\n# Observation scales (applied BEFORE ONNX normalization):\n") for k, v in OBS_SCALES.items(): f.write(f" {k}: {v}\n") f.write(f" command_scale: [2.0, 2.0, 0.25] # for [vx, vy, wz]\n") print(f"✓ Metadata saved to: {meta_path}") return onnx_path def main(): parser = argparse.ArgumentParser(description="Export JAX-trained Go1 policy to ONNX") parser.add_argument("--checkpoint", type=str, default="runs/go1-flat-terrain-walk/skrl/26-06-19_15-05-34-538657_PPO/checkpoints/best_agent.pickle", help="Path to SKRL JAX checkpoint (.pickle)") parser.add_argument("--output", type=str, default="exports_go1_flat", help="Output directory for ONNX model and artifacts") args = parser.parse_args() if not os.path.exists(args.checkpoint): print(f"Error: checkpoint not found: {args.checkpoint}") print("Train first: uv run scripts/train.py --env go1-flat-terrain-walk") return 1 print(f"Loading checkpoint: {args.checkpoint}") onnx_path = export(args.checkpoint, args.output) print(f"\nDone! ONNX model ready for sim2sim deployment:") print(f" {onnx_path}") return 0 if __name__ == "__main__": exit(main())