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