315 lines
11 KiB
Python
315 lines
11 KiB
Python
#!/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())
|