fix: clamp std before distribution, lower init_noise to 0.5, NaN guard
This commit is contained in:
217
scripts/export_dreamwaq_onnx.py
Normal file
217
scripts/export_dreamwaq_onnx.py
Normal file
@@ -0,0 +1,217 @@
|
||||
#!/usr/bin/env python3
|
||||
"""DreamWaQ ONNX export: CENet encoder + Actor → ONNX for MuJoCo deployment.
|
||||
|
||||
Model: 2 inputs, 1 output
|
||||
- observations: (1, 45)
|
||||
- obs_history: (1, 5, 45)
|
||||
→ actions: (1, 12)
|
||||
|
||||
Usage:
|
||||
uv run scripts/export_dreamwaq_onnx.py
|
||||
uv run scripts/export_dreamwaq_onnx.py --checkpoint PATH --vae PATH --output PATH
|
||||
"""
|
||||
import argparse, os, pickle, sys
|
||||
import msgpack
|
||||
import numpy as np
|
||||
import jax, jax.numpy as jnp
|
||||
import flax.linen as nn
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
# Export helpers
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
|
||||
def _decode_flax_array(ext) -> np.ndarray | None:
|
||||
"""Decode flax-serialized msgpack ExtType to numpy array."""
|
||||
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
|
||||
shape = []
|
||||
def _flatten(s):
|
||||
if isinstance(s, list):
|
||||
for x in s: _flatten(x)
|
||||
elif isinstance(s, int):
|
||||
shape.append(s)
|
||||
_flatten(parts[0])
|
||||
dtype_str = parts[1]
|
||||
raw_bytes = parts[2]
|
||||
return np.frombuffer(raw_bytes, dtype=np.dtype(dtype_str)).reshape(shape)
|
||||
|
||||
|
||||
def load_skrl_policy(path):
|
||||
"""Extract actor weights from SKRL checkpoint. Slice first layer: 254→64."""
|
||||
with open(path, 'rb') as f:
|
||||
ckpt = pickle.load(f)
|
||||
raw = msgpack.unpackb(ckpt['policy'])['params']
|
||||
params = {}
|
||||
for name, val in raw.items():
|
||||
if isinstance(val, dict):
|
||||
params[name] = {k: _decode_flax_array(v) for k, v in val.items()}
|
||||
else:
|
||||
params[name] = _decode_flax_array(val)
|
||||
# Slice first Dense layer: (254, 512) → (64, 512)
|
||||
params['Dense_0'] = {
|
||||
'kernel': params['Dense_0']['kernel'][:64, :],
|
||||
'bias': params['Dense_0']['bias'],
|
||||
}
|
||||
return params
|
||||
|
||||
|
||||
def load_state_preprocessor(path):
|
||||
"""Load RunningStandardScaler stats (running_mean, running_variance) for the
|
||||
first 64 dims = [code(19), obs(45)] that feed the actor.
|
||||
|
||||
CRITICAL: the policy was trained on NORMALIZED observations. Deployment must
|
||||
apply: clip((x - mean) / (sqrt(var) + 1e-8), -5, 5) before the actor.
|
||||
"""
|
||||
with open(path, 'rb') as f:
|
||||
ckpt = pickle.load(f)
|
||||
if 'state_preprocessor' not in ckpt:
|
||||
print("[WARN] No state_preprocessor in checkpoint — skipping normalization")
|
||||
return None, None
|
||||
sp = msgpack.unpackb(ckpt['state_preprocessor'], raw=False)
|
||||
mean = _decode_flax_array(sp['running_mean'])[:64].astype(np.float32)
|
||||
var = _decode_flax_array(sp['running_variance'])[:64].astype(np.float32)
|
||||
return mean, var
|
||||
|
||||
|
||||
def export_onnx(actor_params, vae_params, output_path, obs_mean=None, obs_var=None):
|
||||
"""Build PyTorch model from Flax params, export to ONNX.
|
||||
|
||||
If obs_mean/obs_var given, bakes in the state-preprocessor normalization
|
||||
(applied to [code(19), obs(45)] before the actor) — REQUIRED for the policy
|
||||
to behave correctly, since it was trained on normalized observations.
|
||||
"""
|
||||
import torch, torch.nn as tnn
|
||||
|
||||
class DreamWaQTorch(tnn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# CENet encoder (Manaro-Alpha: 225→128→64)
|
||||
self.enc1 = tnn.Linear(225, 128)
|
||||
self.enc2 = tnn.Linear(128, 64)
|
||||
self.latent_mu = tnn.Linear(64, 16)
|
||||
self.vel_mu = tnn.Linear(64, 3)
|
||||
# Actor (64→512→256→128→12)
|
||||
self.act1 = tnn.Linear(64, 512)
|
||||
self.act2 = tnn.Linear(512, 256)
|
||||
self.act3 = tnn.Linear(256, 128)
|
||||
self.act_out = tnn.Linear(128, 12)
|
||||
# State-preprocessor normalization buffers (for [code(19), obs(45)] = 64)
|
||||
self.register_buffer("obs_mean", torch.zeros(64))
|
||||
self.register_buffer("obs_std", torch.ones(64))
|
||||
self.normalize = False
|
||||
|
||||
def forward(self, obs, history):
|
||||
h = history.reshape(history.shape[0], -1)
|
||||
h = tnn.functional.elu(self.enc1(h))
|
||||
h = tnn.functional.elu(self.enc2(h))
|
||||
z = self.latent_mu(h)
|
||||
vel = self.vel_mu(h)
|
||||
x = torch.cat([vel, z, obs], dim=-1)
|
||||
# Apply state-preprocessor normalization (clip((x-mean)/(std+eps), -5, 5))
|
||||
if self.normalize:
|
||||
x = torch.clamp((x - self.obs_mean) / (self.obs_std + 1e-8), -5.0, 5.0)
|
||||
x = tnn.functional.elu(self.act1(x))
|
||||
x = tnn.functional.elu(self.act2(x))
|
||||
x = tnn.functional.elu(self.act3(x))
|
||||
return self.act_out(x)
|
||||
|
||||
model = DreamWaQTorch()
|
||||
if obs_mean is not None and obs_var is not None:
|
||||
model.obs_mean.data = torch.from_numpy(obs_mean.copy())
|
||||
model.obs_std.data = torch.from_numpy(np.sqrt(obs_var).copy())
|
||||
model.normalize = True
|
||||
print("[ONNX] State-preprocessor normalization baked in")
|
||||
|
||||
# Transfer CENet encoder weights (from Flax frozen dict)
|
||||
vp = vae_params['params']
|
||||
model.enc1.weight.data = torch.from_numpy(np.array(vp['enc_fc1']['kernel']).T.copy())
|
||||
model.enc1.bias.data = torch.from_numpy(np.array(vp['enc_fc1']['bias']).copy())
|
||||
model.enc2.weight.data = torch.from_numpy(np.array(vp['enc_fc2']['kernel']).T.copy())
|
||||
model.enc2.bias.data = torch.from_numpy(np.array(vp['enc_fc2']['bias']).copy())
|
||||
model.latent_mu.weight.data = torch.from_numpy(np.array(vp['latent_mu']['kernel']).T.copy())
|
||||
model.latent_mu.bias.data = torch.from_numpy(np.array(vp['latent_mu']['bias']).copy())
|
||||
model.vel_mu.weight.data = torch.from_numpy(np.array(vp['vel_mu']['kernel']).T.copy())
|
||||
model.vel_mu.bias.data = torch.from_numpy(np.array(vp['vel_mu']['bias']).copy())
|
||||
|
||||
# Transfer Actor weights (64-dim, already sliced, from msgpack decoded)
|
||||
ap = actor_params
|
||||
model.act1.weight.data = torch.from_numpy(ap['Dense_0']['kernel'].T.copy())
|
||||
model.act1.bias.data = torch.from_numpy(ap['Dense_0']['bias'].copy())
|
||||
model.act2.weight.data = torch.from_numpy(ap['Dense_1']['kernel'].T.copy())
|
||||
model.act2.bias.data = torch.from_numpy(ap['Dense_1']['bias'].copy())
|
||||
model.act3.weight.data = torch.from_numpy(ap['Dense_2']['kernel'].T.copy())
|
||||
model.act3.bias.data = torch.from_numpy(ap['Dense_2']['bias'].copy())
|
||||
model.act_out.weight.data = torch.from_numpy(ap['Dense_3']['kernel'].T.copy())
|
||||
model.act_out.bias.data = torch.from_numpy(ap['Dense_3']['bias'].copy())
|
||||
|
||||
model.eval()
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
torch.onnx.export(
|
||||
model,
|
||||
(torch.randn(1, 45), torch.randn(1, 5, 45)),
|
||||
output_path,
|
||||
input_names=['observations', 'obs_history'],
|
||||
output_names=['actions'],
|
||||
opset_version=11,
|
||||
dynamic_axes={'observations': {0: 'batch'}, 'obs_history': {0: 'batch'}, 'actions': {0: 'batch'}},
|
||||
)
|
||||
print(f"[ONNX] Exported → {output_path}")
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
# Main
|
||||
# ═══════════════════════════════════════════════════════════════════════
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--checkpoint", default=None, help="SKRL agent checkpoint")
|
||||
p.add_argument("--vae", default=None, help="CENet params .pkl")
|
||||
p.add_argument("--output", default=os.path.join(PROJECT, "exports_go1_dreamwaq", "policy.onnx"))
|
||||
args = p.parse_args()
|
||||
|
||||
run_dir = os.path.join(PROJECT, "runs", "go1-dreamwaq-walk", "skrl")
|
||||
|
||||
# Auto-find checkpoint
|
||||
if not args.checkpoint:
|
||||
runs = sorted([d for d in os.listdir(run_dir) if os.path.isdir(os.path.join(run_dir, d)) and d.startswith("26-")])
|
||||
if runs:
|
||||
ckpt_dir = os.path.join(run_dir, runs[-1], "checkpoints")
|
||||
args.checkpoint = os.path.join(ckpt_dir, "best_agent.pickle")
|
||||
|
||||
# Auto-find VAE params
|
||||
if not args.vae:
|
||||
vae_files = sorted([f for f in os.listdir(run_dir) if f.startswith("vae_") and f.endswith(".pkl")],
|
||||
key=lambda x: int(x.split("_")[1].split(".")[0]))
|
||||
if vae_files:
|
||||
args.vae = os.path.join(run_dir, vae_files[-1])
|
||||
else:
|
||||
# Try cenet_params.pkl (saved at end of training)
|
||||
cpath = os.path.join(run_dir, "cenet_params.pkl")
|
||||
if os.path.exists(cpath):
|
||||
args.vae = cpath
|
||||
|
||||
if not args.checkpoint or not os.path.exists(args.checkpoint):
|
||||
print(f"[ERROR] Checkpoint not found: {args.checkpoint}")
|
||||
sys.exit(1)
|
||||
if not args.vae or not os.path.exists(args.vae):
|
||||
print(f"[ERROR] VAE params not found: {args.vae}")
|
||||
sys.exit(1)
|
||||
|
||||
print(f"Policy: {args.checkpoint}")
|
||||
print(f"VAE: {args.vae}")
|
||||
|
||||
with open(args.vae, 'rb') as f:
|
||||
vae_params = pickle.load(f)
|
||||
actor_params = load_skrl_policy(args.checkpoint)
|
||||
obs_mean, obs_var = load_state_preprocessor(args.checkpoint)
|
||||
export_onnx(actor_params, vae_params, args.output, obs_mean, obs_var)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user