Files
Motrixlab/scripts/export_go1_onnx.py

315 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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())