Files
Motrixlab/scripts/export_dreamwaq_onnx_new.py
2026-07-22 02:17:01 +08:00

87 lines
2.8 KiB
Python
Raw 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 DreamWaQ CENetActorModel to ONNX for MuJoCo sim2sim.
输入: obs(45) + obs_history(225) → 输出: action(12)确定性mean CENet code
用法:
uv run scripts/export_dreamwaq_onnx_new.py
uv run scripts/export_dreamwaq_onnx_new.py --checkpoint runs/.../model_2850.pt
"""
import argparse, os, sys, glob
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from motrix_rl.rslrl.torch.models.cenet_actor import CENetActorModel
from tensordict import TensorDict
PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
class DwaqInfer(torch.nn.Module):
"""确定性推理VAE mean code + Actor MLP。"""
def __init__(self, model: CENetActorModel):
super().__init__()
self.vae = model.vae
self.actor_mlp = model.mlp
def forward(self, obs, obs_history):
"""obs: (N,45), obs_history: (N,225) → action: (N,12)"""
code = self.vae.deterministic_code(obs_history) # (N,19)
latent = torch.cat([code, obs], dim=-1) # (N,64)
return torch.clamp(self.actor_mlp(latent), -4.0, 4.0)
def main():
p = argparse.ArgumentParser()
p.add_argument("--checkpoint", default=None)
p.add_argument("--output", default=None)
args = p.parse_args()
if args.checkpoint:
ckpt = args.checkpoint
else:
models = sorted(glob.glob(os.path.join(PROJECT, "runs/go1-dreamwaq-walk/rslrl/*/model_*.pt")))
if not models:
print("No checkpoints found"); sys.exit(1)
ckpt = models[-1]
print(f"[Export] checkpoint: {ckpt}")
# 加载模型
dummy = TensorDict({
"policy": torch.zeros(1, 45),
"obs_history": torch.zeros(1, 225),
"privileged_obs": torch.zeros(1, 247),
}, batch_size=[1])
model = CENetActorModel(
dummy, {"actor": ["policy", "obs_history"]}, "actor", 12,
hidden_dims=[512, 256, 128], activation="elu", stochastic=True,
init_noise_std=1.0, cenet_in_dim=225, cenet_out_dim=19,
)
data = torch.load(ckpt, map_location="cpu")
if "actor_state_dict" in data:
model.load_state_dict(data["actor_state_dict"])
else:
model.load_state_dict(data)
model.eval()
# 包装为推理模型
infer = DwaqInfer(model)
infer.eval()
# 导出 ONNX
out_path = args.output or os.path.join(os.path.dirname(ckpt), "policy.onnx")
dummy_obs = torch.zeros(1, 45)
dummy_hist = torch.zeros(1, 225)
torch.onnx.export(
infer, (dummy_obs, dummy_hist), out_path,
input_names=["obs", "obs_history"],
output_names=["actions"],
opset_version=11,
dynamic_axes={"obs": {0: "batch"}, "obs_history": {0: "batch"}, "actions": {0: "batch"}},
)
print(f"[Export] saved: {out_path}")
if __name__ == "__main__":
main()