fix: clamp std before distribution, lower init_noise to 0.5, NaN guard
This commit is contained in:
314
scripts/export_go1_onnx.py
Normal file
314
scripts/export_go1_onnx.py
Normal file
@@ -0,0 +1,314 @@
|
||||
#!/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())
|
||||
Reference in New Issue
Block a user