Files
Motrixlab/scripts/export_go1_no_linevel_onnx.py

290 lines
9.9 KiB
Python

#!/usr/bin/env python3
"""Export JAX/Flax-trained SKRL Go1 (no-linevel, 57-dim) policy to ONNX.
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_no_linevel_onnx.py
uv run scripts/export_go1_no_linevel_onnx.py --output ./my_exports
"""
import argparse
import os
import pickle
from pathlib import Path
import msgpack
import numpy as np
def _decode_flax_array(ext) -> np.ndarray | None:
if not hasattr(ext, "code"):
return None
parts = msgpack.unpackb(ext.data, raw=False)
if not isinstance(parts, list) or len(parts) < 3:
return None
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:
with open(ckpt_path, "rb") as f:
data = pickle.load(f)
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)
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 -----------------------------------------------------------
import torch
import torch.nn as nn
class PolicyTorch(nn.Module):
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):
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)
def flax_to_torch_weights(flax_params: dict, obs_dim: int, hidden_dims: list[int], action_dim: int) -> dict:
state_dict = {}
layer_names = sorted([k for k in flax_params if k.startswith("Dense_")])
hidden_dense = layer_names[:-1]
layer_idx = 0
for name in hidden_dense:
layer_params = flax_params[name]
kernel = layer_params["kernel"]
bias = layer_params["bias"]
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
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
# -- Config ------------------------------------------------------------------
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)
# 57-dim observation layout (NO linear velocity):
# [0:3] gyro (scaled *0.25)
# [3:6] gravity vector (body frame)
# [6:18] joint angle deviation (scaled *1.0)
# [18:30] joint velocity (scaled *0.05)
# [30:42] last action (raw)
# [42:45] command [vx*2.0, vy*2.0, wz*0.25]
# [45:57] foot contact forces (body frame, raw)
OBS_SCALES = {
"ang_vel": 0.25,
"dof_pos": 1.0,
"dof_vel": 0.05,
"contact_force": 1.0, # raw, no scaling
}
ACTION_SCALE = 0.05
KP, KD = 80.0, 1.0
CLIP_ACTIONS = 23.7
CLIP_OBS = 100.0
def auto_discover_checkpoint(env_name: str) -> str:
"""Find the latest best_agent checkpoint for the given env."""
base_dir = Path(f"runs/{env_name}/skrl")
if not base_dir.exists():
raise FileNotFoundError(f"No training results found: {base_dir}")
runs = sorted([d for d in base_dir.iterdir() if d.is_dir()], key=lambda d: d.stat().st_mtime, reverse=True)
for run_dir in runs:
ckpt = run_dir / "checkpoints" / "best_agent.pickle"
if ckpt.exists():
return str(ckpt)
raise FileNotFoundError(f"No best_agent.pickle found in {base_dir}")
def export(checkpoint_path: str, output_dir: str):
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)
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}]")
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()
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)} ...")
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}")
# Normalizer stats
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}")
# Metadata
meta_path = os.path.join(output_dir, "metadata.txt")
with open(meta_path, "w") as f:
f.write("# Go1 No-Linevel Terrain Walk - ONNX Policy Metadata\n")
f.write(f"env: go1-stairs-terrain-walk-no-linevel\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("\n# Observation layout (57 dims, NO linear velocity):\n")
f.write(" [0:3] gyro * ang_vel_scale\n")
f.write(" [3:6] gravity (body frame)\n")
f.write(" [6:18] joint_angle_deviation * dof_pos_scale\n")
f.write(" [18:30] joint_vel * dof_vel_scale\n")
f.write(" [30:42] last_actions (raw)\n")
f.write(" [42:45] commands [vx*2.0, vy*2.0, wz*0.25]\n")
f.write(" [45:57] foot_contact_forces (body frame, raw)\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("\n# Observation scales:\n")
for k, v in OBS_SCALES.items():
f.write(f" {k}: {v}\n")
f.write(" command_scale: [2.0, 2.0, 0.25]\n")
print(f"Metadata saved to: {meta_path}")
return onnx_path
def main():
parser = argparse.ArgumentParser(description="Export JAX-trained Go1 no-linevel policy to ONNX")
parser.add_argument("--checkpoint", type=str, default=None,
help="Path to SKRL JAX checkpoint (auto-discovered if not set)")
parser.add_argument("--output", type=str, default="exports_go1_no_linevel",
help="Output directory")
args = parser.parse_args()
if args.checkpoint:
ckpt_path = args.checkpoint
else:
ckpt_path = auto_discover_checkpoint("go1-stairs-terrain-walk-no-linevel")
if not os.path.exists(ckpt_path):
print(f"Error: checkpoint not found: {ckpt_path}")
return 1
print(f"Loading checkpoint: {ckpt_path}")
onnx_path = export(ckpt_path, args.output)
print(f"\nDone! ONNX model ready for sim2sim:")
print(f" {onnx_path}")
return 0
if __name__ == "__main__":
exit(main())