diff --git a/cmd.md b/cmd.md
new file mode 100644
index 0000000..5f29e43
--- /dev/null
+++ b/cmd.md
@@ -0,0 +1,16 @@
+uv run tensorboard --logdir runs/go1-stairs-terrain-walk-no-linevel
+uv run scripts/train.py --env go1-stairs-terrain-walk-no-linevel \
+ --checkpoint runs/go1-stairs-terrain-walk-no-linevel/skrl/26-06-23_01-37-34-357547_PPO/checkpoints/best_agent.pickle
+
+uv run scripts/play.py --env go1-stairs-terrain-walk-no-linevel --num-envs 1
+
+# 只训练楼梯地形,线速度奖励 3 倍,继承 best_agent 权重
+uv run scripts/train.py --env go1-stairs-terrain-walk-no-linevel --force-phase 0 --tracking-linvel-scale 1.0
+
+
+uv run scripts/play.py --env go1-stairs-terrain-walk-no-linevel --num-envs 1 --force-phase 0 --num-envs 1024
+
+
+DREAMWAQ_TERRAIN=flat_stairs uv run scripts/train_dreamwaq_rsl.py \
+ --num-envs 2048 --init-noise-std 2.0 \
+ --resume runs/go1-dreamwaq-walk/rsl_dwaq/06-29_16-01-56/model_3000.pt
diff --git a/deploy_orig/go1_sim2sim.py b/deploy_orig/go1_sim2sim.py
new file mode 100644
index 0000000..e871a74
--- /dev/null
+++ b/deploy_orig/go1_sim2sim.py
@@ -0,0 +1,236 @@
+#!/usr/bin/env python3
+"""Go1 sim2sim MuJoCo viewer — Original 30k flat training.
+
+Usage: python go1_sim2sim.py
+
+Requires: mujoco, onnxruntime, pynput
+Install: pip install mujoco onnxruntime pynput
+
+Controls:
+ W/S: forward/back Q/E: strafe left/right
+ A/D: rotate Space: stop R: reset Esc: quit
+"""
+import numpy as np, mujoco, onnxruntime as ort, os, time, threading, queue
+from mujoco import viewer
+from pynput import keyboard
+
+HERE = os.path.dirname(os.path.abspath(__file__))
+ONNX = os.path.join(HERE, "policy.onnx")
+
+# ── Parameters (original MotrixLab Go1 config) ──
+NUM_OBS = 48
+KP, KD = 80.0, 0.5 # KD=0.5 + MuJoCo joint_damping(0.5) = 1.0 = training kd
+ACTION_SCALE = 0.05
+CLIP = 23.7
+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)
+
+# ── Keyboard ──
+class KB:
+ def __init__(s):
+ s._q = queue.Queue(); s.running = True; s.held = set()
+ def _n(s, k):
+ try:
+ if hasattr(k, 'char') and k.char: return k.char.lower()
+ except: pass
+ return str(k).lower()
+ def _w(s):
+ while s.running:
+ try:
+ et, k = s._q.get(timeout=0.05)
+ n = s._n(k)
+ if et == 'press': s.held.add(n)
+ elif et == 'release': s.held.discard(n)
+ except queue.Empty: pass
+ def init(s):
+ s._l = keyboard.Listener(
+ on_press=lambda k: s._q.put(('press', k)),
+ on_release=lambda k: s._q.put(('release', k)))
+ s._l.start()
+ s._t = threading.Thread(target=s._w, daemon=True); s._t.start()
+ def keys(s): return s.held.copy()
+ def stop(s): s.running = False; s._l.stop()
+
+# ── Main ──
+def main():
+ # The model XML is embedded below
+ xml = '''
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+'''
+
+ model = mujoco.MjModel.from_xml_string(xml)
+ data = mujoco.MjData(model)
+ data.qpos[0:3] = [0, 0, 0.42]
+ data.qpos[3:7] = [1, 0, 0, 0]
+ data.qpos[7:19] = DEFAULT_ANGLES
+ mujoco.mj_forward(model, data)
+
+ session = ort.InferenceSession(ONNX, providers=['CPUExecutionProvider'])
+ print(f"[Go1 sim2sim] ONNX={ONNX}")
+ print(f"[Go1 sim2sim] PD kp={KP} kd={KD+0.5} action_scale={ACTION_SCALE} obs={NUM_OBS}-dim")
+ print(f"[Go1 sim2sim] W/S前后 Q/E左右 A/D旋转 Space停 R重置 Esc退出")
+
+ kb = KB(); kb.init()
+ view = viewer.launch_passive(model, data)
+ step, vx, vy, wz = 0, 0.0, 0.0, 0.0
+ last_a = np.zeros(12, dtype=np.float32)
+
+ while view.is_running():
+ keys = kb.keys()
+ if 'escape' in keys: break
+ if 'r' in keys:
+ data.qpos[0:3] = [0, 0, 0.42]
+ data.qpos[3:7] = [1, 0, 0, 0]
+ data.qpos[7:19] = DEFAULT_ANGLES
+ data.qvel[:] = 0; last_a[:] = 0
+ mujoco.mj_forward(model, data)
+
+ vx = 1.0 if 'w' in keys else (-1.0 if 's' in keys else 0.0)
+ vy = 1.0 if 'q' in keys else (-1.0 if 'e' in keys else 0.0)
+ wz = 1.0 if 'a' in keys else (-1.0 if 'd' in keys else 0.0)
+ if ' ' in keys: vx = vy = wz = 0.0
+
+ if step % 2 == 0: # 100Hz control (MuJoCo dt=0.005)
+ obs = np.zeros(NUM_OBS, dtype=np.float32)
+ sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, "gyro")
+ adr = model.sensor_adr[sid]
+ obs[0:3] = data.sensordata[adr:adr+3] * 0.25
+ R = data.xmat[1].reshape(3, 3)
+ obs[6:9] = (R.T @ np.array([0., 0., -1.])).astype(np.float32)
+ obs[9:21] = (data.qpos[7:19] - DEFAULT_ANGLES) * 1.0
+ obs[21:33] = data.qvel[6:18] * 0.05
+ obs[33:45] = last_a
+ obs[45:48] = np.array([vx, vy, wz]) * np.array([2., 2., 0.25])
+ obs = np.clip(obs, -100., 100.)
+ action = session.run(None, {'observations': obs.reshape(1, -1).astype(np.float32)})[0][0]
+ action = np.clip(action, -CLIP, CLIP)
+ last_a = action.copy()
+
+ target = DEFAULT_ANGLES + action * ACTION_SCALE
+ torques = KP * (target - data.qpos[7:19]) - KD * data.qvel[6:18]
+ data.ctrl[:] = np.clip(torques, -CLIP, CLIP)
+ mujoco.mj_step(model, data)
+ view.sync()
+ step += 1
+ time.sleep(0.001)
+
+ kb.stop(); view.close()
+
+if __name__ == "__main__":
+ main()
diff --git a/deploy_orig/policy.onnx b/deploy_orig/policy.onnx
new file mode 100644
index 0000000..2287c10
Binary files /dev/null and b/deploy_orig/policy.onnx differ
diff --git a/docs/dreamwaq_alignment.md b/docs/dreamwaq_alignment.md
new file mode 100644
index 0000000..3432c89
--- /dev/null
+++ b/docs/dreamwaq_alignment.md
@@ -0,0 +1,189 @@
+# DreamWaQ 复现与对齐记录
+
+> 目标:让本地 `go1-dreamwaq-walk` 完全对齐并复现 [Manaro-Alpha/DreamWaQ](https://github.com/Manaro-Alpha/DreamWaQ.git)(基于 IsaacGym + rsl_rl 的实现),迁移到 MotrixLab(MotrixSim + SKRL-JAX)。
+
+最后更新:2026-06-25
+
+---
+
+## 1. 架构总览
+
+```
+CENet (VAE): history(5×45=225) → [128,64] → latent(16) + vel_est(3) = code(19)
+Decoder: code(19) → [64,128] → obs_recon(45)
+Actor: code(19) + obs(45) = 64 → [512,256,128] → action(12)
+Critic: privileged_obs(235) → [512,256,128] → value(1)
+
+观测(45): ang_vel(3) + gravity(3) + commands(3) + dof_pos(12) + dof_vel(12) + actions(12)
+特权观测(235): obs(45) + base_vel(3) + heights(187, 17×11 网格)
+增广观测(254, 训练用): [code(19) | obs(45) | base_vel(3) | heights(187)]
+```
+
+---
+
+## 2. 🔴 真正让机器人站不起来的 Bug(最关键)
+
+排查发现并非训练不充分,而是两个环境 bug 让**约一半的训练数据是垃圾**:
+
+### Bug A — 机器人生成在地形地下
+- **现象**:reset 时机器人固定生成在绝对高度 z=0.40,但 DreamWaQ 地形高度为 0–1.6m。在高地形格子上,机器人生成在**地表以下最多 1.2m**。
+- **影响**:31/64 的环境在生成瞬间就卡进地形 → 爆炸/立即终止/疯狂抖动 → dof_acc 高达 150 万 → 奖励被 `only_positive_rewards` 裁剪为 0。一半训练样本无效。
+- **修复**:新增 `_sample_terrain_height(xy)`,按生成点的局部地形高度设置 z = `terrain_z + 0.40`。修复后 0/64 在地下。
+- **位置**:[dreamwaq.py](../motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py) `reset()` + `_sample_terrain_height()`
+
+### Bug B — feet_air_time 奖励恒为 0
+- **现象**:腾空时间在 `update_feet_air_time` 中被 `*= ~contacts` 清零,**早于** reward 函数读取 `first_contact = (air_time > 0) * contacts`。落地瞬间 air_time 已是 0 → first_contact 永远为 0 → 步态奖励完全失效。
+- **影响**:机器人没有迈步激励,倾向于站着不动/拖行,学不会正常步态。
+- **修复**:按上游顺序重构 —— 先用 `contact_filt = contact OR last_contact` 算出 `first_contact` 和落地时的 air_time,**再**清零;reward 函数改为读取预计算值(无状态,对 DreamWaQ 的 `_get_reward` 二次调用安全)。
+- **位置**:[walk_np.py](../motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py) `update_feet_air_time()` + `_reward_feet_air_time()`
+
+### Bug C — 初始地形难度过高
+- **现象**:首次 reset 在 0-9 级随机生成(含最难地形)。上游用 `max_init_terrain_level=5`。
+- **修复**:首次生成限制在 0-5 级,由课程学习逐步提升。
+- **位置**:[dreamwaq.py](../motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py) `reset()`
+
+---
+
+## 3. 对齐改动(与上游一致)
+
+### 控制 / 物理(cfg.py + dreamwaq.py)
+| 参数 | 改动前 | 改动后(上游) |
+|---|---|---|
+| stiffness (Kp) | 80 | **28.0** |
+| damping (Kd) | 1.0 | **0.7** |
+| action_scale | 0.05 | **0.25** |
+| sim_dt / ctrl_dt | 0.01 / 0.01 | **0.005 / 0.02**(decimation=4,200Hz 物理 / 50Hz 控制) |
+
+### 默认关节角度(cfg.py,非对称)
+```
+hips: ±0.1 (FL/RL: +0.1, FR/RR: -0.1)
+thighs: 前 0.8 / 后 1.0
+calves: -1.5 (全部)
+```
+
+### 观测顺序(walk_np.py)
+改为上游顺序 `[ang_vel, gravity, commands, dof_pos, dof_vel, actions]`,并加入观测噪声(`add_noise=True, noise_level=1.0`)。
+
+### 命令范围(dreamwaq.py)
+`lin_vel_x/y/yaw` 全部 `[-1, 1]`(之前 vel_x 上限是 2.0)。
+
+### 奖励函数对齐(dreamwaq.py)
+- `smoothness`:基于 `joint_pos_target`(= action×scale + default)的二阶差分,而非原始 action
+- `base_height`:`square(base_z - mean(measured_heights) - 0.30)`,扣除局部地形高度
+- `power_distribution`:`var(torque×vel)` 跨全部 12 关节,而非按腿分组
+- 12 项奖励尺度与上游 Go1 base config 精确一致
+
+### 域随机化(dreamwaq.py)
+motor_strength / Kp_factor / Kd_factor ∈ [0.9, 1.1],每 4 秒重采样(`rand_interval_s=4`),在 `_compute_torques` 中应用。
+
+### 训练(dreamwaq_ppo.py + go1_dreamwaq.py)
+- **关键 bug 修复**:SKRL 字段名 `entropy_coef→entropy_loss_scale`、`clip_param→ratio_clip`、`gamma→discount_factor`。改前 entropy 被静默设为 0.0(字段名不存在)。
+- CENet 解码目标:重构**当前观测**(历史最后一帧),而非预测下一帧(匹配上游 `decode_target = obs_batch`)
+- 正交权重初始化 `orthogonal(sqrt(2))`,应用于 CENet/Actor/Critic 所有 Dense 层
+- PPO 超参:rollouts=24, lr=1e-3, entropy=0.01, gamma=0.99, lam=0.95, clip=0.2, epochs=5, mini_batches=4
+
+### 地形(gen_dreamwaq_terrain.py)
+降低高度尺度(最大 2.73m → 1.66m),更适合 0.33m 高的机器人。生成后更新 XML 的 hfield `size`。
+
+---
+
+## 3.5 🔴🔴 部署缺少观测归一化("站不起来"的真正根因)
+
+这是导致导出/play/sim2sim 中机器人**完全站不起来**的根本原因,独立于训练质量。
+
+- **现象**:训练奖励正常(mean 35 / max 924),但导出 ONNX / `play_dreamwaq.py` / MuJoCo sim2sim 中机器人塌成一团,actor 输出动作幅度高达 12(关节目标偏移 ~3 rad,直接打到限位)。
+- **根因**:SKRL PPO 通过 `_add_runtime_config` 启用了 **`RunningStandardScaler` 状态预处理器**——训练时观测被 `clip((x - running_mean) / (sqrt(var) + 1e-8), -5, 5)` 归一化(running_mean 幅度达 ±3.2,std 0.14–3.2)。但部署管线直接喂**原始观测**给 actor,输入尺度完全错误 → 策略输出垃圾动作 → 倒地。
+- **验证**:加上归一化后,站立从 **0/16 → 16/16 upright**。
+- **修复**:
+ - [export_dreamwaq_onnx.py](../scripts/export_dreamwaq_onnx.py):从 checkpoint 提取 `state_preprocessor` 的 mean/var(前 64 维 = code+obs),烘焙进 ONNX 模型(actor 前做归一化+clip)
+ - [play_dreamwaq.py](../scripts/play_dreamwaq.py):`policy_forward` 中应用同样的归一化
+ - sim2sim 用 ONNX,重新导出后自动修复
+- **注意**:`RunningStandardScaler` 是 SKRL 的附加项,上游 rsl_rl 没有(上游只用固定 obs_scales)。任何脱离 SKRL agent 的部署都**必须**复现这个归一化,否则策略失效。
+
+---
+
+## 3.6 🔴 训练不收敛:初始探索过猛(IsaacGym→MotrixSim reality gap)
+
+修复部署归一化后发现策略**根本没学会站立**——在 MotrixSim 训练环境里 0/64 站立/行走,比"零动作"还差。
+
+- **诊断**:零动作(纯 PD 保持默认姿态)奖励 **1.27/step**,27/64 站立;训练后的策略奖励 **0.035/step**,0/64 站立。**训练把策略训得比啥都不做还差**。
+- **根因**:`initial_log_std=0.0`(std=1)× `action_scale=0.25` = 每步 ±0.25 rad 随机抖动 @50Hz。实测:
+ | 探索 std | 抖动 | 站立率 | 平均奖励 |
+ |---|---|---|---|
+ | 1.0(上游) | 0.25rad | 27/64 | **0.025** |
+ | 0.37 | 0.09rad | 29/64 | 0.512 |
+ | 0.2 | 0.05rad | 23/64 | 0.730 |
+
+ 机器人在所有 std 下都能站(站立率相近),但 std=1 时剧烈抖动产生的 dof_acc/action_rate 惩罚把站立的正奖励**完全淹没**(净奖励 0.025 ≈ 噪声)→ PPO 没有指向站立的梯度 → 收敛到塌倒局部最优。
+- **为何上游能用**:DreamWaQ 的 std=1 是为 IsaacGym 调的;MotrixSim 接触/动力学不同(你能用的 go1 平地任务用的是 5× 更温和的 `action_scale=0.05`)。
+- **修复**:[go1_dreamwaq.py](../motrix_rl/src/motrix_rl/tasks/go1_dreamwaq.py) 设 `initial_log_std=-1.0`(std=0.37,抖动 0.09rad),让机器人早期保持直立、获得清晰的站立奖励信号后再 bootstrap 行走。这是 reality-gap 的必要偏离(上游用 0.0)。
+
+---
+
+## 4. 🛠 原生 play 渲染崩溃修复
+
+- **现象**:`play_dreamwaq.py` 启动 MotrixSim 渲染器时崩溃 `Couldn't get swap chain texture`,而其他 go1 任务的 play 正常。
+- **根因**:JAX 默认预分配 **75% 显存**(12GB 中 ~9GB),饿死了 Vulkan 渲染器。标准 play.py 因渲染器创建与 JAX 初始化间隔充足而不受影响。
+- **修复**:在 `import jax` 前设置 `os.environ["XLA_PYTHON_CLIENT_PREALLOCATE"] = "false"`,并优雅处理 `RenderClosedError`。
+- **注意**:训练占用 GPU 时无法同时跑原生渲染;请在训练结束后查看。
+
+---
+
+## 5. 修改的文件
+
+| 文件 | 改动 |
+|---|---|
+| [motrix_envs/.../go1/cfg.py](../motrix_envs/src/motrix_envs/locomotion/go1/cfg.py) | PD 增益、关节角度、噪声尺度 |
+| [motrix_envs/.../go1/dreamwaq.py](../motrix_envs/src/motrix_envs/locomotion/go1/dreamwaq.py) | **地形生成修复**、ctrl 频率、命令范围、奖励、域随机化、观测噪声 |
+| [motrix_envs/.../go1/walk_np.py](../motrix_envs/src/motrix_envs/locomotion/go1/walk_np.py) | 观测顺序、**feet_air_time 修复** |
+| [motrix_rl/.../skrl/jax/train/dreamwaq_ppo.py](../motrix_rl/src/motrix_rl/skrl/jax/train/dreamwaq_ppo.py) | CENet 解码目标、VAE 缓冲、正交初始化 |
+| [motrix_rl/.../tasks/go1_dreamwaq.py](../motrix_rl/src/motrix_rl/tasks/go1_dreamwaq.py) | **SKRL 字段名修复**、超参 |
+| [scripts/gen_dreamwaq_terrain.py](../scripts/gen_dreamwaq_terrain.py) | 地形高度尺度 |
+| [scripts/play_dreamwaq.py](../scripts/play_dreamwaq.py) | **GPU 预分配修复** |
+| [scripts/dreamwaq_sim2sim_mujoco.py](../scripts/dreamwaq_sim2sim_mujoco.py) | PD/角度/decimation 与训练对齐 |
+| [scripts/export_dreamwaq_onnx.py](../scripts/export_dreamwaq_onnx.py) | CENet 编码器 + Actor → ONNX |
+
+---
+
+## 6. 使用方法
+
+### 训练
+```bash
+uv run scripts/train_dreamwaq.py --num-envs 2048 --timesteps 100M --seed 42
+```
+> 用 2048 envs 避免 4096 在 ~70% 处的 OOM(exit 144)。检查点存于 `runs/go1-dreamwaq-walk/skrl/`。
+
+### 原生查看(MotrixSim,GPU 空闲时)
+```bash
+uv run scripts/play_dreamwaq.py --num-envs 9
+```
+
+### 导出 ONNX
+```bash
+uv run scripts/export_dreamwaq_onnx.py \
+ --checkpoint runs/go1-dreamwaq-walk/skrl//checkpoints/best_agent.pickle \
+ --vae runs/go1-dreamwaq-walk/skrl/vae_.pkl \
+ --output exports_go1_dreamwaq/policy.onnx
+```
+
+### MuJoCo sim2sim
+```bash
+uv run scripts/dreamwaq_sim2sim_mujoco.py --terrain dreamwaq
+# 控件: W/S 前后 Q/E 左右 A/D 旋转 Space 停 R 重置 Esc 退出
+```
+
+### TensorBoard
+```bash
+uv run tensorboard --logdir runs/go1-dreamwaq-walk
+```
+
+---
+
+## 7. 待办 / 注意事项
+
+- **检查点兼容性**:改动前的检查点(含旧 PD/观测顺序/地形 bug)不兼容,应丢弃重训。
+- **未对齐项(有意保留)**:
+ - 奖励未乘 `dt`(上游 `_prepare_reward_function` 会 `scale *= dt`)。这是全局均匀缩放,PPO 优势归一化会抵消,不影响行为;为避免干扰已调好的 value 学习而保留。
+ - VAE 训练为 PPO 更新后独立一步,而非上游的联合优化(功能等价,更易维护)。
+ - 地形几何为自定义生成器,非 IsaacGym `terrain_utils`(概念相似,具体形状不同)。
+- **OOM**:4096 envs 曾在 70% 处被杀(exit 144),改用 2048。
diff --git a/docs/dreamwaq_rsl_handoff.md b/docs/dreamwaq_rsl_handoff.md
new file mode 100644
index 0000000..e8a3ec8
--- /dev/null
+++ b/docs/dreamwaq_rsl_handoff.md
@@ -0,0 +1,168 @@
+# DreamWaQ on MotrixSim — rsl_rl 复现交接文档
+
+> 目标:在 MotrixLab(MotrixSim + numpy 物理)上**忠实复现** [Manaro-Alpha/DreamWaQ](https://github.com/Manaro-Alpha/DreamWaQ)(原版基于 IsaacGym + 定制 rsl_rl-1.0.2),让 Go1 在地形上行走。
+>
+> 最后更新:2026-06-26。当前状态:rsl_rl 路线训练中,趋势健康(站立解决、行走学习中)。
+
+---
+
+## 0. 一句话总览
+
+原版 DreamWaQ = **ActorCritic_DWAQ(CENet+Actor+Critic 一体网络) + PPO 联合训练 VAE + 金字塔地形课程**。我把上游定制的 `rsl_rl-1.0.2` 核心**原样移植**到 MotrixLab,写了个 numpy↔torch 环境适配器,并针对 MotrixSim 的两个 reality gap(hfield 大 z_scale 接触不稳、IsaacGym 超参偏激进)做了适配。
+
+---
+
+## 1. 两条路线(背景)
+
+复现走过两条路线,**rsl_rl 是当前主线**:
+
+| | SKRL-JAX 路线(旧,参考) | **rsl_rl 路线(当前主线)** |
+|---|---|---|
+| 代码 | `motrix_rl/skrl/jax/train/dreamwaq_ppo.py` | `motrix_rl/dwaq_rsl/`(新模块) |
+| 网络 | CENet 拼在 wrapper,Actor/Critic 在 SKRL | 上游 `ActorCritic_DWAQ` 一体网络 |
+| VAE 训练 | **单独一步**(CENet 不参与 PPO 梯度) | **PPO 联合优化**(和上游一致) |
+| 结果 | 站立解决,但**学不会走**(机器人趴/站着不动) | 趋势健康,行走学习中 |
+
+**为什么切到 rsl_rl**:SKRL 路线机器人学不会走,根因是 **VAE 单独训练** —— CENet 的 code(latent+vel)早期是垃圾,actor 用不上。上游把 VAE 和 PPO 联合优化(一个 optimizer step),code 早期就有意义。切到 rsl_rl 后同样的 `std=1.0` 下 reward 从卡死的 0.025 变成快速上涨,这是关键转折。
+
+> SKRL 路线的详细 bug 修复记录见 [dreamwaq_alignment.md](dreamwaq_alignment.md)(地下生成、feet_air_time、部署归一化等,这些 env 层 bug 两条路线都受益)。
+
+---
+
+## 2. 新建模块:`motrix_rl/src/motrix_rl/dwaq_rsl/`
+
+把上游 `rsl_rl-1.0.2` 的核心**原样移植**(纯 torch,去掉 isaacgym 依赖,改相对 import):
+
+| 文件 | 来源 | 作用 |
+|---|---|---|
+| `actor_critic_dwaq.py` | 上游 `modules/actor_critic_DWAQ.py` **原样** | CENet(encoder 225→128→64 + latent/vel 头)+ decoder + Actor[512,256,128] + Critic[512,256,128]。`act(obs, obs_history)` 把 CENet code(19)和 obs(45)拼成 64 维喂 actor |
+| `ppo.py` | 上游 `algorithms/ppo.py` **原样** | PPO + **VAE 联合 loss**:`loss = surrogate + value + autoenc`,其中 `autoenc = MSE(vel, prev_priv[:,45:48]) + MSE(decode, obs) + beta*KL` |
+| `rollout_storage.py` | 上游 `storage/rollout_storage.py` **原样** | 经验存储,含 `prev_critic_obs`(VAE 的 vel_target 来源)、`observation_history` |
+| `utils.py` | 上游 `utils/utils.py` | `split_and_pad_trajectories` |
+| `runner.py` | 上游 `on_policy_runner.py` **适配** | 去掉 rsl_rl 包 import + isaacgym 噪声,改用本地模块。learn() 主循环、save/load、tensorboard 日志 |
+| `env_wrapper.py` | **新写** | `DwaqVecEnv`:MotrixLab `DreamWaQTask`(numpy)→ 上游 VecEnv 接口(torch)。见下 |
+
+### `DwaqVecEnv`(env_wrapper.py)关键点
+
+上游 runner 要求 `env.step()` 返回 **7 元组**:`(obs, privileged_obs, prev_privileged_obs, obs_hist, rewards, dones, extras)`。MotrixLab 的 `DreamWaQTask` 已提供 obs(45)、`info["privileged_obs"]`(235)、`info["obs_history"]`(N,5,45)。wrapper 做:
+- numpy↔torch 转换(MotrixSim 物理在 CPU,网络在 GPU,每步传输)
+- `obs_history` flatten:(N,5,45) → (N,225)
+- 维护 `prev_privileged_obs`(step 前的 privileged,VAE 用其 [45:48]=base_vel 作 vel_target)
+- `extras["time_outs"]` = truncated
+- **rewards/dones 是 (N,) 不是 (N,1)**(上游约定,storage 内部 view 成 (N,1)) — 这是个踩过的坑
+
+### 维度速查
+```
+obs=45, privileged=235(=obs45+base_vel3+heights187), obs_hist=225(=5*45)
+actor_in=64(=code19+obs45), critic_in=235, cenet_in=225, cenet_out=19(=vel3+latent16)
+```
+
+---
+
+## 3. 环境层改动:`motrix_envs/.../go1/dreamwaq.py`
+
+(这些改动 rsl_rl 路线和 SKRL 路线共用同一个 `DreamWaQTask`)
+
+1. **去掉命令课程**(关键):原来我加过一个命令课程(cmd_scale 0.3→1.0),但它让机器人**趴着刷小命令分**(小命令下趴着也能 tracking)。改回**上游 full range [-1,1]**(`commands.curriculum=False`),机器人被要求走,趴着不划算。见 `resample_commands`。
+2. **level 0 起步**:`reset()` 里 init terrain level 全设 0(平地),靠地形难度课程逐级升。原来 0-5 随机会让机器人 spawn 在难地形秒摔。
+3. **中心 spawn**:`_make_origins` 坐标中心化(hfield 中心在世界原点 (0,0)),spawn 在 cell 的 3m 平台中心 ±1m。
+4. **地形高度采样**:`_sample_terrain_height` / `_get_heights` 从 hfield 采样(`hf.height_matrix` + `hf.bound`),用于 spawn 高度 + privileged obs 的 heights。
+5. 奖励对齐上游(smoothness 用 joint_pos_target + step mask、power_distribution 用 `var(abs(power))`、base_height 扣地形高度等)。
+
+---
+
+## 4. 地形:`scripts/gen_dreamwaq_terrain.py` + ⚠️ MotrixSim hfield 坑
+
+### 生成器
+移植了上游 IsaacGym `terrain_utils` 的金字塔地形(纯 numpy):`pyramid_sloped_terrain` / `pyramid_stairs_terrain` / `discrete_obstacles_terrain` / `random_uniform_terrain`,10 行难度 × 20 列类型网格,每格 3m 中心平台,输出 16-bit PNG hfield + XML `` 行。
+
+### ⚠️⚠️ 最大的坑:MotrixSim hfield 在大 z_scale 下接触不稳
+
+**实测结论**(零动作站立 20 步摔倒数):
+
+| hfield z_scale | 机器人 above-terrain | 20步摔倒 |
+|---|---|---|
+| 3.78(上游全高金字塔) | 0.12(**趴**) | 1267 |
+| 1.66 | 0.14(趴) | 1049 |
+| **0.54(矮金字塔)** | **0.35(正常站立)** | **120** |
+| plane geom(对比) | 0.33 | **0** |
+
+- MotrixSim 用 **hfield(高度场)**,陡台阶垂直面只能近似成陡坡,机器人脚(sphere geom)打滑;**z_scale 越大越不稳**。
+- 上游用 **trimesh(三角网格)**,垂直面精确。**调研确认 MotrixSim 不支持非凸 trimesh 碰撞**(mesh 是凸包),接触参数(condim/solref/solimp)调优也无效。
+- **解决**:把金字塔楼梯/坡的高度**缩小 ~4 倍**(`make_terrain` 里 `step_height=0.012+0.020*diff`、`slope=0.10*diff`),让 z_scale≈0.54。保留金字塔结构 + 难度课程,但高度适配 MotrixSim hfield。**这是必要的 reality-gap 妥协**(不完全忠实上游高度,但能训)。
+- `scene_dreamwaq_terrain.xml` 的 hfield `size="85 45 0.54 0.01"`:**第1/2 是半径**(全宽 170×90m),第3 是 z_scale,第4 是底部厚度。
+
+**重新生成地形后必须同步更新 XML 的 z_scale**(gen 脚本会打印 `size="..."` 那行的值)。
+
+---
+
+## 5. 脚本(训练 / 导出 / 可视化)
+
+| 脚本 | 作用 |
+|---|---|
+| `scripts/train_dreamwaq_rsl.py` | rsl_rl 训练。超参 dict 内嵌(对齐上游 Go1RoughCfgPPO)。`--num-envs 2048 --iterations 3000 --init-noise-std 1.0` |
+| `scripts/export_dreamwaq_rsl_onnx.py` | rsl 的 `.pt` checkpoint → ONNX。**rsl 无 state 归一化**(不像 SKRL),所以 ONNX 简单:`obs(1,45)+obs_history(1,5,45)→action(1,12)`,用 mean CENet code |
+| `scripts/gen_dreamwaq_terrain.py` | 生成矮金字塔地形 PNG |
+| `scripts/dreamwaq_sim2sim_mujoco.py` | MuJoCo sim2sim 可视化(用 ONNX)。PD/角度/decimation 已对齐(KP=28,KD=0.7,action_scale=0.25,decimation=4) |
+
+### 命令
+```bash
+# 训练 (后台跑, print 会被 buffer, 用 tensorboard 看)
+PYTHONUNBUFFERED=1 uv run scripts/train_dreamwaq_rsl.py --num-envs 2048 --iterations 3000 --seed 1 --init-noise-std 1.0
+
+# 监控 (print 被 buffer, 看 tensorboard)
+uv run tensorboard --logdir runs/go1-dreamwaq-walk/rsl_dwaq
+
+# 导出最新 checkpoint
+JAX_PLATFORMS=cpu uv run scripts/export_dreamwaq_rsl_onnx.py # 自动找最新 model_*.pt
+
+# MuJoCo 可视化 (本地有显示的机器)
+uv run scripts/dreamwaq_sim2sim_mujoco.py --terrain flat # 平地最干净
+```
+
+---
+
+## 6. 关键决策 + 为什么(reality gap 适配)
+
+复现不是 1:1 照搬,MotrixSim vs IsaacGym 有真实差异,做了几处**有意的适配**(都在代码注释里标了原因):
+
+1. **VAE 联合训练**(用 rsl_rl 而非 SKRL):核心。SKRL 路线 VAE 单独训练→学不会走。
+2. **去命令课程**:命令课程让机器人趴着刷小命令分;上游本来就 full range 无课程。
+3. **矮金字塔地形**(z_scale 0.54):MotrixSim hfield 大 z_scale 接触不稳的妥协。
+4. **level 0 平地起步**:机器人先在平地学走,靠地形课程升级。
+5. 其余超参/网络/奖励**严格对齐上游**(PD 28/0.7、action_scale 0.25、decimation 4、init_noise_std 1.0、rollouts 24、lr 1e-3、entropy 0.01 等)。
+
+---
+
+## 7. 当前状态(2026-06-26)
+
+训练中(`runs/go1-dreamwaq-walk/rsl_dwaq/06-26_00-10-27`),趋势健康:
+
+| 指标 | it59 → it182 |
+|---|---|
+| mean_reward | 0.02 → 2.6 → **18.6** |
+| episode 长度 | 12 → 101 → **143** |
+| noise_std | 1.0 → 0.84 → **0.57**(在学) |
+| base_height | 稳定(**不趴**) |
+
+机器人能站(above ~0.35)+ episode 越来越长 + std 持续降 = 在学走。**需等训练跑完(3000 iter,~2.5h)再导出验证实际行走**(站立率/前进速度跟踪/摔倒率,别只看奖励数字 — 之前吃过"奖励高但实际趴着"的亏)。
+
+---
+
+## 8. 已知问题 / 坑(接手必读)
+
+1. **训练 print 被 buffer**:重定向到文件时 Python print 不实时 flush。**用 tensorboard 监控**,别等 print。(已加 `PYTHONUNBUFFERED=1` 缓解)
+2. **`ps` 的 TIME 是 CPU 累计时间**(跨核),不是 wall time。看 wall 用 `etime` 或 tfevents 时间戳。
+3. **VAE checkpoint 文件混在共享目录**(SKRL 路线遗留):`runs/.../skrl/vae_*.pkl`。rsl 路线 checkpoint 是 `rsl_dwaq//model_*.pt`,别搞混。
+4. **物理在 CPU**(MotrixSim numpy backend),网络在 GPU。每步 numpy↔torch 传输有开销,~16 it/min(2048 envs)。物理是瓶颈。
+5. **MotrixSim hfield 不支持陡台阶**(根本限制):如果要更忠实的高地形,得调研 MotrixSim 是否有 mesh 凸分解地形方案,或接受矮地形。
+6. **验证要看行为不看奖励**:`only_positive_rewards` + 一堆"别动"惩罚,机器人容易学"趴着/站着不动"刷分。验证时务必测**实际前进速度跟踪 + 摔倒率**,不能只看 mean_reward。
+
+---
+
+## 9. 下一步
+
+1. 等当前训练跑完,导出 + MuJoCo 验证实际行走(前进/转向跟踪)。
+2. 若行走 OK:逐步提高地形难度(z_scale),找 MotrixSim hfield 能稳定接触的上限。
+3. 若行走不够好:检查地形课程是否真在升级(机器人有没有从 level 0 走到 level 1+)、调 init_noise_std。
+4. 长期:调研 MotrixSim mesh 地形(凸分解)以更忠实复现陡台阶。
diff --git a/exports_go1_cts/normalizer.npz b/exports_go1_cts/normalizer.npz
new file mode 100644
index 0000000..a11551d
Binary files /dev/null and b/exports_go1_cts/normalizer.npz differ
diff --git a/exports_go1_cts/policy.onnx b/exports_go1_cts/policy.onnx
new file mode 100644
index 0000000..8c01694
Binary files /dev/null and b/exports_go1_cts/policy.onnx differ
diff --git a/exports_go1_cts/policy_full.onnx b/exports_go1_cts/policy_full.onnx
new file mode 100644
index 0000000..7a3195d
Binary files /dev/null and b/exports_go1_cts/policy_full.onnx differ
diff --git a/exports_go1_dreamwaq/policy.onnx b/exports_go1_dreamwaq/policy.onnx
new file mode 100644
index 0000000..ccecef0
Binary files /dev/null and b/exports_go1_dreamwaq/policy.onnx differ
diff --git a/exports_go1_flat/metadata.txt b/exports_go1_flat/metadata.txt
new file mode 100644
index 0000000..8898f86
--- /dev/null
+++ b/exports_go1_flat/metadata.txt
@@ -0,0 +1,21 @@
+# Go1 Flat Terrain Walk - ONNX Policy Metadata
+obs_dim: 45
+action_dim: 12
+hidden_dims: [256, 128, 64]
+observation_layout: gyro(3) + gravity(3) + joint_angle(12) + joint_vel(12) + last_action(12) + command(3)
+ - NO linear velocity in observation
+
+# Joint order: ['FR_hip', 'FR_thigh', 'FR_calf', 'FL_hip', 'FL_thigh', 'FL_calf', 'RR_hip', 'RR_thigh', 'RR_calf', 'RL_hip', 'RL_thigh', 'RL_calf']
+default_angles: [-0.0, 0.8999999761581421, -1.7999999523162842, 0.0, 0.8999999761581421, -1.7999999523162842, -0.0, 0.8999999761581421, -1.7999999523162842, 0.0, 0.8999999761581421, -1.7999999523162842]
+action_scale: 0.05
+kp: 80.0
+kd: 1.0
+clip_actions: 23.7
+clip_observations: 100.0
+
+# Observation scales (applied BEFORE ONNX normalization):
+ lin_vel: 2.0
+ ang_vel: 0.25
+ dof_pos: 1.0
+ dof_vel: 0.05
+ command_scale: [2.0, 2.0, 0.25] # for [vx, vy, wz]
diff --git a/exports_go1_flat/normalizer.npz b/exports_go1_flat/normalizer.npz
new file mode 100644
index 0000000..9110aeb
Binary files /dev/null and b/exports_go1_flat/normalizer.npz differ
diff --git a/exports_go1_flat/policy.onnx b/exports_go1_flat/policy.onnx
new file mode 100644
index 0000000..88619ab
Binary files /dev/null and b/exports_go1_flat/policy.onnx differ
diff --git a/exports_go1_go2style/normalizer.npz b/exports_go1_go2style/normalizer.npz
new file mode 100644
index 0000000..dcbb432
Binary files /dev/null and b/exports_go1_go2style/normalizer.npz differ
diff --git a/exports_go1_go2style/policy.onnx b/exports_go1_go2style/policy.onnx
new file mode 100644
index 0000000..1b98d90
Binary files /dev/null and b/exports_go1_go2style/policy.onnx differ
diff --git a/exports_go1_no_linevel/better1_policy.onnx b/exports_go1_no_linevel/better1_policy.onnx
new file mode 100644
index 0000000..f86599c
Binary files /dev/null and b/exports_go1_no_linevel/better1_policy.onnx differ
diff --git a/exports_go1_no_linevel/better2.onnx b/exports_go1_no_linevel/better2.onnx
new file mode 100644
index 0000000..301ad7a
Binary files /dev/null and b/exports_go1_no_linevel/better2.onnx differ
diff --git a/exports_go1_no_linevel/better3.onnx b/exports_go1_no_linevel/better3.onnx
new file mode 100644
index 0000000..fdd3380
Binary files /dev/null and b/exports_go1_no_linevel/better3.onnx differ
diff --git a/exports_go1_no_linevel/metadata.txt b/exports_go1_no_linevel/metadata.txt
new file mode 100644
index 0000000..c090eb2
--- /dev/null
+++ b/exports_go1_no_linevel/metadata.txt
@@ -0,0 +1,29 @@
+# Go1 No-Linevel Terrain Walk - ONNX Policy Metadata
+env: go1-stairs-terrain-walk-no-linevel
+obs_dim: 57
+action_dim: 12
+hidden_dims: [512, 256, 128]
+
+# Observation layout (57 dims, NO linear velocity):
+ [0:3] gyro * ang_vel_scale
+ [3:6] gravity (body frame)
+ [6:18] joint_angle_deviation * dof_pos_scale
+ [18:30] joint_vel * dof_vel_scale
+ [30:42] last_actions (raw)
+ [42:45] commands [vx*2.0, vy*2.0, wz*0.25]
+ [45:57] foot_contact_forces (body frame, raw)
+
+# Joint order: ['FR_hip', 'FR_thigh', 'FR_calf', 'FL_hip', 'FL_thigh', 'FL_calf', 'RR_hip', 'RR_thigh', 'RR_calf', 'RL_hip', 'RL_thigh', 'RL_calf']
+default_angles: [-0.0, 0.8999999761581421, -1.7999999523162842, 0.0, 0.8999999761581421, -1.7999999523162842, -0.0, 0.8999999761581421, -1.7999999523162842, 0.0, 0.8999999761581421, -1.7999999523162842]
+action_scale: 0.05
+kp: 80.0
+kd: 1.0
+clip_actions: 23.7
+clip_observations: 100.0
+
+# Observation scales:
+ ang_vel: 0.25
+ dof_pos: 1.0
+ dof_vel: 0.05
+ contact_force: 1.0
+ command_scale: [2.0, 2.0, 0.25]
diff --git a/exports_go1_no_linevel/normalizer.npz b/exports_go1_no_linevel/normalizer.npz
new file mode 100644
index 0000000..0a598ca
Binary files /dev/null and b/exports_go1_no_linevel/normalizer.npz differ
diff --git a/exports_go1_orig/policy.onnx b/exports_go1_orig/policy.onnx
new file mode 100644
index 0000000..e275514
Binary files /dev/null and b/exports_go1_orig/policy.onnx differ
diff --git a/exports_orig_comb/normalizer.npz b/exports_orig_comb/normalizer.npz
new file mode 100644
index 0000000..9cb04e2
Binary files /dev/null and b/exports_orig_comb/normalizer.npz differ
diff --git a/exports_orig_comb/policy.onnx b/exports_orig_comb/policy.onnx
new file mode 100644
index 0000000..592c126
Binary files /dev/null and b/exports_orig_comb/policy.onnx differ
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/__init__.py b/motrix_rl/src/motrix_rl/dwaq_rsl/__init__.py
new file mode 100644
index 0000000..045ae1c
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/__init__.py
@@ -0,0 +1,13 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+"""Faithful port of Manaro-Alpha/DreamWaQ's custom rsl_rl-1.0.2 for MotrixLab.
+
+Provides ActorCritic_DWAQ (CENet+Actor+Critic), PPO (with joint VAE training),
+RolloutStorage, OnPolicyRunner, and a MotrixLab NpEnv -> VecEnv adapter.
+"""
+from .actor_critic_dwaq import ActorCritic_DWAQ
+from .ppo import PPO
+from .rollout_storage import RolloutStorage
+from .runner import OnPolicyRunner
+from .env_wrapper import DwaqVecEnv
+
+__all__ = ["ActorCritic_DWAQ", "PPO", "RolloutStorage", "OnPolicyRunner", "DwaqVecEnv"]
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/actor_critic_dwaq.py b/motrix_rl/src/motrix_rl/dwaq_rsl/actor_critic_dwaq.py
new file mode 100644
index 0000000..23105be
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/actor_critic_dwaq.py
@@ -0,0 +1,169 @@
+from __future__ import annotations
+
+import torch
+import torch.nn as nn
+from torch.distributions import Normal
+
+class ActorCritic_DWAQ(nn.Module):
+ def __init__(self, num_actor_obs, num_critic_obs, num_actions, cenet_in_dim, cenet_out_dim, activation="elu", init_noise_std=1.0,):
+ super().__init__()
+
+ self.activation = get_activation(activation)
+ actor_input_dim = num_actor_obs
+ critic_input_dim = num_critic_obs
+
+ self.actor = nn.Sequential(
+ nn.Linear(actor_input_dim,512),
+ self.activation,
+ nn.Linear(512,256),
+ self.activation,
+ nn.Linear(256,128),
+ self.activation,
+ nn.Linear(128,num_actions)
+ )
+
+ self.critic = nn.Sequential(
+ nn.Linear(critic_input_dim,512),
+ self.activation,
+ nn.Linear(512,256),
+ self.activation,
+ nn.Linear(256,128),
+ self.activation,
+ nn.Linear(128,1)
+ )
+
+ self.encoder = nn.Sequential(
+ nn.Linear(cenet_in_dim,128),
+ self.activation,
+ nn.Linear(128,64),
+ self.activation,
+ )
+ self.encode_mean_latent = nn.Linear(64,cenet_out_dim-3)
+ self.encode_logvar_latent = nn.Linear(64,cenet_out_dim-3)
+ self.encode_mean_vel = nn.Linear(64,3)
+ self.encode_logvar_vel = nn.Linear(64,3)
+
+ self.decoder = nn.Sequential(
+ nn.Linear(cenet_out_dim,64),
+ self.activation,
+ nn.Linear(64,128),
+ self.activation,
+ nn.Linear(128,45)
+ )
+
+ self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
+ self.distribution = None
+ # disable args validation for speedup
+ Normal.set_default_validate_args = False
+
+ # seems that we get better performance without init
+ # self.init_memory_weights(self.memory_a, 0.001, 0.)
+ # self.init_memory_weights(self.memory_c, 0.001, 0.)
+
+ @staticmethod
+ # not used at the moment
+ def init_weights(sequential, scales):
+ [
+ torch.nn.init.orthogonal_(module.weight, gain=scales[idx])
+ for idx, module in enumerate(mod for mod in sequential if isinstance(mod, nn.Linear))
+ ]
+
+ def reset(self, dones=None):
+ pass
+
+ def forward(self):
+ raise NotImplementedError
+
+ def reparameterise(self,mean,logvar):
+ var = torch.exp(logvar*0.5)
+ code_temp = torch.randn_like(var)
+ code = mean + var*code_temp
+ return code
+
+ def cenet_forward(self,obs_history):
+ distribution = self.encoder(obs_history)
+ mean_latent = self.encode_mean_latent(distribution)
+ logvar_latent = self.encode_logvar_latent(distribution)
+ # var = torch.exp(logvar_latent*0.5)
+ # code_temp = torch.randn_like(var)
+ # code = mean_latent + var*code_temp
+ # print("latent : ",code[0])
+ mean_vel = self.encode_mean_vel(distribution)
+ logvar_vel = self.encode_mean_vel(distribution)
+ code_latent = self.reparameterise(mean_latent,logvar_latent)
+ code_vel = self.reparameterise(mean_vel,logvar_vel)
+ code = torch.cat((code_vel,code_latent),dim=-1)
+ decode = self.decoder(code)
+ return code,code_vel,decode,mean_vel,logvar_vel,mean_latent,logvar_latent
+
+ @property
+ def action_mean(self):
+ return self.distribution.mean
+
+ @property
+ def action_std(self):
+ return self.distribution.stddev
+
+ @property
+ def entropy(self):
+ return self.distribution.entropy().sum(dim=-1)
+
+ def update_distribution(self, observations):
+ mean = self.actor(observations)
+ # Clamp raw actor output before building the Normal distribution:
+ # on unstable terrain (hfield) the network can produce extreme values
+ # that overflow → NaN → crash the full PPO pipeline.
+ mean = torch.clamp(mean, -10.0, 10.0)
+ self.distribution = Normal(mean, mean * 0.0 + self.std)
+
+ def act(self, observations, obs_history, **kwargs):
+ code,_,decode,_,_,_,_ = self.cenet_forward(obs_history)
+ observations = torch.cat((code,observations),dim=-1)
+ self.update_distribution(observations)
+ return self.distribution.sample()
+
+ def get_actions_log_prob(self, actions):
+ return self.distribution.log_prob(actions).sum(dim=-1)
+
+ def act_inference(self, observations,obs_history):
+ code,_,decode,_,_,_,_ = self.cenet_forward(obs_history)
+ observations = torch.cat((code,observations),dim=-1)
+ actions_mean = self.actor(observations)
+ return actions_mean
+
+ def evaluate(self, critic_observations, **kwargs):
+ value = self.critic(critic_observations)
+ return value
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+def get_activation(act_name):
+ if act_name == "elu":
+ return nn.ELU()
+ elif act_name == "selu":
+ return nn.SELU()
+ elif act_name == "relu":
+ return nn.ReLU()
+ elif act_name == "crelu":
+ return nn.CReLU()
+ elif act_name == "lrelu":
+ return nn.LeakyReLU()
+ elif act_name == "tanh":
+ return nn.Tanh()
+ elif act_name == "sigmoid":
+ return nn.Sigmoid()
+ else:
+ print("invalid activation function!")
+ return None
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/env_wrapper.py b/motrix_rl/src/motrix_rl/dwaq_rsl/env_wrapper.py
new file mode 100644
index 0000000..93968df
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/env_wrapper.py
@@ -0,0 +1,106 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+"""Adapter: MotrixLab DreamWaQ NpEnv -> upstream rsl_rl-1.0.2 VecEnv interface.
+
+Upstream OnPolicyRunner expects env.step() to return a 7-tuple:
+ (obs, privileged_obs, prev_privileged_obs, obs_hist, rewards, dones, extras)
+and env.reset() to return (obs, privileged_obs, prev_privileged_obs, obs_hist).
+
+MotrixLab's DreamWaQTask (numpy, CPU physics) already provides:
+ state.obs -> obs (N, 45)
+ info["privileged_obs"] -> privileged (N, 235) = obs(45)+base_vel(3)+heights(187)
+ info["obs_history"] -> history (N, 5, 45) (flattened to 225)
+ state.terminated/truncated -> dones / time_outs
+
+This wrapper does numpy<->torch transfer and maintains prev_privileged_obs.
+"""
+import numpy as np
+import torch
+
+
+class DwaqVecEnv:
+ """Wraps a MotrixLab DreamWaQ NpEnv for the upstream DreamWaQ rsl_rl runner."""
+
+ def __init__(self, env, device, num_obs=45, num_privileged_obs=235,
+ num_obs_hist=5, num_actions=12, clip_actions=100.0):
+ self._env = env
+ self.device = device
+ self.num_envs = env.num_envs
+ self.num_obs = num_obs
+ self.num_privileged_obs = num_privileged_obs
+ self.num_obs_hist = num_obs_hist
+ self.num_actions = num_actions
+ self._clip_actions = clip_actions
+ self.max_episode_length = int(getattr(env.cfg, "max_episode_steps", None) or 1000)
+
+ n = self.num_envs
+ self.obs_buf = torch.zeros(n, num_obs, device=device)
+ self.obs_hist_buf = torch.zeros(n, num_obs_hist * num_obs, device=device)
+ self.privileged_obs_buf = torch.zeros(n, num_privileged_obs, device=device)
+ self.prev_privileged_obs_buf = torch.zeros(n, num_privileged_obs, device=device)
+ self.rew_buf = torch.zeros(n, device=device)
+ self.reset_buf = torch.ones(n, device=device, dtype=torch.long)
+ self.episode_length_buf = torch.zeros(n, device=device, dtype=torch.long)
+ self.extras = {}
+
+ self._env.init_state()
+ self._sync_from_state(self._env._state)
+
+ # ── helpers ──
+ def _t(self, arr):
+ return torch.from_numpy(np.ascontiguousarray(arr, dtype=np.float32)).to(self.device)
+
+ def _sync_from_state(self, state):
+ self.obs_buf = self._t(state.obs)
+ info = state.info
+ priv = info.get("privileged_obs")
+ if priv is None:
+ priv = np.zeros((self.num_envs, self.num_privileged_obs), dtype=np.float32)
+ self.privileged_obs_buf = self._t(priv)
+ hist = info.get("obs_history") # (N, 5, 45)
+ if hist is None:
+ hist = np.zeros((self.num_envs, self.num_obs_hist, self.num_obs), dtype=np.float32)
+ self.obs_hist_buf = self._t(hist).reshape(self.num_envs, -1)
+ self.episode_length_buf = torch.from_numpy(
+ np.asarray(info.get("steps", np.zeros(self.num_envs)), dtype=np.int64)).to(self.device)
+
+ # ── VecEnv interface ──
+ def step(self, actions):
+ actions = torch.clip(actions, -self._clip_actions, self._clip_actions)
+ actions_np = actions.detach().cpu().numpy().astype(np.float32)
+ # prev_privileged = privileged BEFORE this step (matches upstream)
+ self.prev_privileged_obs_buf = self.privileged_obs_buf.clone()
+ state = self._env.step(actions_np)
+ self._sync_from_state(state)
+
+ done = np.logical_or(state.terminated, state.truncated)
+ self.rew_buf = self._t(state.reward).reshape(-1)
+ self.reset_buf = torch.from_numpy(done.astype(np.int64)).to(self.device)
+ self.extras = {
+ "time_outs": torch.from_numpy(np.asarray(state.truncated, dtype=np.float32)).to(self.device),
+ }
+ # pass per-term reward logging through if present
+ terms = state.info.get("reward_terms")
+ if terms:
+ self.extras["episode"] = {f"rew_{k}": float(v) for k, v in terms.items()}
+ # Log terrain_level so tensorboard tracks curriculum progress.
+ tl = state.info.get("terrain_level")
+ if tl is not None:
+ self.extras["episode"] = self.extras.get("episode", {})
+ self.extras["episode"]["terrain_level"] = float(np.asarray(tl).mean())
+
+ # rewards/dones are (N,) like upstream (storage views them to (N,1) internally)
+ return (self.obs_buf, self.privileged_obs_buf, self.prev_privileged_obs_buf,
+ self.obs_hist_buf, self.rew_buf, self.reset_buf, self.extras)
+
+ def reset(self, env_ids=None):
+ self._env.init_state()
+ self._sync_from_state(self._env._state)
+ self.prev_privileged_obs_buf = torch.zeros_like(self.privileged_obs_buf)
+ return (self.obs_buf, self.privileged_obs_buf,
+ self.prev_privileged_obs_buf, self.obs_hist_buf)
+
+ def get_observations(self):
+ return self.obs_buf, self.obs_hist_buf
+
+ def get_privileged_observations(self):
+ return self.privileged_obs_buf, self.prev_privileged_obs_buf
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/ppo.py b/motrix_rl/src/motrix_rl/dwaq_rsl/ppo.py
new file mode 100644
index 0000000..416ce70
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/ppo.py
@@ -0,0 +1,203 @@
+# SPDX-FileCopyrightText: Copyright (c) 2021 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: BSD-3-Clause
+#
+# Redistribution and use in source and binary forms, with or without
+# modification, are permitted provided that the following conditions are met:
+#
+# 1. Redistributions of source code must retain the above copyright notice, this
+# list of conditions and the following disclaimer.
+#
+# 2. Redistributions in binary form must reproduce the above copyright notice,
+# this list of conditions and the following disclaimer in the documentation
+# and/or other materials provided with the distribution.
+#
+# 3. Neither the name of the copyright holder nor the names of its
+# contributors may be used to endorse or promote products derived from
+# this software without specific prior written permission.
+#
+# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
+# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
+# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
+# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
+# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
+# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
+# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+#
+# Copyright (c) 2021 ETH Zurich, Nikita Rudin
+
+import torch
+import torch.nn as nn
+import torch.optim as optim
+
+from .actor_critic_dwaq import ActorCritic_DWAQ
+from .rollout_storage import RolloutStorage
+
+class PPO:
+ actor_critic: ActorCritic_DWAQ
+ def __init__(self,
+ actor_critic,
+ num_learning_epochs=1,
+ num_mini_batches=1,
+ clip_param=0.2,
+ gamma=0.99,
+ lam=0.95,
+ value_loss_coef=1.0,
+ entropy_coef=0.0,
+ learning_rate=1e-3,
+ max_grad_norm=1.0,
+ use_clipped_value_loss=True,
+ schedule="fixed",
+ desired_kl=0.01,
+ device='cpu',
+ ):
+
+ self.device = device
+
+ self.desired_kl = desired_kl
+ self.schedule = schedule
+ self.learning_rate = learning_rate
+
+ # PPO components
+ self.actor_critic = actor_critic
+ self.actor_critic.to(self.device)
+ self.storage = None # initialized later
+ self.optimizer = optim.Adam(self.actor_critic.parameters(), lr=learning_rate)
+ self.transition = RolloutStorage.Transition()
+
+ # PPO parameters
+ self.clip_param = clip_param
+ self.num_learning_epochs = num_learning_epochs
+ self.num_mini_batches = num_mini_batches
+ self.value_loss_coef = value_loss_coef
+ self.entropy_coef = entropy_coef
+ self.gamma = gamma
+ self.lam = lam
+ self.max_grad_norm = max_grad_norm
+ self.use_clipped_value_loss = use_clipped_value_loss
+
+ def init_storage(self, num_envs, num_transitions_per_env, actor_obs_shape, critic_obs_shape, obs_hist_shape, action_shape):
+ self.storage = RolloutStorage(num_envs, num_transitions_per_env, actor_obs_shape, critic_obs_shape, obs_hist_shape, action_shape, self.device)
+
+ def test_mode(self):
+ self.actor_critic.test()
+
+ def train_mode(self):
+ self.actor_critic.train()
+
+ def act(self, obs, critic_obs, prev_critic_obs, obs_history):
+ # if self.actor_critic.is_recurrent:
+ # self.transition.hidden_states = self.actor_critic.get_hidden_states()
+ # Compute the actions and values
+ self.transition.actions = self.actor_critic.act(obs,obs_history).detach()
+ self.transition.values = self.actor_critic.evaluate(critic_obs).detach()
+ self.transition.actions_log_prob = self.actor_critic.get_actions_log_prob(self.transition.actions).detach()
+ self.transition.action_mean = self.actor_critic.action_mean.detach()
+ self.transition.action_sigma = self.actor_critic.action_std.detach()
+ # need to record obs and critic_obs before env.step()
+ self.transition.observations = obs
+ self.transition.observation_history = obs_history
+ self.transition.critic_observations = critic_obs
+ self.transition.prev_critic_obs = prev_critic_obs
+ return self.transition.actions
+
+ def process_env_step(self, rewards, dones, infos):
+ self.transition.rewards = rewards.clone()
+ self.transition.dones = dones
+ # Bootstrapping on time outs
+ if 'time_outs' in infos:
+ self.transition.rewards += self.gamma * torch.squeeze(self.transition.values * infos['time_outs'].unsqueeze(1).to(self.device), 1)
+
+ # Record the transition
+ self.storage.add_transitions(self.transition)
+ self.transition.clear()
+ self.actor_critic.reset(dones)
+
+ def compute_returns(self, last_critic_obs):
+ last_values= self.actor_critic.evaluate(last_critic_obs).detach()
+ self.storage.compute_returns(last_values, self.gamma, self.lam)
+
+ def update(self,beta=1):
+ mean_value_loss = 0
+ mean_surrogate_loss = 0
+ mean_autoenc_loss = 0
+ # if self.actor_critic.is_recurrent:
+ # generator = self.storage.reccurent_mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
+ # else:
+ generator = self.storage.mini_batch_generator(self.num_mini_batches, self.num_learning_epochs)
+ for obs_batch, critic_obs_batch, prev_critic_obs_batch, obs_hist_batch, actions_batch, target_values_batch, advantages_batch, returns_batch, old_actions_log_prob_batch, \
+ old_mu_batch, old_sigma_batch, hid_states_batch, masks_batch in generator:
+
+
+ self.actor_critic.act(obs_batch, obs_hist_batch, masks=masks_batch, hidden_states=hid_states_batch[0])
+ actions_log_prob_batch = self.actor_critic.get_actions_log_prob(actions_batch)
+ value_batch = self.actor_critic.evaluate(critic_obs_batch, masks=masks_batch, hidden_states=hid_states_batch[1])
+ mu_batch = self.actor_critic.action_mean
+ sigma_batch = self.actor_critic.action_std
+ entropy_batch = self.actor_critic.entropy
+
+ # KL
+ if self.desired_kl != None and self.schedule == 'adaptive':
+ with torch.inference_mode():
+ kl = torch.sum(
+ torch.log(sigma_batch / old_sigma_batch + 1.e-5) + (torch.square(old_sigma_batch) + torch.square(old_mu_batch - mu_batch)) / (2.0 * torch.square(sigma_batch)) - 0.5, axis=-1)
+ kl_mean = torch.mean(kl)
+
+ if kl_mean > self.desired_kl * 2.0:
+ self.learning_rate = max(1e-5, self.learning_rate / 1.5)
+ elif kl_mean < self.desired_kl / 2.0 and kl_mean > 0.0:
+ self.learning_rate = min(1e-2, self.learning_rate * 1.5)
+
+ for param_group in self.optimizer.param_groups:
+ param_group['lr'] = self.learning_rate
+
+
+ #Beta VAE loss
+ code,code_vel,decode,mean_vel,logvar_vel,mean_latent,logvar_latent = self.actor_critic.cenet_forward(obs_hist_batch)
+
+ vel_target = prev_critic_obs_batch[:,45:48]
+ decode_target = obs_batch
+ vel_target.requires_grad = False
+ decode_target.requires_grad = False
+ autoenc_loss = (nn.MSELoss()(code_vel,vel_target) + nn.MSELoss()(decode,decode_target) + beta*(-0.5 * torch.sum(1 + logvar_latent - mean_latent.pow(2) - logvar_latent.exp())))/self.num_mini_batches
+ # estimation_loss = (code[:,0:3] - prev_critic_obs_batch[:,45:48]).pow(2).mean()
+ # reconst_loss = (decode - obs_batch).pow(2).mean()
+ # latent_loss = beta*(-0.5 * torch.sum(1 + logvar - mean.pow(2) - logvar.exp()))/mean.shape[0]
+ # autoenc_loss = estimation_loss + reconst_loss + latent_loss
+ # Surrogate loss
+ ratio = torch.exp(actions_log_prob_batch - torch.squeeze(old_actions_log_prob_batch))
+ surrogate = -torch.squeeze(advantages_batch) * ratio
+ surrogate_clipped = -torch.squeeze(advantages_batch) * torch.clamp(ratio, 1.0 - self.clip_param,
+ 1.0 + self.clip_param)
+ surrogate_loss = torch.max(surrogate, surrogate_clipped).mean()
+
+ # Value function loss
+ if self.use_clipped_value_loss:
+ value_clipped = target_values_batch + (value_batch - target_values_batch).clamp(-self.clip_param,
+ self.clip_param)
+ value_losses = (value_batch - returns_batch).pow(2)
+ value_losses_clipped = (value_clipped - returns_batch).pow(2)
+ value_loss = torch.max(value_losses, value_losses_clipped).mean()
+ else:
+ value_loss = (returns_batch - value_batch).pow(2).mean()
+
+ loss = surrogate_loss + self.value_loss_coef * value_loss - self.entropy_coef * entropy_batch.mean() + autoenc_loss
+
+ # Gradient step
+ self.optimizer.zero_grad()
+ loss.backward()
+ nn.utils.clip_grad_norm_(self.actor_critic.parameters(), self.max_grad_norm)
+ self.optimizer.step()
+
+ mean_value_loss += value_loss.item()
+ mean_surrogate_loss += surrogate_loss.item()
+ mean_autoenc_loss += autoenc_loss.item()
+
+ num_updates = self.num_learning_epochs * self.num_mini_batches
+ mean_value_loss /= num_updates
+ mean_surrogate_loss /= num_updates
+ self.storage.clear()
+
+ return mean_value_loss, mean_surrogate_loss, mean_autoenc_loss
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/rollout_storage.py b/motrix_rl/src/motrix_rl/dwaq_rsl/rollout_storage.py
new file mode 100644
index 0000000..26a99f5
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/rollout_storage.py
@@ -0,0 +1,246 @@
+# SPDX-FileCopyrightText: Copyright (c) 2021 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: BSD-3-Clause
+#
+# Redistribution and use in source and binary forms, with or without
+# modification, are permitted provided that the following conditions are met:
+#
+# 1. Redistributions of source code must retain the above copyright notice, this
+# list of conditions and the following disclaimer.
+#
+# 2. Redistributions in binary form must reproduce the above copyright notice,
+# this list of conditions and the following disclaimer in the documentation
+# and/or other materials provided with the distribution.
+#
+# 3. Neither the name of the copyright holder nor the names of its
+# contributors may be used to endorse or promote products derived from
+# this software without specific prior written permission.
+#
+# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
+# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
+# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
+# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
+# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
+# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
+# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+#
+# Copyright (c) 2021 ETH Zurich, Nikita Rudin
+
+import torch
+import numpy as np
+
+from .utils import split_and_pad_trajectories
+
+class RolloutStorage:
+ class Transition:
+ def __init__(self):
+ self.observations = None
+ self.critic_observations = None
+ self.prev_critic_obs = None
+ self.observation_history = None
+ self.actions = None
+ self.rewards = None
+ self.dones = None
+ self.values = None
+ self.actions_log_prob = None
+ self.action_mean = None
+ self.action_sigma = None
+ self.hidden_states = None
+
+ def clear(self):
+ self.__init__()
+
+ def __init__(self, num_envs, num_transitions_per_env, obs_shape, privileged_obs_shape, obs_hist_shape, actions_shape, device="cpu"):
+
+ self.device = device
+
+ self.obs_shape = obs_shape
+ self.privileged_obs_shape = privileged_obs_shape
+ self.actions_shape = actions_shape
+
+ # Core
+ self.observations = torch.zeros(num_transitions_per_env, num_envs, *obs_shape, device=self.device)
+ if privileged_obs_shape[0] is not None:
+ self.privileged_observations = torch.zeros(num_transitions_per_env, num_envs, *privileged_obs_shape, device=self.device)
+ else:
+ self.privileged_observations = None
+
+ self.prev_critic_obs = torch.zeros(num_transitions_per_env, num_envs, *privileged_obs_shape, device=self.device)
+ self.observation_history = torch.zeros(num_transitions_per_env, num_envs, *obs_hist_shape, device=self.device)
+ self.rewards = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
+ self.actions = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
+ self.dones = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device).byte()
+
+ # For PPO
+ self.actions_log_prob = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
+ self.values = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
+ self.returns = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
+ self.advantages = torch.zeros(num_transitions_per_env, num_envs, 1, device=self.device)
+ self.mu = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
+ self.sigma = torch.zeros(num_transitions_per_env, num_envs, *actions_shape, device=self.device)
+
+ self.num_transitions_per_env = num_transitions_per_env
+ self.num_envs = num_envs
+
+ # rnn
+ self.saved_hidden_states_a = None
+ self.saved_hidden_states_c = None
+
+ self.step = 0
+
+ def add_transitions(self, transition: Transition):
+ if self.step >= self.num_transitions_per_env:
+ raise AssertionError("Rollout buffer overflow")
+ self.observations[self.step].copy_(transition.observations)
+ if self.privileged_observations is not None: self.privileged_observations[self.step].copy_(transition.critic_observations)
+ self.prev_critic_obs[self.step].copy_(transition.prev_critic_obs)
+ self.observation_history[self.step].copy_(transition.observation_history)
+ self.actions[self.step].copy_(transition.actions)
+ self.rewards[self.step].copy_(transition.rewards.view(-1, 1))
+ self.dones[self.step].copy_(transition.dones.view(-1, 1))
+ self.values[self.step].copy_(transition.values)
+ self.actions_log_prob[self.step].copy_(transition.actions_log_prob.view(-1, 1))
+ self.mu[self.step].copy_(transition.action_mean)
+ self.sigma[self.step].copy_(transition.action_sigma)
+ self._save_hidden_states(transition.hidden_states)
+ self.step += 1
+
+ def _save_hidden_states(self, hidden_states):
+ if hidden_states is None or hidden_states==(None, None):
+ return
+ # make a tuple out of GRU hidden state sto match the LSTM format
+ hid_a = hidden_states[0] if isinstance(hidden_states[0], tuple) else (hidden_states[0],)
+ hid_c = hidden_states[1] if isinstance(hidden_states[1], tuple) else (hidden_states[1],)
+
+ # initialize if needed
+ if self.saved_hidden_states_a is None:
+ self.saved_hidden_states_a = [torch.zeros(self.observations.shape[0], *hid_a[i].shape, device=self.device) for i in range(len(hid_a))]
+ self.saved_hidden_states_c = [torch.zeros(self.observations.shape[0], *hid_c[i].shape, device=self.device) for i in range(len(hid_c))]
+ # copy the states
+ for i in range(len(hid_a)):
+ self.saved_hidden_states_a[i][self.step].copy_(hid_a[i])
+ self.saved_hidden_states_c[i][self.step].copy_(hid_c[i])
+
+
+ def clear(self):
+ self.step = 0
+
+ def compute_returns(self, last_values, gamma, lam):
+ advantage = 0
+ for step in reversed(range(self.num_transitions_per_env)):
+ if step == self.num_transitions_per_env - 1:
+ next_values = last_values
+ else:
+ next_values = self.values[step + 1]
+ next_is_not_terminal = 1.0 - self.dones[step].float()
+ delta = self.rewards[step] + next_is_not_terminal * gamma * next_values - self.values[step]
+ advantage = delta + next_is_not_terminal * gamma * lam * advantage
+ self.returns[step] = advantage + self.values[step]
+
+ # Compute and normalize the advantages
+ self.advantages = self.returns - self.values
+ self.advantages = (self.advantages - self.advantages.mean()) / (self.advantages.std() + 1e-8)
+
+ def get_statistics(self):
+ done = self.dones
+ done[-1] = 1
+ flat_dones = done.permute(1, 0, 2).reshape(-1, 1)
+ done_indices = torch.cat((flat_dones.new_tensor([-1], dtype=torch.int64), flat_dones.nonzero(as_tuple=False)[:, 0]))
+ trajectory_lengths = (done_indices[1:] - done_indices[:-1])
+ return trajectory_lengths.float().mean(), self.rewards.mean()
+
+ def mini_batch_generator(self, num_mini_batches, num_epochs=8):
+ batch_size = self.num_envs * self.num_transitions_per_env
+ mini_batch_size = batch_size // num_mini_batches
+ indices = torch.randperm(num_mini_batches*mini_batch_size, requires_grad=False, device=self.device)
+
+ observations = self.observations.flatten(0, 1)
+ if self.privileged_observations is not None:
+ critic_observations = self.privileged_observations.flatten(0, 1)
+ else:
+ critic_observations = observations
+
+ prev_critic_obs = self.prev_critic_obs.flatten(0,1)
+ obs_history = self.observation_history.flatten(0,1)
+ actions = self.actions.flatten(0, 1)
+ values = self.values.flatten(0, 1)
+ returns = self.returns.flatten(0, 1)
+ old_actions_log_prob = self.actions_log_prob.flatten(0, 1)
+ advantages = self.advantages.flatten(0, 1)
+ old_mu = self.mu.flatten(0, 1)
+ old_sigma = self.sigma.flatten(0, 1)
+
+ for epoch in range(num_epochs):
+ for i in range(num_mini_batches):
+
+ start = i*mini_batch_size
+ end = (i+1)*mini_batch_size
+ batch_idx = indices[start:end]
+
+ obs_batch = observations[batch_idx]
+ critic_observations_batch = critic_observations[batch_idx]
+ prev_critic_obs_batch = prev_critic_obs[batch_idx]
+ obs_hist_batch = obs_history[batch_idx]
+ actions_batch = actions[batch_idx]
+ target_values_batch = values[batch_idx]
+ returns_batch = returns[batch_idx]
+ old_actions_log_prob_batch = old_actions_log_prob[batch_idx]
+ advantages_batch = advantages[batch_idx]
+ old_mu_batch = old_mu[batch_idx]
+ old_sigma_batch = old_sigma[batch_idx]
+ yield obs_batch, critic_observations_batch,prev_critic_obs_batch,obs_hist_batch, actions_batch, target_values_batch, advantages_batch, returns_batch, \
+ old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (None, None), None
+
+ # for RNNs only
+ def reccurent_mini_batch_generator(self, num_mini_batches, num_epochs=8):
+
+ padded_obs_trajectories, trajectory_masks = split_and_pad_trajectories(self.observations, self.dones)
+ if self.privileged_observations is not None:
+ padded_critic_obs_trajectories, _ = split_and_pad_trajectories(self.privileged_observations, self.dones)
+ else:
+ padded_critic_obs_trajectories = padded_obs_trajectories
+
+ mini_batch_size = self.num_envs // num_mini_batches
+ for ep in range(num_epochs):
+ first_traj = 0
+ for i in range(num_mini_batches):
+ start = i*mini_batch_size
+ stop = (i+1)*mini_batch_size
+
+ dones = self.dones.squeeze(-1)
+ last_was_done = torch.zeros_like(dones, dtype=torch.bool)
+ last_was_done[1:] = dones[:-1]
+ last_was_done[0] = True
+ trajectories_batch_size = torch.sum(last_was_done[:, start:stop])
+ last_traj = first_traj + trajectories_batch_size
+
+ masks_batch = trajectory_masks[:, first_traj:last_traj]
+ obs_batch = padded_obs_trajectories[:, first_traj:last_traj]
+ critic_obs_batch = padded_critic_obs_trajectories[:, first_traj:last_traj]
+
+ actions_batch = self.actions[:, start:stop]
+ old_mu_batch = self.mu[:, start:stop]
+ old_sigma_batch = self.sigma[:, start:stop]
+ returns_batch = self.returns[:, start:stop]
+ advantages_batch = self.advantages[:, start:stop]
+ values_batch = self.values[:, start:stop]
+ old_actions_log_prob_batch = self.actions_log_prob[:, start:stop]
+
+ # reshape to [num_envs, time, num layers, hidden dim] (original shape: [time, num_layers, num_envs, hidden_dim])
+ # then take only time steps after dones (flattens num envs and time dimensions),
+ # take a batch of trajectories and finally reshape back to [num_layers, batch, hidden_dim]
+ last_was_done = last_was_done.permute(1, 0)
+ hid_a_batch = [ saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj].transpose(1, 0).contiguous()
+ for saved_hidden_states in self.saved_hidden_states_a ]
+ hid_c_batch = [ saved_hidden_states.permute(2, 0, 1, 3)[last_was_done][first_traj:last_traj].transpose(1, 0).contiguous()
+ for saved_hidden_states in self.saved_hidden_states_c ]
+ # remove the tuple for GRU
+ hid_a_batch = hid_a_batch[0] if len(hid_a_batch)==1 else hid_a_batch
+ hid_c_batch = hid_c_batch[0] if len(hid_c_batch)==1 else hid_a_batch
+
+ yield obs_batch, critic_obs_batch, actions_batch, values_batch, advantages_batch, returns_batch, \
+ old_actions_log_prob_batch, old_mu_batch, old_sigma_batch, (hid_a_batch, hid_c_batch), masks_batch
+
+ first_traj = last_traj
\ No newline at end of file
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/runner.py b/motrix_rl/src/motrix_rl/dwaq_rsl/runner.py
new file mode 100644
index 0000000..f97cd14
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/runner.py
@@ -0,0 +1,180 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+"""DreamWaQ OnPolicyRunner — adapted from upstream rsl_rl-1.0.2.
+
+Faithful port of Manaro-Alpha/DreamWaQ's on_policy_runner.py, with the
+rsl_rl-package imports replaced by local (relative) imports and the isaacgym
+noise removed. Drives the PPO + CENet (VAE) joint training loop.
+"""
+import time
+import os
+from collections import deque
+import statistics
+
+from torch.utils.tensorboard import SummaryWriter
+import torch
+
+from .ppo import PPO
+from .actor_critic_dwaq import ActorCritic_DWAQ
+
+
+class OnPolicyRunner:
+ def __init__(self, env, train_cfg, log_dir=None, device="cpu"):
+ self.cfg = train_cfg["runner"]
+ self.alg_cfg = train_cfg["algorithm"]
+ self.policy_cfg = train_cfg["policy"]
+ self.device = device
+ self.env = env
+
+ num_critic_obs = env.num_privileged_obs if env.num_privileged_obs is not None else env.num_obs
+ cenet_in_dim = env.num_obs_hist * env.num_obs # 5*45 = 225
+ cenet_out_dim = 19 # vel(3) + latent(16)
+
+ actor_critic = ActorCritic_DWAQ(
+ env.num_obs + cenet_out_dim, # actor input: obs(45) + code(19) = 64
+ num_critic_obs, # critic input: privileged_obs
+ env.num_actions, # 12
+ cenet_in_dim, # 225
+ cenet_out_dim, # 19
+ **self.policy_cfg,
+ ).to(self.device)
+
+ self.alg = PPO(actor_critic, device=self.device, **self.alg_cfg)
+ self.num_steps_per_env = self.cfg["num_steps_per_env"]
+ self.save_interval = self.cfg["save_interval"]
+
+ self.alg.init_storage(
+ env.num_envs, self.num_steps_per_env,
+ [env.num_obs], [env.num_privileged_obs],
+ [env.num_obs_hist * env.num_obs], [env.num_actions],
+ )
+
+ self.log_dir = log_dir
+ self.writer = None
+ self.tot_timesteps = 0
+ self.tot_time = 0
+ self.current_learning_iteration = 0
+
+ _, _, _, _ = self.env.reset()
+
+ def learn(self, num_learning_iterations, init_at_random_ep_len=False):
+ if self.log_dir is not None and self.writer is None:
+ self.writer = SummaryWriter(log_dir=self.log_dir, flush_secs=10)
+ if init_at_random_ep_len:
+ self.env.episode_length_buf = torch.randint_like(
+ self.env.episode_length_buf, high=int(self.env.max_episode_length))
+
+ obs, obs_hist = self.env.get_observations()
+ privileged_obs, prev_critic_obs = self.env.get_privileged_observations()
+ critic_obs = privileged_obs if privileged_obs is not None else obs
+ obs, critic_obs, prev_critic_obs, obs_hist = (
+ obs.to(self.device), critic_obs.to(self.device),
+ prev_critic_obs.to(self.device), obs_hist.to(self.device))
+ self.alg.actor_critic.train()
+
+ ep_infos = []
+ rewbuffer = deque(maxlen=100)
+ lenbuffer = deque(maxlen=100)
+ cur_reward_sum = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
+ cur_episode_length = torch.zeros(self.env.num_envs, dtype=torch.float, device=self.device)
+
+ tot_iter = self.current_learning_iteration + num_learning_iterations
+ for it in range(self.current_learning_iteration, tot_iter):
+ start = time.time()
+ with torch.inference_mode():
+ for i in range(self.num_steps_per_env):
+ actions = self.alg.act(obs, critic_obs, prev_critic_obs, obs_hist)
+ obs, privileged_obs, prev_privileged_obs, obs_hist, rewards, dones, infos = self.env.step(actions)
+ critic_obs = privileged_obs if privileged_obs is not None else obs
+ prev_critic_obs = prev_privileged_obs
+ obs, critic_obs, prev_critic_obs, obs_hist, rewards, dones = (
+ obs.to(self.device), critic_obs.to(self.device),
+ prev_critic_obs.to(self.device), obs_hist.to(self.device),
+ rewards.to(self.device), dones.to(self.device))
+ self.alg.process_env_step(rewards, dones, infos)
+
+ if self.log_dir is not None:
+ if "episode" in infos:
+ ep_infos.append(infos["episode"])
+ cur_reward_sum += rewards
+ cur_episode_length += 1
+ new_ids = (dones > 0).nonzero(as_tuple=False)
+ rewbuffer.extend(cur_reward_sum[new_ids][:, 0].cpu().numpy().tolist())
+ lenbuffer.extend(cur_episode_length[new_ids][:, 0].cpu().numpy().tolist())
+ cur_reward_sum[new_ids] = 0
+ cur_episode_length[new_ids] = 0
+
+ stop = time.time()
+ collection_time = stop - start
+ start = stop
+ self.alg.compute_returns(critic_obs)
+
+ mean_value_loss, mean_surrogate_loss, mean_autoenc_loss = self.alg.update()
+ stop = time.time()
+ learn_time = stop - start
+ if self.log_dir is not None:
+ self.log(locals())
+ if it % self.save_interval == 0:
+ self.save(os.path.join(self.log_dir, "model_{}.pt".format(it)))
+ ep_infos.clear()
+
+ self.current_learning_iteration += num_learning_iterations
+ self.save(os.path.join(self.log_dir, "model_{}.pt".format(self.current_learning_iteration)))
+
+ def log(self, locs, width=80, pad=35):
+ self.tot_timesteps += self.num_steps_per_env * self.env.num_envs
+ self.tot_time += locs["collection_time"] + locs["learn_time"]
+ iteration_time = locs["collection_time"] + locs["learn_time"]
+
+ if locs["ep_infos"]:
+ for key in locs["ep_infos"][0]:
+ infotensor = torch.tensor([], device=self.device)
+ for ep_info in locs["ep_infos"]:
+ v = ep_info[key]
+ if not isinstance(v, torch.Tensor):
+ v = torch.Tensor([v])
+ if len(v.shape) == 0:
+ v = v.unsqueeze(0)
+ infotensor = torch.cat((infotensor, v.to(self.device)))
+ self.writer.add_scalar("Episode/" + key, torch.mean(infotensor), locs["it"])
+ mean_std = self.alg.actor_critic.std.mean()
+ fps = int(self.num_steps_per_env * self.env.num_envs / (locs["collection_time"] + locs["learn_time"]))
+
+ self.writer.add_scalar("Loss/value_function", locs["mean_value_loss"], locs["it"])
+ self.writer.add_scalar("Loss/surrogate", locs["mean_surrogate_loss"], locs["it"])
+ self.writer.add_scalar("Loss/autoenc_function", locs["mean_autoenc_loss"], locs["it"])
+ self.writer.add_scalar("Loss/learning_rate", self.alg.learning_rate, locs["it"])
+ self.writer.add_scalar("Policy/mean_noise_std", mean_std.item(), locs["it"])
+ self.writer.add_scalar("Perf/total_fps", fps, locs["it"])
+ if len(locs["rewbuffer"]) > 0:
+ self.writer.add_scalar("Train/mean_reward", statistics.mean(locs["rewbuffer"]), locs["it"])
+ self.writer.add_scalar("Train/mean_episode_length", statistics.mean(locs["lenbuffer"]), locs["it"])
+
+ if it_is_log := (locs["it"] % 10 == 0):
+ rew = statistics.mean(locs["rewbuffer"]) if len(locs["rewbuffer"]) > 0 else 0.0
+ elen = statistics.mean(locs["lenbuffer"]) if len(locs["lenbuffer"]) > 0 else 0.0
+ print(f"[it {locs['it']}/{self.current_learning_iteration + locs['num_learning_iterations']}] "
+ f"fps={fps} val_loss={locs['mean_value_loss']:.3f} surr={locs['mean_surrogate_loss']:.4f} "
+ f"vae={locs['mean_autoenc_loss']:.3f} std={mean_std.item():.2f} "
+ f"rew={rew:.2f} eplen={elen:.0f}")
+
+ def save(self, path, infos=None):
+ torch.save({
+ "model_state_dict": self.alg.actor_critic.state_dict(),
+ "optimizer_state_dict": self.alg.optimizer.state_dict(),
+ "iter": self.current_learning_iteration,
+ "infos": infos,
+ }, path)
+
+ def load(self, path, load_optimizer=True):
+ loaded_dict = torch.load(path, map_location=self.device)
+ self.alg.actor_critic.load_state_dict(loaded_dict["model_state_dict"])
+ if load_optimizer:
+ self.alg.optimizer.load_state_dict(loaded_dict["optimizer_state_dict"])
+ self.current_learning_iteration = loaded_dict["iter"]
+ return loaded_dict["infos"]
+
+ def get_inference_policy(self, device=None):
+ self.alg.actor_critic.eval()
+ if device is not None:
+ self.alg.actor_critic.to(device)
+ return self.alg.actor_critic.act_inference
diff --git a/motrix_rl/src/motrix_rl/dwaq_rsl/utils.py b/motrix_rl/src/motrix_rl/dwaq_rsl/utils.py
new file mode 100644
index 0000000..b6affab
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/dwaq_rsl/utils.py
@@ -0,0 +1,71 @@
+# SPDX-FileCopyrightText: Copyright (c) 2021 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: BSD-3-Clause
+#
+# Redistribution and use in source and binary forms, with or without
+# modification, are permitted provided that the following conditions are met:
+#
+# 1. Redistributions of source code must retain the above copyright notice, this
+# list of conditions and the following disclaimer.
+#
+# 2. Redistributions in binary form must reproduce the above copyright notice,
+# this list of conditions and the following disclaimer in the documentation
+# and/or other materials provided with the distribution.
+#
+# 3. Neither the name of the copyright holder nor the names of its
+# contributors may be used to endorse or promote products derived from
+# this software without specific prior written permission.
+#
+# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
+# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
+# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
+# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
+# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
+# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
+# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
+# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
+# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
+# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+#
+# Copyright (c) 2021 ETH Zurich, Nikita Rudin
+
+import torch
+
+def split_and_pad_trajectories(tensor, dones):
+ """ Splits trajectories at done indices. Then concatenates them and padds with zeros up to the length og the longest trajectory.
+ Returns masks corresponding to valid parts of the trajectories
+ Example:
+ Input: [ [a1, a2, a3, a4 | a5, a6],
+ [b1, b2 | b3, b4, b5 | b6]
+ ]
+
+ Output:[ [a1, a2, a3, a4], | [ [True, True, True, True],
+ [a5, a6, 0, 0], | [True, True, False, False],
+ [b1, b2, 0, 0], | [True, True, False, False],
+ [b3, b4, b5, 0], | [True, True, True, False],
+ [b6, 0, 0, 0] | [True, False, False, False],
+ ] | ]
+
+ Assumes that the inputy has the following dimension order: [time, number of envs, aditional dimensions]
+ """
+ dones = dones.clone()
+ dones[-1] = 1
+ # Permute the buffers to have order (num_envs, num_transitions_per_env, ...), for correct reshaping
+ flat_dones = dones.transpose(1, 0).reshape(-1, 1)
+
+ # Get length of trajectory by counting the number of successive not done elements
+ done_indices = torch.cat((flat_dones.new_tensor([-1], dtype=torch.int64), flat_dones.nonzero()[:, 0]))
+ trajectory_lengths = done_indices[1:] - done_indices[:-1]
+ trajectory_lengths_list = trajectory_lengths.tolist()
+ # Extract the individual trajectories
+ trajectories = torch.split(tensor.transpose(1, 0).flatten(0, 1),trajectory_lengths_list)
+ padded_trajectories = torch.nn.utils.rnn.pad_sequence(trajectories)
+
+
+ trajectory_masks = trajectory_lengths > torch.arange(0, tensor.shape[0], device=tensor.device).unsqueeze(1)
+ return padded_trajectories, trajectory_masks
+
+def unpad_trajectories(trajectories, masks):
+ """ Does the inverse operation of split_and_pad_trajectories()
+ """
+ # Need to transpose before and after the masking to have proper reshaping
+ return trajectories.transpose(1, 0)[masks.transpose(1, 0)].view(-1, trajectories.shape[0], trajectories.shape[-1]).transpose(1, 0)
\ No newline at end of file
diff --git a/motrix_rl/src/motrix_rl/rslrl/torch/models/cenet_actor.py b/motrix_rl/src/motrix_rl/rslrl/torch/models/cenet_actor.py
index ce310e5..cbbb1bf 100644
--- a/motrix_rl/src/motrix_rl/rslrl/torch/models/cenet_actor.py
+++ b/motrix_rl/src/motrix_rl/rslrl/torch/models/cenet_actor.py
@@ -116,15 +116,15 @@ class CENetActorModel(MLPModel):
self._last_cenet_output = None
def _update_distribution(self, obs: torch.Tensor) -> None:
- """覆盖父类 — 确保 std 始终为正,防止 NaN。"""
- super()._update_distribution(obs)
- # 如果 std 因数值问题变负,clamp 到最小值
+ """覆盖父类 — 强制 std > 0 再创建 Normal 分布(防止 NaN)。"""
+ # 先 clamp std,再调父类创建分布
if self.stochastic and not self.state_dependent_std:
with torch.no_grad():
if self.noise_std_type == "scalar":
self.std.clamp_(min=1e-6)
elif self.noise_std_type == "log":
self.log_std.clamp_(min=-20.0, max=10.0)
+ super()._update_distribution(obs)
def _get_latent_dim(self) -> int:
"""Actor 实际输入:code(19) + policy(45) = 64。"""
diff --git a/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py b/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py
index 2cb8f05..bdb73db 100644
--- a/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py
+++ b/motrix_rl/src/motrix_rl/rslrl/torch/train/dreamwaq_ppo.py
@@ -81,6 +81,9 @@ class DreamWaQPPO(PPO):
returns_batch, old_actions_log_prob_batch, old_mu_batch,
old_sigma_batch, hid_states_batch, masks_batch,
) in generator:
+ # NaN 检测 — 数据有 NaN 就跳过这个 batch
+ if torch.isnan(obs_batch["policy"]).any() or torch.isnan(obs_batch["obs_history"]).any():
+ continue
# ── 标准 PPO 前向 ──
self.actor(obs_batch, masks=masks_batch, stochastic_output=True)
self.critic(obs_batch, masks=masks_batch)
diff --git a/motrix_rl/src/motrix_rl/skrl/jax/train/cts_ppo.py b/motrix_rl/src/motrix_rl/skrl/jax/train/cts_ppo.py
new file mode 100644
index 0000000..b98093c
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/skrl/jax/train/cts_ppo.py
@@ -0,0 +1,306 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+"""CTS PPO v2 — Teacher/Student encoders with distillation loss.
+
+Matches go2_rl_gym CTS training flow:
+ Phase 1: Standard PPO update with teacher-computed latent
+ Phase 2: Distillation: MSE(teacher_latent, student_latent) on student envs
+
+Architecture:
+ Teacher: privileged_obs(76) → MLP(512,256) → L2Norm → latent(32)
+ Student: history_obs(225) → MLP(512,256,256) → L2Norm → latent(32)
+ Actor: student_obs(45) + latent(32) = 77 → [512,256,128] → action(12)
+ Critic: student_obs(45) + latent(32) = 77 → [512,256,128] → value(1)
+
+The wrapper computes latent online (teacher for 75% envs, student for 25%)
+and augments observations to 77-dim for standard PPO. After each PPO update,
+a distillation step trains the student encoder to match the teacher's latent.
+"""
+
+import flax.linen as nn
+import jax
+import jax.numpy as jnp
+import numpy as np
+import optax
+
+from skrl.agents.jax.ppo import PPO as BasePPO
+from skrl.memories.jax import RandomMemory
+from skrl.trainers.jax import SequentialTrainer
+from skrl.utils import set_seed
+
+from motrix_envs import registry as env_registry
+from motrix_rl import registry, utils
+from motrix_rl.skrl import get_log_dir
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Encoder Modules (Flax)
+# ═══════════════════════════════════════════════════════════════════════════
+
+class TeacherEncoder(nn.Module):
+ latent_dim: int = 32
+ hidden_dims: tuple = (512, 256)
+
+ @nn.compact
+ def __call__(self, x):
+ for d in self.hidden_dims:
+ x = nn.elu(nn.Dense(d)(x))
+ return nn.Dense(self.latent_dim)(x)
+
+
+class StudentEncoder(nn.Module):
+ latent_dim: int = 32
+ hidden_dims: tuple = (512, 256, 256)
+
+ @nn.compact
+ def __call__(self, x):
+ for d in self.hidden_dims:
+ x = nn.elu(nn.Dense(d)(x))
+ return nn.Dense(self.latent_dim)(x)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# CTS Env Wrapper
+# ═══════════════════════════════════════════════════════════════════════════
+
+class CTSEnvWrapper:
+ """Wraps CTS env. Computes latent online, returns augmented obs (77-dim)."""
+
+ def __init__(self, env, teacher_ratio=0.75, latent_dim=32,
+ teacher_enc=None, teacher_params=None,
+ student_enc=None, student_params=None):
+ self.__env = env
+ self._teacher_ratio = teacher_ratio
+ self._latent_dim = latent_dim
+ self._num_envs = env.num_envs
+ self._t_enc = teacher_enc
+ self._t_params = teacher_params
+ self._s_enc = student_enc
+ self._s_params = student_params
+ from gymnasium import spaces
+ self._obs_space = spaces.Box(-np.inf, np.inf,
+ (45 + latent_dim,), dtype=np.float32)
+
+ @property
+ def num_envs(self): return self._num_envs
+ @property
+ def num_teacher(self): return int(self._num_envs * self._teacher_ratio)
+ @property
+ def teacher_mask(self):
+ m = np.zeros(self._num_envs, dtype=np.bool)
+ m[:self.num_teacher] = True
+ return m
+ @property
+ def student_mask(self): return ~self.teacher_mask
+ @property
+ def privileged_obs(self):
+ return self.__env._state.info.get("privileged_obs",
+ np.zeros((self._num_envs, 76), dtype=np.float32))
+ @property
+ def history_obs(self): return self.__env.history_obs
+ @property
+ def observation_space(self): return self._obs_space
+ @property
+ def action_space(self): return self.__env.action_space
+ @property
+ def device(self): return None
+ @property
+ def num_agents(self): return 1
+ @property
+ def state_preprocessor(self): return None
+ @property
+ def value_preprocessor(self): return None
+ @property
+ def world_info_idx(self): return None
+ @property
+ def random_timestep(self): return False
+ @property
+ def running(self): return True
+ @property
+ def _env(self): return self.__env
+
+ def _augment(self, states):
+ n = states.shape[0]
+ latent = np.zeros((n, self._latent_dim), dtype=np.float32)
+ t_mask = self.teacher_mask
+ if t_mask.any() and self._t_params is not None:
+ priv = jnp.array(self.privileged_obs[t_mask])
+ latent[t_mask] = np.array(self._t_enc.apply(self._t_params, priv))
+ s_mask = self.student_mask
+ if s_mask.any() and self._s_params is not None:
+ hist = jnp.array(self.history_obs[s_mask])
+ latent[s_mask] = np.array(self._s_enc.apply(self._s_params, hist))
+ return np.concatenate([states, latent], axis=-1)
+
+ def reset(self):
+ state = self.__env.init_state()
+ return self._augment(state.obs), state.info
+
+ def step(self, actions):
+ state = self.__env.step(np.array(actions))
+ return (self._augment(state.obs), state.reward.reshape(-1, 1),
+ state.terminated.reshape(-1, 1),
+ state.truncated.reshape(-1, 1), state.info)
+
+ def render(self, *a, **kw): pass
+ def close(self): pass
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# CTS Agent — adds distillation step after PPO update
+# ═══════════════════════════════════════════════════════════════════════════
+
+class CTSAgent(BasePPO):
+ """PPO agent with post-update distillation on student encoder."""
+
+ def __init__(self, models, memory, cfg, observation_space, action_space, device,
+ teacher_enc=None, student_enc=None, cts_wrapper=None,
+ teacher_params=None, student_params=None,
+ distill_coef=0.1, student_lr=1e-3):
+ super().__init__(models=models, memory=memory, cfg=cfg,
+ observation_space=observation_space,
+ action_space=action_space, device=device)
+ self._t_enc = teacher_enc
+ self._s_enc = student_enc
+ self._t_params = teacher_params
+ self._s_params = student_params
+ self._wrapper = cts_wrapper
+ self._distill_coef = distill_coef
+ self._latent_dim = getattr(teacher_enc, 'latent_dim', 32) if teacher_enc else 32
+
+ # Separate optimizer for student encoder
+ if student_params is not None:
+ self._s_optimizer = optax.adam(student_lr)
+ self._s_opt_state = self._s_optimizer.init(student_params)
+ else:
+ self._s_optimizer = None
+
+ def update(self):
+ """PPO update + distillation phase."""
+ result = super().update()
+
+ # Phase 2: Distillation — MSE(teacher_latent, student_latent) on student envs
+ if self._s_params is not None and self._s_optimizer is not None:
+ self._distillation_step()
+
+ return result
+
+ def _distillation_step(self):
+ """Train student encoder to match teacher latent.
+
+ go2_rl_gym Phase 2: only on student env samples.
+ Loss = || teacher_enc(privileged_obs) - student_enc(history_obs) ||²
+ """
+ wrapper = self._wrapper
+ if wrapper is None:
+ return
+
+ s_mask = wrapper.student_mask
+ if not s_mask.any():
+ return
+
+ # Get current student env observations
+ priv = jnp.array(wrapper.privileged_obs[s_mask]) # [S, 76]
+ hist = jnp.array(wrapper.history_obs[s_mask]) # [S, 225]
+
+ # L2 normalize both before computing MSE (match go2_rl_gym)
+ def l2_norm(x):
+ return x / (jnp.linalg.norm(x, axis=-1, keepdims=True) + 1e-8)
+
+ def distill_loss(s_params):
+ t_latent = l2_norm(jax.lax.stop_gradient(
+ self._t_enc.apply(self._t_params, priv)))
+ s_latent = l2_norm(self._s_enc.apply(s_params, hist))
+ return jnp.mean(jnp.square(t_latent - s_latent))
+
+ loss, grads = jax.value_and_grad(distill_loss)(self._s_params)
+ updates, self._s_opt_state = self._s_optimizer.update(
+ grads, self._s_opt_state, self._s_params)
+ self._s_params = optax.apply_updates(self._s_params, updates)
+
+ # Track
+ self.tracking_data["Loss / Distillation"].append(loss)
+
+
+# ═══════════════════════════════════════════════════════════════════════════
+# Trainer
+# ═══════════════════════════════════════════════════════════════════════════
+
+class CTSTrainer:
+ def __init__(self, env_name="go1-cts-flat-walk-go2style",
+ sim_backend=None, enable_render=False, cfg_override=None):
+ self._env_name = env_name
+ self._sim_backend = sim_backend
+ rlcfg = registry.default_rl_cfg(env_name, "skrl", backend="jax")
+ if cfg_override:
+ rlcfg = utils.cfg_override(rlcfg, cfg_override)
+ self._rlcfg = rlcfg
+
+ def train(self):
+ cfg = self._rlcfg
+ set_seed(cfg.runner.seed)
+
+ env = env_registry.make(self._env_name, num_envs=cfg.num_envs)
+ env_cfg = env._cfg
+ latent_dim = getattr(env_cfg, 'latent_dim', 32)
+ priv_dim = getattr(env_cfg, 'num_privileged_obs', 76)
+ hist_len = getattr(env_cfg, 'history_length', 5)
+ student_dim = getattr(env_cfg, 'num_student_obs', 45)
+ teacher_ratio = getattr(env_cfg, 'teacher_env_ratio', 0.75)
+
+ # Init Flax encoders
+ rng = jax.random.PRNGKey(cfg.runner.seed)
+ rng, t_key, s_key = jax.random.split(rng, 3)
+ teacher_enc = TeacherEncoder(latent_dim=latent_dim)
+ student_enc = StudentEncoder(latent_dim=latent_dim)
+ t_params = teacher_enc.init(t_key, jnp.ones((1, priv_dim)))
+ s_params = student_enc.init(s_key, jnp.ones((1, hist_len * student_dim)))
+
+ # CTS wrapper
+ wrapper = CTSEnvWrapper(env, teacher_ratio=teacher_ratio, latent_dim=latent_dim,
+ teacher_enc=teacher_enc, teacher_params=t_params,
+ student_enc=student_enc, student_params=s_params)
+ env.init_state()
+
+ from gymnasium import spaces
+ aug_obs = spaces.Box(-np.inf, np.inf, (student_dim + latent_dim,), dtype=np.float32)
+ wrapper._obs_space = aug_obs
+
+ nT = int(cfg.num_envs * teacher_ratio)
+ nS = cfg.num_envs - nT
+ print(f"[CTS v2] {cfg.num_envs} envs ({nT}T/{nS}S), latent={latent_dim}, "
+ f"priv={priv_dim}, hist={hist_len}×{student_dim}, distill_coef=0.1")
+
+ # Standard model creation
+ from motrix_rl.skrl.jax.train.ppo import Trainer as StdTrainer, _add_runtime_config
+ std = StdTrainer(self._env_name, sim_backend=None, enable_render=False)
+ std._rlcfg = cfg
+ models = std._make_model(wrapper, cfg)
+
+ ppo_cfg = cfg.runner.agent.to_dict()
+ _add_runtime_config(ppo_cfg, wrapper,
+ log_dir=get_log_dir(self._env_name, rllib="skrl", agent_name="CTS_PPO"))
+
+ mem_size = cfg.runner.memory.memory_size
+ if mem_size == -1:
+ mem_size = ppo_cfg.get("rollouts", 24)
+ memory = RandomMemory(memory_size=mem_size, num_envs=cfg.num_envs, device=None)
+
+ agent = CTSAgent(
+ models=models, memory=memory, cfg=ppo_cfg,
+ observation_space=aug_obs, action_space=env.action_space,
+ device=None,
+ teacher_enc=teacher_enc, student_enc=student_enc,
+ cts_wrapper=wrapper,
+ teacher_params=t_params, student_params=s_params,
+ distill_coef=0.1, student_lr=1e-3,
+ )
+
+ trainer_cfg = {"timesteps": cfg.runner.trainer.timesteps, "headless": True}
+ SequentialTrainer(cfg=trainer_cfg, env=wrapper, agents=agent).train()
+
+ # Save encoder params for deployment
+ import pickle as _pk
+ _log = get_log_dir(self._env_name, rllib="skrl", agent_name="CTS_PPO")
+ _pk.dump(t_params, open(f"{_log}/teacher_params.pkl", "wb"))
+ _pk.dump(s_params, open(f"{_log}/student_params.pkl", "wb"))
+ print(f"[CTS] Encoder weights saved → {_log}")
diff --git a/motrix_rl/src/motrix_rl/skrl/jax/train/dreamwaq_ppo.py b/motrix_rl/src/motrix_rl/skrl/jax/train/dreamwaq_ppo.py
new file mode 100644
index 0000000..092ad38
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/skrl/jax/train/dreamwaq_ppo.py
@@ -0,0 +1,540 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+"""DreamWaQ PPO — aligned with Manaro-Alpha/DreamWaQ.
+
+CENet (VAE): history(225) → [128,64] → latent(16) + vel_est(3) = code(19)
+Decoder: code(19) → [64,128] → next_obs(45)
+Actor: code(19) + obs(45) = 64 → [512,256,128] → action(12)
+Critic: privileged_obs(235) → [512,256,128] → value(1)
+
+VAE Loss: reconstruction_MSE + velocity_MSE + beta * KL
+"""
+
+import flax.linen as nn
+import jax
+import jax.numpy as jnp
+import numpy as np
+import optax
+import pickle, os
+
+from skrl.agents.jax.ppo import PPO as BasePPO
+from skrl.memories.jax import RandomMemory
+from skrl.models.jax import DeterministicMixin, GaussianMixin, Model
+from skrl.trainers.jax import SequentialTrainer
+from skrl.utils import set_seed
+
+from motrix_envs import registry as env_registry
+from motrix_envs.locomotion.go1.dreamwaq import DreamWaQTask
+from motrix_rl import registry, utils
+from motrix_rl.skrl import get_log_dir
+
+# Ensure RL config registered
+import motrix_rl.tasks.go1_dreamwaq # noqa: F401
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# CENet (VAE) — exact Manaro-Alpha architecture
+# ═══════════════════════════════════════════════════════════════════════
+
+class CENet(nn.Module):
+ """VAE for terrain imagination + velocity estimation.
+
+ Encoder: 225 → Linear(128) → ELU → Linear(64) → ELU
+ 4 heads on 64-dim bottleneck:
+ - latent_mu, latent_logvar: 64 → 16 each
+ - vel_mu, vel_logvar: 64 → 3 each
+ Decoder: (latent+vel=19) → Linear(64) → ELU → Linear(128) → ELU → Linear(45)
+ """
+ num_obs: int = 45
+ num_history: int = 5
+ num_latent: int = 16 # latent z dim
+ cenet_out_dim: int = 19 # code = vel(3) + latent(16)
+
+ def setup(self):
+ # Orthogonal weight init matching upstream torch.nn.init.orthogonal_
+ kernel_init = nn.initializers.orthogonal(jnp.sqrt(2))
+
+ # Encoder (matches Manaro-Alpha: 225 → 128 → 64)
+ self.enc_fc1 = nn.Dense(128, kernel_init=kernel_init)
+ self.enc_fc2 = nn.Dense(64, kernel_init=kernel_init)
+
+ # 4 output heads
+ self.latent_mu = nn.Dense(self.num_latent, kernel_init=kernel_init)
+ self.latent_logvar = nn.Dense(self.num_latent, kernel_init=kernel_init)
+ self.vel_mu = nn.Dense(3, kernel_init=kernel_init)
+ self.vel_logvar = nn.Dense(3, kernel_init=kernel_init)
+
+ # Decoder (matches Manaro-Alpha: 19 → 64 → 128 → 45)
+ self.dec_fc1 = nn.Dense(64, kernel_init=kernel_init)
+ self.dec_fc2 = nn.Dense(128, kernel_init=kernel_init)
+ self.dec_out = nn.Dense(self.num_obs, kernel_init=kernel_init)
+
+ def encode(self, obs_history):
+ """obs_history: (B, T, 45) or (B, T*45) → (latent_mu, latent_logvar, vel_mu, vel_logvar)."""
+ x = obs_history.reshape(obs_history.shape[0], -1)
+ x = nn.elu(self.enc_fc1(x))
+ x = nn.elu(self.enc_fc2(x))
+ return (self.latent_mu(x), self.latent_logvar(x),
+ self.vel_mu(x), self.vel_logvar(x))
+
+ def reparameterize(self, mu, logvar, rng):
+ std = jnp.exp(0.5 * logvar)
+ return mu + std * jax.random.normal(rng, mu.shape)
+
+ def decode(self, z, vel):
+ x = jnp.concatenate([z, vel], axis=-1)
+ x = nn.elu(self.dec_fc1(x))
+ x = nn.elu(self.dec_fc2(x))
+ return self.dec_out(x)
+
+ def __call__(self, obs_history, rng):
+ lm, ll, vm, vl = self.encode(obs_history)
+ rng_z, rng_v = jax.random.split(rng)
+ z = self.reparameterize(lm, ll, rng_z)
+ vel = self.reparameterize(vm, vl, rng_v)
+ return {
+ 'z': z, 'vel_est': vel,
+ 'latent_mu': lm, 'latent_logvar': ll,
+ 'vel_mu': vm, 'vel_logvar': vl,
+ 'obs_pred': self.decode(z, vel),
+ }
+
+ def sample(self, obs_history, rng):
+ """Sample (z, vel) for rollout (adds noise)."""
+ lm, ll, vm, vl = self.encode(obs_history)
+ rng_z, rng_v = jax.random.split(rng)
+ z = self.reparameterize(lm, ll, rng_z)
+ vel = self.reparameterize(vm, vl, rng_v)
+ return z, vel
+
+ def inference(self, obs_history):
+ """Mean (z, vel) for deployment (no noise)."""
+ lm, ll, vm, vl = self.encode(obs_history)
+ return lm, vm
+
+ def loss_fn(self, obs_history, current_obs, base_vel, rng, kl_weight=1.0):
+ """VAE loss: reconstructs current obs (last history frame), NOT next obs.
+
+ Upstream: decode_target = obs_batch (the same timestep's observation,
+ not a future prediction). This is an autoencoder over current observation.
+ """
+ lm, ll, vm, vl = self.encode(obs_history)
+ rng_z, rng_v = jax.random.split(rng)
+ z = self.reparameterize(lm, ll, rng_z)
+ vel = self.reparameterize(vm, vl, rng_v)
+ pred = self.decode(z, vel)
+
+ recons = jnp.mean(jnp.square(pred - current_obs), axis=-1)
+ vel_loss = jnp.mean(jnp.square(vel - base_vel), axis=-1)
+ kld = -0.5 * jnp.sum(1 + ll - jnp.square(lm) - jnp.exp(ll), axis=-1)
+
+ return {
+ 'loss': jnp.mean(recons + vel_loss + kl_weight * kld),
+ 'recons_loss': jnp.mean(recons),
+ 'vel_loss': jnp.mean(vel_loss),
+ 'kld_loss': jnp.mean(kld),
+ }
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# Environment Wrapper
+# ═══════════════════════════════════════════════════════════════════════
+
+class DreamWaQWrapper:
+ """Wraps DreamWaQTask, computes CENet code online.
+
+ Augmented observation layout (254-dim):
+ [code_vel(3) | code_latent(16) | obs(45) | base_vel(3) | heights(187)]
+
+ Actor uses: [0:64] = code(19) + obs(45)
+ Critic uses: [19:254] = obs(45) + base_vel(3) + heights(187)
+ """
+
+ def __init__(self, env, cenet, cenet_params, rng=None, vae_buf_size=24):
+ self._env = env
+ self._cenet = cenet
+ self._cenet_params = cenet_params
+ self._num_envs = env.num_envs
+ self._rng = rng if rng is not None else jax.random.PRNGKey(0)
+ self._code_dim = 19 # vel(3) + latent(16)
+ self._aug_dim = self._code_dim + 45 + 3 + 187 # 254
+
+ from gymnasium import spaces
+ self._obs_space = spaces.Box(-np.inf, np.inf, (self._aug_dim,), dtype=np.float32)
+
+ # VAE training buffer
+ self._vae_buf = []
+ self._vae_buf_size = vae_buf_size
+
+ @property
+ def num_envs(self): return self._num_envs
+ @property
+ def observation_space(self): return self._obs_space
+ @property
+ def action_space(self): return self._env.action_space
+ @property
+ def device(self): return None
+ @property
+ def num_agents(self): return 1
+ @property
+ def state_preprocessor(self): return None
+ @property
+ def value_preprocessor(self): return None
+ @property
+ def world_info_idx(self): return None
+ @property
+ def random_timestep(self): return False
+ @property
+ def running(self): return True
+
+ @property
+ def privileged_obs(self):
+ return self._env._state.info.get("privileged_obs",
+ np.zeros((self._num_envs, 235), dtype=np.float32))
+
+ @property
+ def obs_history(self):
+ return self._env._state.info.get("obs_history",
+ np.zeros((self._num_envs, 5, 45), dtype=np.float32))
+
+ @property
+ def base_vel(self):
+ return self._env._state.info.get("base_vel",
+ np.zeros((self._num_envs, 3), dtype=np.float32))
+
+ @property
+ def vae_training_data(self):
+ return self._vae_buf
+
+ def _augment(self, obs):
+ """Compute CENet code and build augmented obs."""
+ hist = jnp.array(self.obs_history)
+ self._rng, rng_key = jax.random.split(self._rng)
+ z, vel = self._cenet.apply(self._cenet_params, hist, rng_key, method=self._cenet.sample)
+ code = np.concatenate([np.array(vel), np.array(z)], axis=-1) # (N, 19)
+
+ priv = self.privileged_obs
+ heights = priv[:, 45+3:] # (N, 187)
+ base_vel_norm = self.base_vel * np.array([2.0, 2.0, 1.0], dtype=np.float32)
+
+ return np.concatenate([code, obs, base_vel_norm, heights], axis=-1)
+
+ def reset(self):
+ state = self._env.init_state()
+ self._vae_buf = []
+ return self._augment(state.obs), state.info
+
+ def step(self, actions):
+ old_hist = self.obs_history.copy()
+ state = self._env.step(np.array(actions, dtype=np.float32))
+
+ # Store for VAE training: history → current_obs (last frame of history)
+ # Upstream: decode target = obs_batch (current observation, not next)
+ self._vae_buf.append((old_hist, old_hist[:, -1, :].copy(), self.base_vel.copy()))
+ if len(self._vae_buf) > self._vae_buf_size:
+ self._vae_buf.pop(0)
+
+ return (self._augment(state.obs),
+ state.reward.reshape(-1, 1).astype(np.float32),
+ state.terminated.reshape(-1, 1).astype(np.float32),
+ state.truncated.reshape(-1, 1).astype(np.float32),
+ state.info)
+
+ def render(self, *a, **kw): pass
+ def close(self): pass
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# DreamWaQ Agent — PPO + CENet training
+# ═══════════════════════════════════════════════════════════════════════
+
+class DreamWaQAgent(BasePPO):
+ """PPO agent with asymmetric critic + CENet (VAE) training."""
+
+ def __init__(self, models, memory, cfg, observation_space, action_space, device,
+ cenet=None, cenet_params=None, wrapper=None,
+ kl_weight=1.0, vae_lr=1e-3, save_dir=None):
+ super().__init__(models=models, memory=memory, cfg=cfg,
+ observation_space=observation_space,
+ action_space=action_space, device=device)
+ self._cenet = cenet
+ self._cenet_params = cenet_params
+ self._wrapper = wrapper
+ self._kl_weight = kl_weight
+ self._save_dir = save_dir
+ self._vae_count = 0
+ self._rng = jax.random.PRNGKey(42)
+
+ if cenet_params is not None:
+ self._vae_opt = optax.adam(vae_lr)
+ self._vae_opt_state = self._vae_opt.init(cenet_params)
+ else:
+ self._vae_opt = None
+
+ # Override SKRL's internal _update (NOT 'update')
+ def _update(self, timestep: int, timesteps: int):
+ super()._update(timestep, timesteps)
+
+ # Track detailed reward terms
+ wrapper = self._wrapper
+ if wrapper is not None:
+ try:
+ terms = wrapper._env._state.info.get("reward_terms", {})
+ for name, val in terms.items():
+ self.tracking_data[f"Reward / {name}"].append(float(val))
+ except Exception:
+ pass
+
+ # CENet training after each PPO update
+ if self._cenet_params is not None and self._vae_opt is not None:
+ self._vae_step()
+
+ def _vae_step(self):
+ """Train CENet: reconstruct current obs (last history frame) + velocity supervision + KL."""
+ wrapper = self._wrapper
+ if wrapper is None or len(wrapper.vae_training_data) == 0:
+ return
+
+ hist_np, current_obs_np, base_vel_np = wrapper.vae_training_data[-1]
+ hist = jnp.array(hist_np)
+ current_obs = jnp.array(current_obs_np)
+ base_vel = jnp.array(base_vel_np)
+
+ n = wrapper.num_envs
+ valid = jnp.linalg.norm(hist.reshape(n, -1), axis=-1) > 1e-6
+ if not valid.any():
+ return
+
+ self._rng, vae_rng = jax.random.split(self._rng)
+
+ def loss_fn(p):
+ return self._cenet.apply(
+ p, hist, current_obs, base_vel, vae_rng, self._kl_weight,
+ method=self._cenet.loss_fn)['loss']
+
+ loss, grads = jax.value_and_grad(loss_fn)(self._cenet_params)
+ updates, self._vae_opt_state = self._vae_opt.update(
+ grads, self._vae_opt_state, self._cenet_params)
+ self._cenet_params = optax.apply_updates(self._cenet_params, updates)
+
+ # Sync wrapper
+ wrapper._cenet_params = self._cenet_params
+
+ self._vae_count += 1
+ if self._save_dir and self._vae_count % 200 == 0:
+ pickle.dump(self._cenet_params,
+ open(f"{self._save_dir}/vae_{self._vae_count}.pkl", "wb"))
+
+ self.tracking_data["Loss / VAE"].append(loss)
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# Trainer
+# ═══════════════════════════════════════════════════════════════════════
+
+class DreamWaQTrainer:
+ def __init__(self, env_name="go1-dreamwaq-walk", cfg_override=None):
+ self._env_name = env_name
+ rlcfg = registry.default_rl_cfg(env_name, "skrl", backend="jax")
+ if cfg_override:
+ rlcfg = utils.cfg_override(rlcfg, cfg_override)
+ self._rlcfg = rlcfg
+
+ def train(self):
+ cfg = self._rlcfg
+ set_seed(cfg.runner.seed)
+
+ env = env_registry.make(self._env_name, num_envs=cfg.num_envs)
+ env_cfg = env._cfg
+
+ # Init CENet
+ rng = jax.random.PRNGKey(cfg.runner.seed)
+ rng, crng = jax.random.split(rng)
+ cenet = CENet(num_obs=45, num_history=env_cfg.num_history,
+ num_latent=env_cfg.num_latent,
+ cenet_out_dim=env_cfg.cenet_out_dim)
+ cenet_params = cenet.init(crng, jnp.ones((1, env_cfg.num_history, 45)), jax.random.PRNGKey(0))
+
+ # Wrapper
+ wrapper = DreamWaQWrapper(env, cenet, cenet_params,
+ rng=jax.random.PRNGKey(cfg.runner.seed + 1))
+ env.init_state()
+
+ print(f"[DreamWaQ] {cfg.num_envs} envs, CENet[128,64], latent=16+3=19, "
+ f"aug_dim={wrapper.observation_space.shape[0]}")
+
+ # Build models
+ from motrix_rl.skrl.jax.train.ppo import Trainer as StdTrainer, _add_runtime_config
+
+ policy_cfg = cfg.runner.models.policy
+ value_cfg = cfg.runner.models.value
+
+ class Policy(GaussianMixin, Model):
+ def __init__(self, observation_space, action_space, device=None, **kwargs):
+ Model.__init__(self, observation_space, action_space, device, **kwargs)
+ GaussianMixin.__init__(self, policy_cfg.clip_actions,
+ policy_cfg.clip_log_std,
+ policy_cfg.min_log_std, policy_cfg.max_log_std,
+ policy_cfg.reduction)
+
+ @nn.compact
+ def __call__(self, inputs, role):
+ kernel_init = nn.initializers.orthogonal(jnp.sqrt(2))
+ x = inputs["states"]
+ # Actor: code(19) + obs(45) = first 64 dims
+ x = x[:, :64]
+ for d in policy_cfg.hiddens:
+ x = nn.elu(nn.Dense(d, kernel_init=kernel_init)(x))
+ x = nn.Dense(self.num_actions, kernel_init=kernel_init)(x)
+ log_std = self.param("log_std",
+ lambda _: jnp.full(self.num_actions, float(policy_cfg.initial_log_std)))
+ return x, log_std, {}
+
+ class Value(DeterministicMixin, Model):
+ def __init__(self, observation_space, action_space, device=None, **kwargs):
+ Model.__init__(self, observation_space, action_space, device, **kwargs)
+ DeterministicMixin.__init__(self, value_cfg.clip_actions)
+
+ @nn.compact
+ def __call__(self, inputs, role):
+ kernel_init = nn.initializers.orthogonal(jnp.sqrt(2))
+ x = inputs["states"]
+ # Critic: obs(45) + base_vel(3) + heights(187) = 235
+ # Layout: [code(19) | obs(45) | base_vel(3) | heights(187)]
+ x_c = jnp.concatenate([x[:, 19:64], x[:, 64:254]], axis=-1)
+ for d in value_cfg.hiddens:
+ x_c = nn.elu(nn.Dense(d, kernel_init=kernel_init)(x_c))
+ return nn.Dense(1, kernel_init=kernel_init)(x_c), {}
+
+ models = {
+ "policy": Policy(wrapper.observation_space, wrapper.action_space, None),
+ "value": Value(wrapper.observation_space, wrapper.action_space, None),
+ }
+ for role, m in models.items():
+ m.init_state_dict(role)
+
+ # Agent
+ ppo_cfg = cfg.runner.agent.to_dict()
+ log_dir = get_log_dir(self._env_name, rllib="skrl", agent_name="DreamWaQ_PPO")
+ _add_runtime_config(ppo_cfg, wrapper, log_dir=log_dir)
+
+ mem_size = cfg.runner.memory.memory_size
+ if mem_size == -1:
+ mem_size = ppo_cfg.get("rollouts", 24)
+ memory = RandomMemory(memory_size=mem_size, num_envs=cfg.num_envs, device=None)
+
+ agent = DreamWaQAgent(
+ models=models, memory=memory, cfg=ppo_cfg,
+ observation_space=wrapper.observation_space,
+ action_space=wrapper.action_space, device=None,
+ cenet=cenet, cenet_params=cenet_params, wrapper=wrapper,
+ kl_weight=1.0, vae_lr=1e-3, save_dir=log_dir,
+ )
+
+ trainer_cfg = {"timesteps": cfg.runner.trainer.timesteps, "headless": True}
+ SequentialTrainer(cfg=trainer_cfg, env=wrapper, agents=agent).train()
+
+ # Save final VAE params
+ pickle.dump(agent._cenet_params, open(f"{log_dir}/cenet_params.pkl", "wb"))
+ print(f"[DreamWaQ] CENet weights saved → {log_dir}")
+
+ def play(self, policy_path=None):
+ """Render env with trained policy (uses wrapper + CENet)."""
+ from motrix_rl.skrl.jax.wrap_np import SkrlNpWrapper as RenderWrapper
+ cfg = self._rlcfg
+ raw_env = env_registry.make(self._env_name, num_envs=min(64, cfg.play_num_envs))
+ env = RenderWrapper(raw_env) # adds render() method
+ env_cfg = raw_env._cfg
+
+ import msgpack
+ rng = jax.random.PRNGKey(cfg.runner.seed)
+ rng, crng = jax.random.split(rng)
+ cenet = CENet()
+ log_dir = get_log_dir(self._env_name, rllib="skrl", agent_name="DreamWaQ_PPO")
+
+ # Load CENet params
+ if policy_path and os.path.exists(policy_path):
+ # Load full agent checkpoint and extract
+ with open(policy_path, 'rb') as f:
+ ckpt = pickle.load(f)
+ # Load policy params
+ policy_raw = msgpack.unpackb(ckpt['policy'])
+ policy_params = {}
+ for name, val in policy_raw['params'].items():
+ if isinstance(val, dict):
+ policy_params[name] = {
+ k: _decode_msgpack_array(v) for k, v in val.items()
+ }
+ else:
+ policy_params[name] = _decode_msgpack_array(val)
+
+ # Try loading VAE from base skrl dir
+ log_dir = get_log_dir(self._env_name, rllib="skrl", agent_name="DreamWaQ_PPO")
+ vae_path = os.path.join(log_dir, "cenet_params.pkl")
+ if not os.path.exists(vae_path):
+ # Find latest vae_N.pkl
+ vae_files = sorted([f for f in os.listdir(log_dir) if f.startswith("vae_")])
+ if vae_files:
+ vae_path = os.path.join(log_dir, vae_files[-1])
+ if os.path.exists(vae_path):
+ with open(vae_path, 'rb') as f:
+ cenet_params = pickle.load(f)
+ print(f"[Play] VAE loaded: {vae_path}")
+ else:
+ print("[Play] WARNING: VAE params not found, using random init")
+ cenet_params = cenet.init(crng, jnp.ones((1, 5, 45)), jax.random.PRNGKey(0))
+ else:
+ print("[Play] No policy, random actions")
+ policy_params = None
+ cenet_params = cenet.init(crng, jnp.ones((1, 5, 45)), jax.random.PRNGKey(0))
+
+ wrapper = DreamWaQWrapper(raw_env, cenet, cenet_params,
+ rng=jax.random.PRNGKey(cfg.runner.seed + 1))
+ raw_env.init_state()
+
+ import time
+ n = raw_env._num_envs
+ print(f"[Play] {n} envs, press Ctrl+C to stop")
+ try:
+ while True:
+ if policy_params is not None:
+ aug_obs = wrapper._augment(raw_env._state.obs)
+ actions = np.array(_policy_forward(jnp.array(aug_obs), policy_params))
+ else:
+ actions = np.zeros((n, 12), dtype=np.float32)
+
+ wrapper.step(actions)
+ env.render()
+ time.sleep(0.01)
+ except KeyboardInterrupt:
+ env.close()
+ print("[Play] Stopped")
+
+
+def _decode_msgpack_array(ext) -> np.ndarray | None:
+ """Decode flax-serialized msgpack ExtType."""
+ import msgpack
+ 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])
+ return np.frombuffer(parts[2], dtype=np.dtype(parts[1])).reshape(shape)
+
+
+def _policy_forward(x, params):
+ """Forward pass through policy (first 64 dims only, apply elu)."""
+ import jax.numpy as jnp
+ import jax
+ w0, b0 = jnp.array(params['Dense_0']['kernel']), jnp.array(params['Dense_0']['bias'])
+ w1, b1 = jnp.array(params['Dense_1']['kernel']), jnp.array(params['Dense_1']['bias'])
+ w2, b2 = jnp.array(params['Dense_2']['kernel']), jnp.array(params['Dense_2']['bias'])
+ w3, b3 = jnp.array(params['Dense_3']['kernel']), jnp.array(params['Dense_3']['bias'])
+ x = x[:, :64] # code(19) + obs(45)
+ x = jax.nn.elu(x @ w0 + b0)
+ x = jax.nn.elu(x @ w1 + b1)
+ x = jax.nn.elu(x @ w2 + b2)
+ return x @ w3 + b3
diff --git a/motrix_rl/src/motrix_rl/skrl/jax/train/ppo.py b/motrix_rl/src/motrix_rl/skrl/jax/train/ppo.py
index 6f421d6..c2e0149 100644
--- a/motrix_rl/src/motrix_rl/skrl/jax/train/ppo.py
+++ b/motrix_rl/src/motrix_rl/skrl/jax/train/ppo.py
@@ -185,6 +185,7 @@ class Trainer:
sim_backend: str = None,
enable_render: bool = False,
cfg_override: dict = None,
+ env_cfg_override: dict = None,
) -> None:
rlcfg = registry.default_rl_cfg(env_name, "skrl", backend="jax")
if cfg_override is not None:
@@ -193,13 +194,18 @@ class Trainer:
self._env_name = env_name
self._sim_backend = sim_backend
self._enable_render = enable_render
+ self._env_cfg_override = env_cfg_override
- def train(self) -> None:
+ def train(self, checkpoint: str = None) -> None:
"""
Start training the agent.
+
+ Args:
+ checkpoint: Optional path to a checkpoint to resume from.
"""
rlcfg = self._rlcfg
- env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.num_envs)
+ env = env_registry.make(self._env_name, sim_backend=self._sim_backend,
+ num_envs=rlcfg.num_envs, env_cfg_override=self._env_cfg_override)
set_seed(rlcfg.runner.seed)
skrl_env = wrap_env(env, self._enable_render)
@@ -209,6 +215,12 @@ class Trainer:
# Add runtime-specific configuration
_add_runtime_config(ppo_cfg, skrl_env, log_dir=get_log_dir(self._env_name, rllib="skrl", agent_name="PPO"))
agent = self._make_agent(models, skrl_env, ppo_cfg, rlcfg.runner.memory)
+
+ if checkpoint:
+ agent.load(checkpoint)
+ logger = __import__("logging").getLogger(__name__)
+ logger.info(f"Resumed from checkpoint: {checkpoint}")
+
cfg_trainer = {
"timesteps": rlcfg.runner.trainer.timesteps,
"headless": not self._enable_render,
@@ -220,7 +232,17 @@ class Trainer:
import time
rlcfg = self._rlcfg
- env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.play_num_envs)
+ # Enable play_mode + merge constructor-level overrides (e.g. force_phase)
+ play_override = {"play_mode": True}
+ if self._env_cfg_override:
+ play_override.update(self._env_cfg_override)
+ try:
+ env = env_registry.make(
+ self._env_name, sim_backend=self._sim_backend,
+ num_envs=rlcfg.play_num_envs, env_cfg_override=play_override,
+ )
+ except ValueError:
+ env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.play_num_envs)
set_seed(rlcfg.runner.seed)
env = wrap_env(env, self._enable_render)
diff --git a/motrix_rl/src/motrix_rl/skrl/torch/train/ppo.py b/motrix_rl/src/motrix_rl/skrl/torch/train/ppo.py
index d4cb614..03b8be1 100644
--- a/motrix_rl/src/motrix_rl/skrl/torch/train/ppo.py
+++ b/motrix_rl/src/motrix_rl/skrl/torch/train/ppo.py
@@ -185,6 +185,7 @@ class Trainer:
sim_backend: str = None,
enable_render: bool = False,
cfg_override: dict = None,
+ env_cfg_override: dict = None,
) -> None:
rlcfg = registry.default_rl_cfg(env_name, "skrl", backend="torch")
if cfg_override is not None:
@@ -193,13 +194,18 @@ class Trainer:
self._env_name = env_name
self._sim_backend = sim_backend
self._enable_render = enable_render
+ self._env_cfg_override = env_cfg_override
- def train(self) -> None:
+ def train(self, checkpoint: str = None) -> None:
"""
Start training the agent.
+
+ Args:
+ checkpoint: Optional path to a checkpoint to resume from.
"""
rlcfg = self._rlcfg
- env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.num_envs)
+ env = env_registry.make(self._env_name, sim_backend=self._sim_backend,
+ num_envs=rlcfg.num_envs, env_cfg_override=self._env_cfg_override)
set_seed(rlcfg.runner.seed)
skrl_env = wrap_env(env, self._enable_render)
models = self._make_model(skrl_env, rlcfg)
@@ -208,6 +214,12 @@ class Trainer:
# Add runtime-specific configuration
_add_runtime_config(ppo_cfg, skrl_env, log_dir=get_log_dir(self._env_name, rllib="skrl", agent_name="PPO"))
agent = self._make_agent(models, skrl_env, ppo_cfg, rlcfg.runner.memory)
+
+ if checkpoint:
+ agent.load(checkpoint)
+ logger = __import__("logging").getLogger(__name__)
+ logger.info(f"Resumed from checkpoint: {checkpoint}")
+
cfg_trainer = {
"timesteps": rlcfg.runner.trainer.timesteps,
"headless": not self._enable_render,
@@ -219,7 +231,17 @@ class Trainer:
import time
rlcfg = self._rlcfg
- env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.play_num_envs)
+ # Enable play_mode if the env config supports it (for multi-terrain random spawn)
+ play_override = {"play_mode": True}
+ if self._env_cfg_override:
+ play_override.update(self._env_cfg_override)
+ try:
+ env = env_registry.make(
+ self._env_name, sim_backend=self._sim_backend,
+ num_envs=rlcfg.play_num_envs, env_cfg_override=play_override,
+ )
+ except ValueError:
+ env = env_registry.make(self._env_name, sim_backend=self._sim_backend, num_envs=rlcfg.play_num_envs)
set_seed(rlcfg.runner.seed)
env = wrap_env(env, self._enable_render)
models = self._make_model(env, rlcfg)
diff --git a/motrix_rl/src/motrix_rl/tasks/go1.py b/motrix_rl/src/motrix_rl/tasks/go1.py
index 741bbae..99d2b1e 100644
--- a/motrix_rl/src/motrix_rl/tasks/go1.py
+++ b/motrix_rl/src/motrix_rl/tasks/go1.py
@@ -135,7 +135,7 @@ class rslrl:
runner.actor.class_name = (
"motrix_rl.rslrl.torch.models.cenet_actor:CENetActorModel")
runner.actor.hidden_dims = [512, 256, 128]
- runner.actor.init_noise_std = 1.0
+ runner.actor.init_noise_std = 0.5 # 降噪声先学站稳,再探索行走
# Critic:标准 MLPModel,输入 privileged_obs
runner.critic.class_name = "MLPModel"
diff --git a/motrix_rl/src/motrix_rl/tasks/go1_dreamwaq.py b/motrix_rl/src/motrix_rl/tasks/go1_dreamwaq.py
new file mode 100644
index 0000000..acc5711
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/tasks/go1_dreamwaq.py
@@ -0,0 +1,42 @@
+"""DreamWaQ training config — Manaro-Alpha aligned.
+
+Actor [512,256,128], Critic [512,256,128], CENet [128,64].
+lr=1e-3, vae_lr=1e-3, KL_weight=1.0, entropy=0.01.
+"""
+import motrix_envs.locomotion.go1.dreamwaq # noqa: F401
+from motrix_rl.registry import rlcfg
+from motrix_rl.skrl.config import SkrlCfg
+
+
+def _apply(cfg: SkrlCfg):
+ # PPO hyperparams (exact DreamWaQ match)
+ cfg.runner.agent.rollouts = 24 # num_steps_per_env
+ cfg.runner.agent.learning_epochs = 5 # num_learning_epochs
+ cfg.runner.agent.mini_batches = 4 # num_mini_batches
+ cfg.runner.agent.learning_rate = 1e-3
+ cfg.runner.agent.entropy_loss_scale = 0.01
+ cfg.runner.agent.ratio_clip = 0.2
+ cfg.runner.agent.discount_factor = 0.99
+ cfg.runner.agent.lam = 0.95
+
+ # Network (DreamWaQ: [512,256,128])
+ cfg.runner.models.policy.hiddens = [512, 256, 128]
+ cfg.runner.models.policy.hidden_activation = ["elu"]
+ # initial_log_std: upstream uses 0.0 (std=1) in IsaacGym, but in MotrixSim with
+ # action_scale=0.25 that gives ±0.25 rad jitter/step → robot falls immediately and
+ # never learns to stand. -1.0 (std=0.37, ±0.09 rad) lets it stay upright early and
+ # bootstrap walking. (MotrixSim<->IsaacGym reality-gap adjustment.)
+ cfg.runner.models.policy.initial_log_std = -1.0
+ cfg.runner.models.value.hiddens = [512, 256, 128]
+ cfg.runner.models.value.hidden_activation = ["elu"]
+
+ cfg.runner.memory.memory_size = 24
+ cfg.runner.agent.experiment.checkpoint_interval = 200
+ cfg.runner.agent.experiment.write_interval = 200
+
+
+@rlcfg("go1-dreamwaq-walk")
+class DreamWaQ(SkrlCfg):
+ def __init__(self, **kwargs):
+ super().__init__(**kwargs)
+ _apply(self)
diff --git a/motrix_rl/src/motrix_rl/tasks/go1_go2style.py b/motrix_rl/src/motrix_rl/tasks/go1_go2style.py
new file mode 100644
index 0000000..64ff766
--- /dev/null
+++ b/motrix_rl/src/motrix_rl/tasks/go1_go2style.py
@@ -0,0 +1,64 @@
+"""Training configs for Go1 go2style environments.
+
+Importing triggers registration of: flat, stairs, combined.
+"""
+from dataclasses import dataclass
+
+import motrix_envs.locomotion.go1.go2style_flat # noqa: F401
+import motrix_envs.locomotion.go1.go2style_stairs # noqa: F401
+import motrix_envs.locomotion.go1.go2style_combined # noqa: F401
+import motrix_envs.locomotion.go1.go2style_cts # noqa: F401
+import motrix_envs.locomotion.go1.orig_combined # noqa: F401
+
+from motrix_rl.registry import rlcfg
+from motrix_rl.skrl.config import SkrlCfg
+
+
+def _apply_go2style_ppo(cfg: SkrlCfg):
+ """Apply go2_rl_gym-style PPO params to config."""
+ cfg.runner.models.policy.hiddens = [512, 256, 128]
+ cfg.runner.models.value.hiddens = [512, 256, 128]
+ cfg.runner.agent.rollouts = 24
+ cfg.runner.agent.learning_epochs = 5
+ cfg.runner.agent.mini_batches = 4
+ cfg.runner.agent.learning_rate = 1e-3
+ cfg.runner.trainer.timesteps = 60000
+
+
+def _apply_orig_ppo(cfg: SkrlCfg):
+ """Match original Go1 stairs config: [512,256,128], 30k steps."""
+ cfg.runner.models.policy.hiddens = [512, 256, 128]
+ cfg.runner.models.value.hiddens = [512, 256, 128]
+ cfg.runner.agent.rollouts = 24
+ cfg.runner.agent.learning_epochs = 5
+ cfg.runner.agent.mini_batches = 4
+ cfg.runner.agent.learning_rate = 1e-3
+ cfg.runner.trainer.timesteps = 30000 # same as original stairs
+
+
+class skrl:
+ @rlcfg("go1-flat-terrain-walk-go2style")
+ @dataclass
+ class Flat(SkrlCfg):
+ def __post_init__(self): _apply_go2style_ppo(self)
+
+ @rlcfg("go1-stairs-terrain-walk-go2style")
+ @dataclass
+ class Stairs(SkrlCfg):
+ def __post_init__(self): _apply_go2style_ppo(self)
+
+ @rlcfg("go1-combined-terrain-walk-go2style")
+ @dataclass
+ class Combined(SkrlCfg):
+ def __post_init__(self): _apply_go2style_ppo(self)
+
+ @rlcfg("go1-cts-flat-walk-go2style")
+ @dataclass
+ class CTS(SkrlCfg):
+ def __post_init__(self): _apply_go2style_ppo(self)
+
+ @rlcfg("go1-orig-combined-walk")
+ @dataclass
+ class OrigCombined(SkrlCfg):
+ """Original Go1 (PD 80, 45-dim, [256,128,64]) on flat+stairs."""
+ def __post_init__(self): _apply_orig_ppo(self)
diff --git a/scripts/dreamwaq_sim2sim_mujoco.py b/scripts/dreamwaq_sim2sim_mujoco.py
new file mode 100644
index 0000000..a6bd190
--- /dev/null
+++ b/scripts/dreamwaq_sim2sim_mujoco.py
@@ -0,0 +1,264 @@
+#!/usr/bin/env python3
+"""DreamWaQ MuJoCo sim2sim — VAE encoder + Actor, 5-frame history buffer.
+
+Usage:
+ uv run scripts/dreamwaq_sim2sim_mujoco.py # flat
+ uv run scripts/dreamwaq_sim2sim_mujoco.py --terrain rough
+ uv run scripts/dreamwaq_sim2sim_mujoco.py --onnx path/to/policy.onnx
+
+Controls:
+ W/S: forward/back Q/E: left/right A/D: rotate
+ Space: stop R: reset Esc: quit
+"""
+import numpy as np
+import mujoco
+from mujoco import viewer
+import onnxruntime as ort
+import os, sys, threading, queue, argparse, time, signal
+
+g_exit_requested = False
+signal.signal(signal.SIGINT, lambda *a: globals().update(g_exit_requested=True))
+
+# ═══════════════════════════════════════════════════════════════════════
+_PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+XML_DIR = os.path.join(_PROJECT, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls")
+DEFAULT_ONNX = os.path.join(_PROJECT, "exports_go1_dreamwaq", "policy.onnx")
+
+# ── DreamWaQ params (matching training: PD 28/0.7, action_scale 0.25, ctrl_dt=0.02) ──
+NUM_OBS = 45
+NUM_ACTIONS = 12
+HISTORY_LEN = 5
+ACTION_SCALE = 0.25
+KP = 28.0
+KD = 0.7
+CLIP_ACTIONS = 23.7
+CLIP_OBS = 100.0
+MAX_VX, MAX_VY, MAX_WZ = 1.0, 1.0, 1.0
+
+# DreamWaQ default joint angles — MUST match MuJoCo XML joint order:
+# qpos[7:19] = FR_hip,FR_thigh,FR_calf, FL_hip,FL_thigh,FL_calf, RR_hip,RR_thigh,RR_calf, RL_hip,RL_thigh,RL_calf
+DEFAULT_ANGLES = np.array([
+ -0.1, 0.8, -1.5, # FR
+ 0.1, 0.8, -1.5, # FL
+ -0.1, 1.0, -1.5, # RR
+ 0.1, 1.0, -1.5, # RL
+], dtype=np.float32)
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# Keyboard
+# ═══════════════════════════════════════════════════════════════════════
+from pynput import keyboard
+
+class KB:
+ def __init__(self):
+ self._q = queue.Queue(); self.running = True
+ self.held = set(); self._t = None; self._l = None
+ def _n(self, k):
+ try:
+ if hasattr(k, 'char') and k.char: return k.char.lower()
+ except: pass
+ return str(k).lower()
+ def _w(self):
+ while self.running:
+ try:
+ et, k = self._q.get(timeout=0.05)
+ n = self._n(k)
+ if et == 'press': self.held.add(n)
+ elif et == 'release': self.held.discard(n)
+ except queue.Empty: pass
+ def init(self):
+ self._l = keyboard.Listener(on_press=lambda k: self._q.put(('press', k)),
+ on_release=lambda k: self._q.put(('release', k)))
+ self._l.start()
+ self._t = threading.Thread(target=self._w, daemon=True); self._t.start()
+ def stop(self): self.running = False; self._l.stop()
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# Sensor
+# ═══════════════════════════════════════════════════════════════════════
+def get_sensor(m, d, name):
+ sid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_SENSOR, name)
+ if sid < 0: return None
+ adr = m.sensor_adr[sid]; dim = m.sensor_dim[sid]
+ return d.sensordata[adr:adr+dim].copy()
+
+
+def compute_obs(model, data, commands, last_action):
+ """DreamWaQ observation (Manaro-Alpha order):
+ ang_vel(3) + gravity(3) + commands(3) + joint_pos(12) + joint_vel(12) + actions(12) = 45
+ """
+ obs = np.zeros(NUM_OBS, dtype=np.float32)
+ # ang_vel [0:3]
+ g = get_sensor(model, data, "gyro")
+ obs[0:3] = (g if g is not None else data.qvel[3:6]) * 0.25
+ # gravity [3:6] (read from MuJoCo model, matching training)
+ grav_world = model.opt.gravity.copy()
+ grav_world = grav_world / np.linalg.norm(grav_world) # normalize
+ R = data.xmat[1].reshape(3, 3)
+ obs[3:6] = (R.T @ grav_world).astype(np.float32)
+ # commands [6:9]
+ obs[6:9] = commands * np.array([2.0, 2.0, 0.25], dtype=np.float32)
+ # joint_pos [9:21]
+ obs[9:21] = (data.qpos[7:19] - DEFAULT_ANGLES) * 1.0
+ # joint_vel [21:33]
+ obs[21:33] = data.qvel[6:18] * 0.05
+ # last_action [33:45]
+ obs[33:45] = last_action
+ return np.clip(obs, -CLIP_OBS, CLIP_OBS)
+
+
+# ═══════════════════════════════════════════════════════════════════════
+# Main
+# ═══════════════════════════════════════════════════════════════════════
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--onnx", default=DEFAULT_ONNX)
+ p.add_argument("--terrain", default="flat", choices=["flat", "rough", "stairs", "dreamwaq", "stairs_test", "stairs_box", "flat_stairs"])
+ p.add_argument("--level", type=int, default=0, help="terrain difficulty level 0-9 (0=flat, 9=hardest)")
+ args = p.parse_args()
+
+ # Select XML scene
+ terrain_map = {
+ "flat": "scene_motor_actuator.xml",
+ "rough": "scene_rough_terrain.xml",
+ "stairs": "scene_stairs_terrain.xml",
+ "dreamwaq": "scene_dreamwaq_terrain.xml",
+ "stairs_test": "scene_stairs_test.xml",
+ "stairs_box": "scene_stairs_box.xml",
+ "flat_stairs": "scene_flat_stairs.xml",
+ }
+ xml_file = os.path.join(XML_DIR, terrain_map[args.terrain])
+
+ if not os.path.exists(args.onnx):
+ print(f"[ERROR] ONNX not found: {args.onnx}")
+ print("Run: uv run scripts/export_dreamwaq_onnx.py (after training completes)")
+ sys.exit(1)
+
+ os.chdir(XML_DIR)
+ with open(xml_file) as f:
+ model = mujoco.MjModel.from_xml_string(f.read())
+ data = mujoco.MjData(model)
+
+ # Spawn pose. Hfield heights: MuJoCo z = gp[2] + sbase + (hd * ztop).
+ # The stairs_test terrain has sbase=0, flat platform z=0; just lift by clearance.
+ if args.terrain == "flat_stairs":
+ lvl = max(0, min(1, args.level))
+ col = np.random.randint(0, 4)
+ spawn_y = 4.0 - lvl * 8.0 # level 0 flat at y=+4, level 1 stairs at y=-4
+ spawn_x = -12.0 + col * 8.0 # platform center (cell center x)
+ elif args.terrain == "stairs_test":
+ spawn_x, spawn_y = -7.5, -4.0 # flat approach before first step (1m zone)
+ elif args.terrain == "stairs_box":
+ spawn_x, spawn_y = -2.0, 0.0 # flat ground before stairs
+ elif args.terrain == "dreamwaq":
+ lvl = max(0, min(9, args.level))
+ col = np.random.randint(0, 4) # NUM_COLS=4
+ spawn_y = 36.0 - lvl * 8.0 # level 0 flat at y=+36, level 9 stairs at y=-36
+ spawn_x = -12.0 + col * 8.0 + 4.0 # centre of cell
+ print(f"[Level {lvl}] type={col} spawn=({spawn_x:.1f}, {spawn_y:.1f})")
+ else:
+ spawn_x, spawn_y = (0.0, 0.0)
+
+ # When DISPLAY is a virtual framebuffer (Xvfb), MuJoCo headless rendering is
+ # handled transparently; on a real display this opens a normal GUI window.
+ # No explicit headless flag needed — MuJoCo glfw detects the display type.
+
+ def hfield_z(mx, my):
+ gid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_GEOM, "floor")
+ if gid < 0 or model.geom_type[gid] != mujoco.mjtGeom.mjGEOM_HFIELD:
+ return 0.0
+ hf = model.geom_dataid[gid]
+ nrow, ncol = int(model.hfield_nrow[hf]), int(model.hfield_ncol[hf])
+ sx, sy, ztop, sbase = model.hfield_size[hf]
+ adr = model.hfield_adr[hf]
+ hd = model.hfield_data[adr:adr + nrow * ncol].reshape(nrow, ncol)
+ gp = model.geom_pos[gid]
+ col = int(np.clip(((mx - gp[0]) / sx * 0.5 + 0.5) * (ncol - 1), 0, ncol - 1))
+ row = int(np.clip(((my - gp[1]) / sy * 0.5 + 0.5) * (nrow - 1), 0, nrow - 1))
+ return float(gp[2] + sbase + hd[row, col] * ztop)
+
+ spawn_z = hfield_z(spawn_x, spawn_y) + 0.45 # standing clearance above terrain
+
+ def reset_state():
+ data.qpos[:] = 0
+ data.qpos[0:3] = [spawn_x, spawn_y, spawn_z]; data.qpos[3:7] = [1, 0, 0, 0]
+ data.qpos[7:19] = DEFAULT_ANGLES; data.qvel[:] = 0
+ mujoco.mj_forward(model, data)
+
+ reset_state()
+
+ # ONNX (2 inputs: observations + obs_history)
+ session = ort.InferenceSession(args.onnx, providers=['CPUExecutionProvider'])
+ print(f"[DreamWaQ] {args.onnx}")
+ print(f"[Terrain] {args.terrain}")
+ print(f"[CTRL] W/S前后 Q/E左右 A/D旋转 Space停 R重置 Esc退出")
+
+ kb = KB(); kb.init()
+ view = viewer.launch_passive(model, data)
+
+ # Camera tracking: follow the trunk body
+ trunk_id = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "trunk")
+ view.cam.lookat = data.body(trunk_id).xpos.copy()
+ view.cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING
+ view.cam.trackbodyid = trunk_id
+
+ step = 0
+ vx, vy, wz = 0.0, 0.0, 0.0
+ last_action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ history = np.zeros((1, HISTORY_LEN, NUM_OBS), dtype=np.float32)
+ decimation = 4 # MuJoCo dt=0.005, policy dt=0.02 (DreamWaQ aligned)
+
+ loop_t0 = time.time()
+
+ while view.is_running() and not g_exit_requested:
+ keys = kb.held
+ if 'escape' in keys: break
+ if 'r' in keys:
+ reset_state()
+ last_action[:] = 0; history[:] = 0
+ print("[R] Reset")
+ if ' ' in keys: vx = vy = wz = 0.0
+
+ vx = MAX_VX if 'w' in keys else (-MAX_VX if 's' in keys else 0.0)
+ vy = MAX_VY if 'q' in keys else (-MAX_VY if 'e' in keys else 0.0)
+ wz = MAX_WZ if 'a' in keys else (-MAX_WZ if 'd' in keys else 0.0)
+
+ if step % decimation == 0:
+ cmd = np.array([vx, vy, wz], dtype=np.float32)
+ obs = compute_obs(model, data, cmd, last_action)
+
+ # Shift history + add new obs
+ history = np.concatenate([history[:, 1:, :], obs.reshape(1, 1, -1)], axis=1)
+
+ # ONNX inference
+ outputs = session.run(None, {
+ 'observations': obs.reshape(1, -1).astype(np.float32),
+ 'obs_history': history.astype(np.float32),
+ })
+ action = outputs[0][0]
+ action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
+ last_action = action.copy()
+
+ # PD control
+ target = DEFAULT_ANGLES + action * ACTION_SCALE
+ torques = KP * (target - data.qpos[7:19]) - KD * data.qvel[6:18]
+ data.ctrl[:] = np.clip(torques, -CLIP_ACTIONS, CLIP_ACTIONS)
+ mujoco.mj_step(model, data)
+ view.sync()
+
+ # Time sync (policy at 50Hz = 0.02s per step)
+ expected = step * 0.02
+ elapsed = time.time() - loop_t0
+ if elapsed < expected:
+ time.sleep(expected - elapsed)
+
+ step += 1
+
+ kb.stop(); view.close()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/eval_go1_commands.py b/scripts/eval_go1_commands.py
new file mode 100644
index 0000000..12e70aa
--- /dev/null
+++ b/scripts/eval_go1_commands.py
@@ -0,0 +1,272 @@
+# Copyright (C) 2020-2025 Motphys Technology Co., Ltd. All Rights Reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+# ==============================================================================
+
+"""Evaluate trained Go1 policy with specific velocity commands.
+
+Tests three command patterns:
+ 1. Forward-Backward (前后往返): vx oscillates +1.0 <-> -1.0
+ 2. Left-Right (左右往返): vy oscillates +1.0 <-> -1.0
+ 3. Rotation (旋转): wz oscillates +1.0 <-> -1.0
+
+For each pattern, the command toggles direction every N seconds.
+Metrics (tracking errors, etc.) are logged to CSV files for analysis.
+"""
+
+import logging
+import time
+from pathlib import Path
+
+import numpy as np
+from absl import app, flags
+
+from motrix_envs import registry as env_registry
+from motrix_rl import registry, utils
+from motrix_rl.skrl.jax import wrap_env
+from motrix_rl.skrl.jax.train.ppo import Trainer as SkrlJaxTrainer
+
+logger = logging.getLogger(__name__)
+
+_ENV = flags.DEFINE_string("env", "go1-flat-terrain-walk", "The env to evaluate")
+_POLICY = flags.DEFINE_string(
+ "policy",
+ None,
+ "Path to policy checkpoint. Auto-discovers the latest best_agent.pickle if not specified.",
+)
+_DURATION = flags.DEFINE_float("duration", 4.0, "Seconds per command direction before toggling")
+_OUTDIR = flags.DEFINE_string("outdir", None, "Output directory for CSV logs")
+_RENDER = flags.DEFINE_bool("render", False, "Enable rendering (may crash on headless)")
+
+# ── Helpers ─────────────────────────────────────────────────────────────
+
+
+def _find_best_policy(env_name: str) -> Path:
+ """Auto-discover the best SKRL policy for the given env."""
+ base = Path(f"runs/{env_name}/skrl")
+ if not base.exists():
+ raise FileNotFoundError(f"No training runs found at {base}")
+
+ runs = sorted([d for d in base.iterdir() if d.is_dir()], key=lambda d: d.stat().st_mtime, reverse=True)
+ if not runs:
+ raise FileNotFoundError(f"No training runs found at {base}")
+
+ ckpt_dir = runs[0] / "checkpoints"
+ best = list(ckpt_dir.glob("best_agent.*"))
+ if best:
+ return best[0]
+
+ # Fallback: highest timestep
+ ckpts = list(ckpt_dir.glob("agent_*.pickle"))
+ if not ckpts:
+ raise FileNotFoundError(f"No checkpoints found in {ckpt_dir}")
+
+ def _ts(p):
+ try:
+ return int(p.stem.split("_")[1])
+ except (IndexError, ValueError):
+ return 0
+
+ return max(ckpts, key=_ts)
+
+
+def _save_log(outdir: Path, label: str, records: list[dict]) -> None:
+ """Save records to a CSV file."""
+ outdir.mkdir(parents=True, exist_ok=True)
+ import csv
+
+ path = outdir / f"{label}.csv"
+ with open(path, "w", newline="") as f:
+ writer = csv.DictWriter(f, fieldnames=records[0].keys())
+ writer.writeheader()
+ writer.writerows(records)
+ logger.info(f"Saved {len(records)} records → {path}")
+
+
+# ── Command generators ──────────────────────────────────────────────────
+
+
+def _command_forward_backward(t: float, period: float):
+ """vx oscillates +1.0 / -1.0, vy=0, wz=0."""
+ half = period / 2.0
+ phase = (t % period) / half # 0..1 forward, 1..2 backward
+ vx = 1.0 if phase < 1.0 else -1.0
+ return np.array([vx, 0.0, 0.0], dtype=np.float32)
+
+
+def _command_left_right(t: float, period: float):
+ """vy oscillates +1.0 / -1.0, vx=0, wz=0."""
+ half = period / 2.0
+ phase = (t % period) / half
+ vy = 1.0 if phase < 1.0 else -1.0
+ return np.array([0.0, vy, 0.0], dtype=np.float32)
+
+
+def _command_rotation(t: float, period: float):
+ """wz oscillates +1.0 / -1.0, vx=0, vy=0."""
+ half = period / 2.0
+ phase = (t % period) / half
+ wz = 1.0 if phase < 1.0 else -1.0
+ return np.array([0.0, 0.0, wz], dtype=np.float32)
+
+
+# ── Main evaluation logic ───────────────────────────────────────────────
+
+
+def _run_pattern(
+ trainer: SkrlJaxTrainer,
+ policy_path: str,
+ pattern_name: str,
+ command_fn,
+ period: float,
+ total_steps: int,
+ ctrl_dt: float,
+) -> list[dict]:
+ """Run one command pattern and return tracking records."""
+ logger.info(f"--- {pattern_name} ---")
+
+ rlcfg = trainer._rlcfg
+ env = env_registry.make(trainer._env_name, sim_backend=trainer._sim_backend, num_envs=1)
+ env = wrap_env(env, enable_render=False)
+
+ # Build fresh agent for this run
+ models = trainer._make_model(env, rlcfg)
+ ppo_cfg = rlcfg.runner.agent.to_dict()
+ from motrix_rl.skrl.jax.train.ppo import _add_runtime_config
+ _add_runtime_config(ppo_cfg, env)
+ agent = trainer._make_agent(models, env, ppo_cfg, rlcfg.runner.memory)
+ agent.load(policy_path)
+
+ obs, info = env.reset()
+ state = env._env.state
+ records = []
+
+ for step in range(total_steps):
+ t = step * ctrl_dt
+
+ # Override command in env state
+ cmd = command_fn(t, period)
+ state.info["commands"] = cmd.reshape(1, -1)
+
+ # Recompute observation with the new command
+ new_obs = env._env._get_obs(state.data, state.info)
+ state = state.replace(obs=new_obs)
+ obs = new_obs
+
+ # Agent inference
+ outputs = agent.act(obs, timestep=0, timesteps=0)
+ actions = outputs[-1].get("mean_actions", outputs[0])
+ obs, reward, terminated, truncated, info = env.step(actions)
+ state = env._env.state
+
+ # Collect tracking data
+ lin_vel = env._env.get_local_linvel(state.data)[0] # [vx, vy, vz] (body frame)
+ gyro = env._env.get_gyro(state.data)[0]
+
+ tracking_err_xy = np.linalg.norm(cmd[:2] - lin_vel[:2])
+ tracking_err_yaw = abs(cmd[2] - gyro[2])
+
+ records.append({
+ "step": step,
+ "time": round(t, 3),
+ "cmd_vx": round(float(cmd[0]), 4),
+ "cmd_vy": round(float(cmd[1]), 4),
+ "cmd_wz": round(float(cmd[2]), 4),
+ "actual_vx": round(float(lin_vel[0]), 4),
+ "actual_vy": round(float(lin_vel[1]), 4),
+ "actual_vz": round(float(lin_vel[2]), 4),
+ "actual_wz": round(float(gyro[2]), 4),
+ "tracking_err_xy": round(float(tracking_err_xy), 6),
+ "tracking_err_yaw": round(float(tracking_err_yaw), 6),
+ "reward": round(float(reward[0][0]), 6),
+ })
+
+ if step % 100 == 0:
+ logger.info(
+ f" [{pattern_name}] step {step:5d}/{total_steps} "
+ f"cmd=[{cmd[0]:+.1f},{cmd[1]:+.1f},{cmd[2]:+.1f}] "
+ f"actual_v=[{lin_vel[0]:+.3f},{lin_vel[1]:+.3f},{gyro[2]:+.3f}] "
+ f"track_err_xy={tracking_err_xy:.4f} "
+ f"reward={float(reward[0][0]):.4f}"
+ )
+
+ env.close()
+ return records
+
+
+def main(argv):
+ env_name = _ENV.value
+ ctrl_dt = 0.01 # matches Go1WalkNpEnvCfg.ctrl_dt
+ period = _DURATION.value * 2.0 # full cycle: forward + backward
+
+ # Resolve policy path
+ if _POLICY.present:
+ policy_path = _POLICY.value
+ else:
+ policy_path = str(_find_best_policy(env_name))
+ logger.info(f"Policy: {policy_path}")
+
+ # Build trainer (used to construct models & agent)
+ rlcfg = registry.default_rl_cfg(env_name, "skrl", backend="jax")
+ trainer = SkrlJaxTrainer(env_name, sim_backend=None, enable_render=False)
+ trainer._rlcfg = rlcfg
+
+ # Total steps per pattern = enough full cycles
+ cycles = 3
+ total_steps = int(cycles * period / ctrl_dt)
+
+ # Output directory
+ if _OUTDIR.present:
+ outdir = Path(_OUTDIR.value)
+ else:
+ ts = time.strftime("%y-%m-%d_%H-%M-%S")
+ outdir = Path(f"runs/{env_name}/eval_{ts}")
+ logger.info(f"Output directory: {outdir}")
+
+ patterns = [
+ ("forward_backward", _command_forward_backward),
+ ("left_right", _command_left_right),
+ ("rotation", _command_rotation),
+ ]
+
+ summary = {}
+ for label, cmd_fn in patterns:
+ records = _run_pattern(
+ trainer, policy_path, label, cmd_fn,
+ period=period, total_steps=total_steps, ctrl_dt=ctrl_dt,
+ )
+ _save_log(outdir, label, records)
+
+ # Summary stats (steady-state, skip first 2s for stabilization)
+ warmup = int(2.0 / ctrl_dt)
+ steady = records[warmup:]
+ if steady:
+ avg_err = np.mean([r["tracking_err_xy"] for r in steady])
+ avg_reward = np.mean([r["reward"] for r in steady])
+ else:
+ avg_err, avg_reward = float("nan"), float("nan")
+ summary[label] = {"avg_tracking_err_xy": avg_err, "avg_reward": avg_reward}
+ logger.info(f" [{label}] steady-state avg tracking_err_xy = {avg_err:.4f}, avg_reward = {avg_reward:.4f}")
+
+ # Print summary
+ print("\n" + "=" * 70)
+ print("EVALUATION SUMMARY")
+ print("=" * 70)
+ for label, stats in summary.items():
+ print(f" {label:25s} track_err_xy={stats['avg_tracking_err_xy']:.4f} avg_reward={stats['avg_reward']:.4f}")
+ print(f"\nDetailed CSV logs saved to: {outdir}")
+ print("=" * 70)
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/scripts/export_dreamwaq_onnx.py b/scripts/export_dreamwaq_onnx.py
new file mode 100644
index 0000000..ad41804
--- /dev/null
+++ b/scripts/export_dreamwaq_onnx.py
@@ -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()
diff --git a/scripts/export_dreamwaq_rsl_onnx.py b/scripts/export_dreamwaq_rsl_onnx.py
new file mode 100644
index 0000000..bfa863a
--- /dev/null
+++ b/scripts/export_dreamwaq_rsl_onnx.py
@@ -0,0 +1,88 @@
+#!/usr/bin/env python3
+"""Export a DreamWaQ rsl_rl checkpoint (ActorCritic_DWAQ) to ONNX for MuJoCo.
+
+rsl_rl has NO state preprocessor, so no normalization is needed (unlike SKRL).
+The exported model uses the MEAN CENet code (deterministic deploy).
+
+ONNX: inputs observations(1,45) + obs_history(1,5,45) -> actions(1,12)
+(matches scripts/dreamwaq_sim2sim_mujoco.py interface)
+
+Usage:
+ uv run scripts/export_dreamwaq_rsl_onnx.py # auto-find latest
+ uv run scripts/export_dreamwaq_rsl_onnx.py --checkpoint runs/.../model_700.pt
+"""
+import argparse, glob, os, sys
+import torch
+import torch.nn as tnn
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from motrix_rl.dwaq_rsl.actor_critic_dwaq import ActorCritic_DWAQ
+
+PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+
+
+class DwaqInfer(tnn.Module):
+ """Deterministic inference: CENet mean code + actor."""
+ def __init__(self, ac: ActorCritic_DWAQ):
+ super().__init__()
+ self.encoder = ac.encoder
+ self.encode_mean_vel = ac.encode_mean_vel
+ self.encode_mean_latent = ac.encode_mean_latent
+ self.actor = ac.actor
+
+ def forward(self, obs, obs_history):
+ h = self.encoder(obs_history.reshape(obs_history.shape[0], -1)) # (B,225)->(B,64)
+ vel = self.encode_mean_vel(h) # (B,3) mean velocity estimate
+ latent = self.encode_mean_latent(h) # (B,16) mean latent
+ code = torch.cat([vel, latent], dim=-1) # (B,19) = [vel, latent]
+ x = torch.cat([code, obs], dim=-1) # (B,64) = [code, obs]
+ return self.actor(x) # (B,12)
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--checkpoint", default=None)
+ p.add_argument("--output", default=os.path.join(PROJECT, "exports_go1_dreamwaq", "policy.onnx"))
+ p.add_argument("--num-obs", type=int, default=45)
+ p.add_argument("--num-priv", type=int, default=235)
+ p.add_argument("--num-hist", type=int, default=5)
+ p.add_argument("--num-act", type=int, default=12)
+ args = p.parse_args()
+
+ if args.checkpoint is None:
+ runs = sorted(glob.glob(os.path.join(PROJECT, "runs", "go1-dreamwaq-walk", "rsl_dwaq", "*")),
+ key=os.path.getmtime)
+ if not runs:
+ print("[ERROR] no rsl_dwaq runs found"); sys.exit(1)
+ models = glob.glob(os.path.join(runs[-1], "model_*.pt"))
+ args.checkpoint = max(models, key=os.path.getmtime)
+
+ print(f"[rsl-ONNX] checkpoint: {args.checkpoint}")
+ cenet_out = 19
+ ac = ActorCritic_DWAQ(
+ args.num_obs + cenet_out, # actor in = 64
+ args.num_priv, # critic in = 235
+ args.num_act, # 12
+ args.num_hist * args.num_obs, # cenet in = 225
+ cenet_out, # 19
+ )
+ ckpt = torch.load(args.checkpoint, map_location="cpu")
+ ac.load_state_dict(ckpt["model_state_dict"])
+ ac.eval()
+
+ model = DwaqInfer(ac).eval()
+ os.makedirs(os.path.dirname(args.output), exist_ok=True)
+ torch.onnx.export(
+ model,
+ (torch.zeros(1, args.num_obs), torch.zeros(1, args.num_hist, args.num_obs)),
+ args.output,
+ 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"[rsl-ONNX] exported -> {args.output} (no normalization; mean CENet code)")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/export_go1_no_linevel_onnx.py b/scripts/export_go1_no_linevel_onnx.py
new file mode 100644
index 0000000..43dcc64
--- /dev/null
+++ b/scripts/export_go1_no_linevel_onnx.py
@@ -0,0 +1,289 @@
+#!/usr/bin/env python3
+"""Export JAX/Flax-trained SKRL Go1 (no-linevel, 57-dim) policy to ONNX.
+
+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_no_linevel_onnx.py
+ uv run scripts/export_go1_no_linevel_onnx.py --output ./my_exports
+"""
+
+import argparse
+import os
+import pickle
+from pathlib import Path
+
+import msgpack
+import numpy as np
+
+
+def _decode_flax_array(ext) -> np.ndarray | None:
+ 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
+
+ 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:
+ with open(ckpt_path, "rb") as f:
+ data = pickle.load(f)
+
+ 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)
+
+ 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 -----------------------------------------------------------
+
+import torch
+import torch.nn as nn
+
+
+class PolicyTorch(nn.Module):
+ 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):
+ 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)
+
+
+def flax_to_torch_weights(flax_params: dict, obs_dim: int, hidden_dims: list[int], action_dim: int) -> dict:
+ state_dict = {}
+ layer_names = sorted([k for k in flax_params if k.startswith("Dense_")])
+
+ hidden_dense = layer_names[:-1]
+ layer_idx = 0
+ for name in hidden_dense:
+ layer_params = flax_params[name]
+ kernel = layer_params["kernel"]
+ bias = layer_params["bias"]
+ 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
+
+ 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
+
+
+# -- Config ------------------------------------------------------------------
+
+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)
+
+# 57-dim observation layout (NO linear velocity):
+# [0:3] gyro (scaled *0.25)
+# [3:6] gravity vector (body frame)
+# [6:18] joint angle deviation (scaled *1.0)
+# [18:30] joint velocity (scaled *0.05)
+# [30:42] last action (raw)
+# [42:45] command [vx*2.0, vy*2.0, wz*0.25]
+# [45:57] foot contact forces (body frame, raw)
+
+OBS_SCALES = {
+ "ang_vel": 0.25,
+ "dof_pos": 1.0,
+ "dof_vel": 0.05,
+ "contact_force": 1.0, # raw, no scaling
+}
+
+ACTION_SCALE = 0.05
+KP, KD = 80.0, 1.0
+CLIP_ACTIONS = 23.7
+CLIP_OBS = 100.0
+
+
+def auto_discover_checkpoint(env_name: str) -> str:
+ """Find the latest best_agent checkpoint for the given env."""
+ base_dir = Path(f"runs/{env_name}/skrl")
+ if not base_dir.exists():
+ raise FileNotFoundError(f"No training results found: {base_dir}")
+
+ runs = sorted([d for d in base_dir.iterdir() if d.is_dir()], key=lambda d: d.stat().st_mtime, reverse=True)
+ for run_dir in runs:
+ ckpt = run_dir / "checkpoints" / "best_agent.pickle"
+ if ckpt.exists():
+ return str(ckpt)
+
+ raise FileNotFoundError(f"No best_agent.pickle found in {base_dir}")
+
+
+def export(checkpoint_path: str, output_dir: str):
+ 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)
+
+ 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}]")
+
+ 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()
+
+ 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)} ...")
+
+ 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}")
+
+ # Normalizer stats
+ 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}")
+
+ # Metadata
+ meta_path = os.path.join(output_dir, "metadata.txt")
+ with open(meta_path, "w") as f:
+ f.write("# Go1 No-Linevel Terrain Walk - ONNX Policy Metadata\n")
+ f.write(f"env: go1-stairs-terrain-walk-no-linevel\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("\n# Observation layout (57 dims, NO linear velocity):\n")
+ f.write(" [0:3] gyro * ang_vel_scale\n")
+ f.write(" [3:6] gravity (body frame)\n")
+ f.write(" [6:18] joint_angle_deviation * dof_pos_scale\n")
+ f.write(" [18:30] joint_vel * dof_vel_scale\n")
+ f.write(" [30:42] last_actions (raw)\n")
+ f.write(" [42:45] commands [vx*2.0, vy*2.0, wz*0.25]\n")
+ f.write(" [45:57] foot_contact_forces (body frame, raw)\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("\n# Observation scales:\n")
+ for k, v in OBS_SCALES.items():
+ f.write(f" {k}: {v}\n")
+ f.write(" command_scale: [2.0, 2.0, 0.25]\n")
+ print(f"Metadata saved to: {meta_path}")
+
+ return onnx_path
+
+
+def main():
+ parser = argparse.ArgumentParser(description="Export JAX-trained Go1 no-linevel policy to ONNX")
+ parser.add_argument("--checkpoint", type=str, default=None,
+ help="Path to SKRL JAX checkpoint (auto-discovered if not set)")
+ parser.add_argument("--output", type=str, default="exports_go1_no_linevel",
+ help="Output directory")
+ args = parser.parse_args()
+
+ if args.checkpoint:
+ ckpt_path = args.checkpoint
+ else:
+ ckpt_path = auto_discover_checkpoint("go1-stairs-terrain-walk-no-linevel")
+
+ if not os.path.exists(ckpt_path):
+ print(f"Error: checkpoint not found: {ckpt_path}")
+ return 1
+
+ print(f"Loading checkpoint: {ckpt_path}")
+ onnx_path = export(ckpt_path, args.output)
+ print(f"\nDone! ONNX model ready for sim2sim:")
+ print(f" {onnx_path}")
+ return 0
+
+
+if __name__ == "__main__":
+ exit(main())
diff --git a/scripts/export_go1_onnx.py b/scripts/export_go1_onnx.py
new file mode 100644
index 0000000..e0264ea
--- /dev/null
+++ b/scripts/export_go1_onnx.py
@@ -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())
diff --git a/scripts/gen_flat_stairs.py b/scripts/gen_flat_stairs.py
new file mode 100644
index 0000000..521401b
--- /dev/null
+++ b/scripts/gen_flat_stairs.py
@@ -0,0 +1,119 @@
+#!/usr/bin/env python3
+"""Generate 2-level terrain: level 0=flat, level 1=pyramid stairs.
+
+1cm = 1px. OpenCV draws concentric filled rectangles (outside→in, 10 steps).
+Higher values overwrite lower = convex pyramid (stairs up toward center).
+Lower values overwrite higher = concave pyramid (stairs down from center).
+
+Usage:
+ uv run scripts/gen_flat_stairs.py # convex (stairs UP)
+ uv run scripts/gen_flat_stairs.py --concave # concave (stairs DOWN)
+ uv run scripts/gen_flat_stairs.py --step-h 0.10 --num-steps 5
+"""
+import cv2, numpy as np, os, argparse
+
+# ═══ fixed params ═══
+HS = 0.01 # 1cm/px
+VS = 0.005 # height unit = 0.5cm
+CELL_M = 8.0 # 8m cell
+BORDER_M = 5.0 # 5m border
+NUM_ROWS = 2 # flat + stairs
+NUM_COLS = 4 # columns
+
+CELL_PX = int(CELL_M / HS) # 800
+BORDER_PX = int(BORDER_M / HS) # 500
+PLATFORM_PX = int(1.0 / HS) # 1m platform = 100px
+
+TOT_ROWS = NUM_ROWS * CELL_PX + 2 * BORDER_PX
+TOT_COLS = NUM_COLS * CELL_PX + 2 * BORDER_PX
+
+# ═══ stairs params (override via CLI) ═══
+NUM_STEPS = 10 # 10 steps
+STEP_H_CM = 20 # 20cm rise per step
+STEP_D_CM = 20 # 20cm tread per step
+
+STEP_H_VS = int(STEP_H_CM / 100.0 / VS) # 0.20 / 0.005 = 40
+STEP_D_PX = int(STEP_D_CM / 100.0 / HS) # 0.20 / 0.01 = 20
+
+
+def draw_pyramid(canvas, x0, y0, num_steps, step_d_px, step_h_vs, concave=False):
+ """Draw concentric rectangles from outside→in.
+
+ Convex: edge=0 → platform=max (stairs up toward center)
+ Concave: raise whole cell to max, then draw pit: edge=max → platform=0
+ """
+ cx, cy = x0 + CELL_PX // 2, y0 + CELL_PX // 2
+ p2 = PLATFORM_PX // 2
+ h_max = step_h_vs * num_steps
+
+ # Fill cell to cell boundary with reference-plane height
+ half_max = CELL_PX // 2 # extend to cell edge
+ cv2.rectangle(canvas, (cx - half_max, cy - half_max),
+ (cx + half_max, cy + half_max), int(h_max), -1)
+
+ if concave:
+ # Pit: rings going DOWN from reference plane
+ for i in range(num_steps + 1):
+ half = p2 + (num_steps - i) * step_d_px
+ x1, y1 = cx - half, cy - half
+ x2, y2 = cx + half, cy + half
+ h = h_max - step_h_vs * i
+ cv2.rectangle(canvas, (x1, y1), (x2, y2), int(h), -1)
+ else:
+ # Mound: rings going UP from reference plane
+ for i in range(num_steps + 1):
+ half = p2 + (num_steps - i) * step_d_px
+ x1, y1 = cx - half, cy - half
+ x2, y2 = cx + half, cy + half
+ h = h_max + step_h_vs * i
+ cv2.rectangle(canvas, (x1, y1), (x2, y2), int(h), -1)
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--concave", action="store_true", help="concave pyramid (stairs down from center)")
+ p.add_argument("--step-h", type=float, default=0.20, help="step rise (m)")
+ p.add_argument("--step-d", type=float, default=0.20, help="step tread (m)")
+ p.add_argument("--num-steps", type=int, default=10, help="number of steps")
+ args = p.parse_args()
+
+ step_h_vs = int(args.step_h / VS)
+ step_d_px = int(args.step_d / HS)
+ total_h = step_h_vs * args.num_steps * VS
+
+ print(f"Building {TOT_COLS}×{TOT_ROWS}px ({TOT_COLS*HS:.0f}×{TOT_ROWS*HS:.0f}m)")
+ print(f" type={'concave' if args.concave else 'convex'} "
+ f"steps={args.num_steps} rise={step_h_vs*VS*100:.0f}cm "
+ f"tread={step_d_px*HS*100:.0f}cm total_h={total_h*100:.0f}cm")
+
+ canvas = np.zeros((TOT_ROWS, TOT_COLS), dtype=np.uint16)
+
+ for row in range(NUM_ROWS):
+ for col in range(NUM_COLS):
+ x0 = BORDER_PX + col * CELL_PX
+ y0 = BORDER_PX + row * CELL_PX
+ if row == 1:
+ # alternate convex/concave across cols
+ concave_cell = (col % 2 == 1)
+ draw_pyramid(canvas, x0, y0, args.num_steps,
+ step_d_px, step_h_vs, concave_cell)
+
+ hf_m = canvas.astype(np.float32) * VS
+ z_min, z_max = float(hf_m.min()), float(hf_m.max())
+ z_range = max(z_max - z_min, 0.001)
+
+ png = ((hf_m - z_min) / z_range * 65535).astype(np.uint16)
+
+ out_dir = os.path.join(os.path.dirname(__file__), "..",
+ "motrix_envs", "src", "motrix_envs", "locomotion",
+ "go1", "xmls", "assets")
+ os.makedirs(out_dir, exist_ok=True)
+ out_path = os.path.join(out_dir, "flat_stairs.png")
+ cv2.imwrite(out_path, png)
+ print(f" saved: {out_path}")
+ print(f" XML: size=\"{TOT_COLS*HS/2:.1f} {TOT_ROWS*HS/2:.1f} "
+ f"{z_range:.3f} {max(z_min,0.001):.3f}\"")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/gen_stairs_box.py b/scripts/gen_stairs_box.py
new file mode 100644
index 0000000..3f1bba1
--- /dev/null
+++ b/scripts/gen_stairs_box.py
@@ -0,0 +1,95 @@
+#!/usr/bin/env python3
+"""Generate box-geom stairs XML for MuJoCo sim2sim.
+
+Each step is a separate box with vertical rises — much steeper than hfield.
+
+Usage:
+ uv run scripts/gen_stairs_box.py # default: 10 steps × 6cm = 60cm
+ uv run scripts/gen_stairs_box.py --step-height 0.04 --num-steps 5
+ uv run scripts/gen_stairs_box.py --step-height 0.10 --num-steps 8 --step-depth 0.4
+"""
+import argparse, os
+
+TPL = '''
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+{steps}
+
+
+
+
+
+'''
+
+STEP_TPL = ' \n'
+PLAT_TPL = ' \n'
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--step-height", type=float, default=0.06, help="rise per step [m]")
+ p.add_argument("--step-depth", type=float, default=0.30, help="tread depth per step [m]")
+ p.add_argument("--num-steps", type=int, default=10, help="number of steps")
+ p.add_argument("--box-thickness", type=float, default=0.03, help="box half-height [m]")
+ args = p.parse_args()
+
+ h = args.step_height
+ d = args.step_depth
+ n = args.num_steps
+ sz = args.box_thickness # half-height of each box
+
+ steps_xml = ""
+ for i in range(n):
+ x = i * d
+ z = i * h + sz # center of box = step top surface - sz
+ steps_xml += STEP_TPL.format(n=i, sx=d/2, sz=sz, x=x, z=z)
+
+ # Platform at top
+ plat_x = n * d + 0.5
+ plat_z = n * h + sz
+ steps_xml += PLAT_TPL.format(sx=0.5, sz=sz, x=plat_x, z=plat_z)
+
+ # Fill box under stairs
+ total_depth = n * d
+ total_height = n * h
+ fill_sx = total_depth / 2
+ fill_sz = total_height / 2
+ fill_x = total_depth / 2
+ fill_z = -fill_sz
+
+ out = TPL.format(steps=steps_xml.rstrip(),
+ fill_sx=fill_sx, fill_sz=fill_sz,
+ fill_x=fill_x, fill_z=fill_z)
+
+ out_dir = os.path.join(os.path.dirname(__file__), "..",
+ "motrix_envs", "src", "motrix_envs", "locomotion",
+ "go1", "xmls")
+ out_path = os.path.join(out_dir, "scene_stairs_box.xml")
+ with open(out_path, "w") as f:
+ f.write(out)
+
+ max_h = n * h
+ print(f"Generated {n} steps × {h*100:.0f}cm = {max_h*100:.0f}cm total")
+ print(f" step depth: {d*100:.0f}cm box thickness: {sz*200:.0f}cm")
+ print(f" saved: {out_path}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/gen_stairs_test.py b/scripts/gen_stairs_test.py
new file mode 100644
index 0000000..ae79dd1
--- /dev/null
+++ b/scripts/gen_stairs_test.py
@@ -0,0 +1,108 @@
+#!/usr/bin/env python3
+"""Generate standard linear stairs for MuJoCo sim2sim testing.
+
+Each cell: flat 2m approach -> N linear steps -> flat 2m platform -> N steps down -> flat edge.
+Treads are horizontal, rises are vertical (1px = 0.1m wide, acceptable for hfield).
+
+Usage:
+ uv run scripts/gen_stairs_test.py --step-height 0.07 --step-depth 0.31
+"""
+import numpy as np, os, argparse
+from PIL import Image
+
+HS = 0.1 # horizontal scale [m/px]
+VS = 0.005 # vertical scale [m/unit]
+CELL_M = 8.0
+PLATFORM_M = 4.0 # bigger flat platform -> fewer steps
+CELL_PX = int(CELL_M / HS) # 80
+PLATFORM_PX = int(PLATFORM_M / HS) # 40
+BORDER_M = 2.0
+BORDER_PX = int(BORDER_M / HS) # 20
+NUM_CELLS = 2
+TOT_PX = NUM_CELLS * CELL_PX + 2 * BORDER_PX
+TOTAL_M = TOT_PX * HS
+
+np.random.seed(42)
+
+
+def make_linear_stairs(step_height_m, step_depth_m=0.31):
+ """Linear stairs: flat approach -> N steps up -> flat platform -> edge.
+ Each step has a flat horizontal tread and (essentially) vertical rise."""
+ t = np.zeros((CELL_PX, CELL_PX), dtype=np.int16)
+ sd = int(step_depth_m / HS) # tread depth in px
+ sh = int(step_height_m / VS) # rise height in pixel units
+
+ # How many steps fit on each side of the platform?
+ avail = (CELL_PX - PLATFORM_PX) // 2
+ n_steps = avail // max(sd, 1)
+ if n_steps < 1:
+ n_steps = 1
+
+ edge = (CELL_PX - PLATFORM_PX - n_steps * sd) // 2 # remaining flat on each side
+
+ # Draw steps going UP from left (in +x direction)
+ # Each step: flat tread at current height, then rise to next height
+ x = edge
+ h = 0
+ for i in range(n_steps):
+ x_next = x + sd
+ t[:, x:x_next] = h # tread at current height
+ x = x_next
+ h += sh
+
+ # Platform (flat at max height)
+ plat_start = x
+ plat_end = plat_start + PLATFORM_PX
+ t[:, plat_start:plat_end] = h
+
+ # Continue stairs going DOWN on the right (optional: mirror)
+ x = plat_end
+ for i in range(n_steps):
+ h -= sh
+ x_next = x + sd
+ t[:, x:x_next] = h
+ x = x_next
+
+ return t
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--step-height", type=float, default=0.15,
+ help="step rise height in metres (default 0.15)")
+ p.add_argument("--step-depth", type=float, default=0.50,
+ help="step tread depth in metres (default 0.50)")
+ args = p.parse_args()
+
+ print(f"Linear stairs: step_h={args.step_height:.2f}m tread={args.step_depth:.2f}m platform={PLATFORM_M:.0f}m")
+
+ hf_raw = np.zeros((TOT_PX, TOT_PX), dtype=np.int16)
+ for i in range(NUM_CELLS):
+ for j in range(NUM_CELLS):
+ cell = make_linear_stairs(args.step_height, args.step_depth)
+ y0 = BORDER_PX + i * CELL_PX
+ x0 = BORDER_PX + j * CELL_PX
+ hf_raw[y0:y0 + CELL_PX, x0:x0 + CELL_PX] = cell
+
+ hf_m = hf_raw.astype(np.float32) * VS
+ z_min = float(hf_m.min())
+ z_max = float(hf_m.max())
+ z_range = max(z_max - z_min, 0.001)
+
+ print(f" height: [{z_min:.3f}, {z_max:.3f}]m z_scale={z_range:.3f} max={z_max*100:.0f}cm")
+
+ png = ((hf_m - z_min) / z_range * 65535.0).astype(np.uint16)
+
+ out_dir = os.path.join(os.path.dirname(__file__), "..",
+ "motrix_envs", "src", "motrix_envs", "locomotion",
+ "go1", "xmls", "assets")
+ os.makedirs(out_dir, exist_ok=True)
+ out_path = os.path.join(out_dir, "stairs_test.png")
+ Image.fromarray(png).save(out_path)
+ print(f" saved: {out_path}")
+ sbase = max(z_min, 0.001)
+ print(f" XML: size=\"{TOTAL_M/2:.1f} {TOTAL_M/2:.1f} {z_range:.3f} {sbase:.3f}\"")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/go1_no_linevel_sim2sim_mujoco.py b/scripts/go1_no_linevel_sim2sim_mujoco.py
new file mode 100644
index 0000000..f33e179
--- /dev/null
+++ b/scripts/go1_no_linevel_sim2sim_mujoco.py
@@ -0,0 +1,428 @@
+#!/usr/bin/env python3
+"""MuJoCo sim2sim for go1-stairs-terrain-walk-no-linevel (57-dim obs, no linvel).
+
+Loads the ONNX policy exported by export_go1_no_linevel_onnx.py and runs
+inference in MuJoCo with PD control.
+
+Usage:
+ # Default: combined flat+rough+stairs terrain, random spawn
+ uv run scripts/go1_no_linevel_sim2sim_mujoco.py
+
+ # Specific terrain
+ uv run scripts/go1_no_linevel_sim2sim_mujoco.py --terrain flat
+ uv run scripts/go1_no_linevel_sim2sim_mujoco.py --terrain rough
+ uv run scripts/go1_no_linevel_sim2sim_mujoco.py --terrain stairs
+
+ # Custom ONNX path
+ uv run scripts/go1_no_linevel_sim2sim_mujoco.py --onnx ./exports_go1_no_linevel/policy.onnx
+
+Keyboard controls:
+ W/S - forward/backward
+ A/D - turn left/right
+ Q/E - strafe left/right
+ Space - stop
+ R - reset robot
+ 1/2/3 - switch terrain (flat/rough/stairs)
+ Esc - quit
+"""
+
+import numpy as np
+import mujoco
+from mujoco import viewer
+import os
+import threading
+import signal
+import queue
+import argparse
+import time
+
+g_exit_requested = False
+
+
+def signal_handler(signum, frame):
+ global g_exit_requested
+ g_exit_requested = True
+
+
+signal.signal(signal.SIGINT, signal_handler)
+
+# ============================================================
+# Paths
+# ============================================================
+_PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+DEFAULT_ONNX_PATH = os.path.join(_PROJECT_DIR, "exports_go1_no_linevel", "policy.onnx")
+MOTRIX_XML_DIR = os.path.join(_PROJECT_DIR, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls")
+
+# ============================================================
+# MotrixLab parameters (matching cfg.py + walk_stairs_terrain_no_linevel.py)
+# ============================================================
+NUM_OBS = 57 # 57-dim: NO linear velocity, WITH contact forces
+NUM_ACTIONS = 12
+OBS_SCALES = {"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_OBSERVATIONS = 100.0
+MAX_LIN_VEL_X = 1.0
+MAX_LIN_VEL_Y = 1.0
+MAX_ANG_VEL = 1.0
+
+# ============================================================
+# Joint names and order
+# ============================================================
+POLICY_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",
+]
+
+DEFAULT_JOINT_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)
+
+FEET = ["FR", "FL", "RR", "RL"]
+
+# Terrain spawn positions (world Y)
+TERRAIN_SPAWN = {
+ "flat": np.array([0.0, 54.0, 0.42], dtype=np.float64),
+ "rough": np.array([0.0, 32.0, 0.42], dtype=np.float64),
+ "stairs": np.array([0.0, 0.0, 0.42], dtype=np.float64),
+}
+
+# ============================================================
+# Keyboard input
+# ============================================================
+from pynput import keyboard
+
+
+class KeyboardReader:
+ def __init__(self):
+ self._event_queue = queue.Queue()
+ self.running = True
+ self.shared_keys_held = set()
+ self.shared_one_shot = set()
+ self._reader_thread = None
+ self._listener = None
+
+ def _normalize_key(self, key):
+ try:
+ if hasattr(key, "char") and key.char is not None:
+ return key.char.lower()
+ except Exception:
+ pass
+ key_str = str(key)
+ if key_str == "Key.esc":
+ return "escape"
+ elif key_str == "Key.space":
+ return "space"
+ elif key_str.startswith("Key."):
+ return key_str.lower()
+ return key_str.lower()
+
+ def _reader_worker(self):
+ while self.running:
+ try:
+ event_type, key = self._event_queue.get(timeout=0.05)
+ k = self._normalize_key(key)
+ if event_type == "press":
+ self.shared_keys_held.add(k)
+ self.shared_one_shot.discard(k)
+ elif event_type == "release":
+ self.shared_keys_held.discard(k)
+ self.shared_one_shot.discard(k)
+ except queue.Empty:
+ pass
+
+ def init(self):
+ def on_press(key):
+ self._event_queue.put(("press", key))
+
+ def on_release(key):
+ self._event_queue.put(("release", key))
+
+ try:
+ self._listener = keyboard.Listener(on_press=on_press, on_release=on_release)
+ self._listener.start()
+ self._reader_thread = threading.Thread(target=self._reader_worker, daemon=True)
+ self._reader_thread.start()
+ print("[INFO] Keyboard listener started")
+ except Exception as e:
+ print(f"[WARN] Cannot init keyboard: {e}")
+
+ def is_key_pressed(self, key):
+ k = self._normalize_key(key) if isinstance(key, str) else self._normalize_key(key)
+ if k not in self.shared_keys_held or k in self.shared_one_shot:
+ return False
+ self.shared_one_shot.add(k)
+ return True
+
+ def is_key_held(self, key):
+ k = self._normalize_key(key) if isinstance(key, str) else self._normalize_key(key)
+ return k in self.shared_keys_held
+
+ def restore(self):
+ self.running = False
+ if self._listener:
+ self._listener.stop()
+
+
+# ============================================================
+# Sensor reading
+# ============================================================
+def get_sensor(model, data, name):
+ sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, name)
+ if sid < 0:
+ return None
+ adr = model.sensor_adr[sid]
+ dim = model.sensor_dim[sid]
+ return data.sensordata[adr : adr + dim].copy()
+
+
+def read_contact_forces(model, data, base_rot):
+ """Read foot contact forces (12-dim, body frame) from MuJoCo contact sensors.
+
+ Tries _stairs, _rough, _flat suffixes for each foot, picking the first
+ sensor that returns non-zero data. In MuJoCo, `data="force"` returns a
+ scalar (normal force). We construct a 3D force vector by projecting onto
+ the body-frame Z axis as an approximation.
+ """
+ forces = np.zeros(12, dtype=np.float32)
+ for i, foot in enumerate(FEET):
+ f_scalar = 0.0
+ for suffix in ["_stairs", "_rough", "_flat"]:
+ name = f"{foot}_foot_contact{suffix}"
+ v = get_sensor(model, data, name)
+ if v is not None and np.abs(v[0]) > 1e-6:
+ f_scalar = v[0]
+ break
+ # Assume contact force is approximately vertical (world Z),
+ # rotate into body frame
+ force_world = np.array([0.0, 0.0, f_scalar], dtype=np.float64)
+ force_body = base_rot.T @ force_world
+ forces[i * 3 : i * 3 + 3] = force_body.astype(np.float32)
+ return forces
+
+
+def compute_observations(model, data, commands, last_actions, base_rot):
+ """Compute 57-dim observation matching go1-stairs-terrain-walk-no-linevel.
+
+ Layout (57 dims, NO linear velocity):
+ [0:3] gyro (ang_vel * 0.25)
+ [3:6] gravity vector (body frame)
+ [6:18] joint angle deviation (dof_pos * 1.0)
+ [18:30] joint velocity (dof_vel * 0.05)
+ [30:42] last actions (raw)
+ [42:45] commands [vx*2.0, vy*2.0, wz*0.25]
+ [45:57] foot contact forces (body frame, raw)
+ """
+ obs = np.zeros(NUM_OBS, dtype=np.float32)
+
+ # Gyro
+ gyro = get_sensor(model, data, "gyro")
+ if gyro is not None:
+ obs[0:3] = gyro * OBS_SCALES["ang_vel"]
+ else:
+ obs[0:3] = data.qvel[3:6] * OBS_SCALES["ang_vel"]
+
+ # Gravity vector (body frame)
+ gravity_world = np.array([0.0, 0.0, -1.0], dtype=np.float64)
+ local_gravity = base_rot.T @ gravity_world
+ obs[3:6] = local_gravity.astype(np.float32)
+
+ # Joint position deviation
+ joint_pos = data.qpos[7:19]
+ dof_pos_rel = (joint_pos - DEFAULT_JOINT_ANGLES) * OBS_SCALES["dof_pos"]
+ obs[6:18] = dof_pos_rel
+
+ # Joint velocity
+ joint_vel = data.qvel[6:18]
+ obs[18:30] = joint_vel * OBS_SCALES["dof_vel"]
+
+ # Last actions
+ obs[30:42] = last_actions
+
+ # Commands (scale matching MotrixLab: [2.0, 2.0, 0.25])
+ obs[42] = commands[0] * 2.0
+ obs[43] = commands[1] * 2.0
+ obs[44] = commands[2] * 0.25
+
+ # Contact forces
+ # obs[45:57] = read_contact_forces(model, data, base_rot) # disabled: test with zeros
+ obs[45:57] = np.zeros(12, dtype=np.float32)
+
+ obs = np.clip(obs, -CLIP_OBSERVATIONS, CLIP_OBSERVATIONS)
+ return obs
+
+
+# ============================================================
+# Main
+# ============================================================
+def main():
+ import onnxruntime as ort
+
+ parser = argparse.ArgumentParser(description="MotrixLab Go1 No-Linevel Policy Inference in MuJoCo")
+ parser.add_argument("--onnx", type=str, default=DEFAULT_ONNX_PATH)
+ parser.add_argument(
+ "--terrain", type=str, default="combined",
+ choices=["flat", "rough", "stairs", "combined"],
+ help="Terrain type (combined = flat+rough+stairs in one scene)",
+ )
+ args = parser.parse_args()
+
+ os.chdir(MOTRIX_XML_DIR)
+
+ # Select XML
+ if args.terrain == "combined":
+ xml_file = f"{MOTRIX_XML_DIR}/scene_combined_flat_rough_stairs.xml"
+ elif args.terrain == "rough":
+ xml_file = f"{MOTRIX_XML_DIR}/scene_rough_terrain.xml"
+ elif args.terrain == "stairs":
+ xml_file = f"{MOTRIX_XML_DIR}/scene_stairs_terrain.xml"
+ else:
+ xml_file = f"{MOTRIX_XML_DIR}/scene_motor_actuator.xml"
+
+ with open(xml_file, "r") as f:
+ xml_content = f.read()
+
+ model = mujoco.MjModel.from_xml_string(xml_content)
+ data = mujoco.MjData(model)
+
+ print(f"[INFO] Terrain: {args.terrain}")
+ print(f"[INFO] Model: {model.nbody} bodies, {model.nq} DoF, {model.nu} actuators")
+ print(f"[INFO] Timestep: {model.opt.timestep}")
+
+ # Initial spawn (default to stairs at origin - visible to default camera)
+ current_terrain = args.terrain if args.terrain != "combined" else "stairs"
+ spawn_xyz = TERRAIN_SPAWN.get(current_terrain, TERRAIN_SPAWN["stairs"]).copy()
+ # Start higher so robot drops onto terrain safely
+ spawn_xyz[2] = 1.0
+
+ data.qpos[0:3] = spawn_xyz
+ data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0])
+ data.qpos[7:19] = DEFAULT_JOINT_ANGLES
+ data.qvel[:] = 0.0
+ data.ctrl[:] = 0.0
+ mujoco.mj_forward(model, data)
+
+ # Load ONNX
+ session = ort.InferenceSession(args.onnx, providers=["CPUExecutionProvider"])
+ print(f"[INFO] ONNX loaded: {args.onnx}")
+
+ # Main loop
+ ctrl_dt = 0.01
+ num_steps_per_inference = int(ctrl_dt / model.opt.timestep)
+ print(f"[INFO] Inference every {num_steps_per_inference} sim steps")
+
+ step_count = 0
+ inference_step = 0
+ commands = np.zeros(3, dtype=np.float32)
+ last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+
+ keyboard_reader = KeyboardReader()
+ keyboard_reader.init()
+
+ viewer_handle = viewer.launch_passive(model, data)
+ print("[INFO] Viewer launched!")
+ print("[KEYS] WASD=move, QE=strafe, Space=stop, R=reset, 1/2/3=terrain, Esc=quit")
+
+ loop_start_time = time.time()
+ terrain_changed = False
+
+ while viewer_handle.is_running() and not g_exit_requested:
+ # --- Keyboard input ---
+ x_vel, y_vel, yaw_vel = 0.0, 0.0, 0.0
+
+ if keyboard_reader.is_key_held("w"):
+ x_vel = MAX_LIN_VEL_X
+ elif keyboard_reader.is_key_held("s"):
+ x_vel = -MAX_LIN_VEL_X
+
+ if keyboard_reader.is_key_held("q"):
+ y_vel = MAX_LIN_VEL_Y
+ elif keyboard_reader.is_key_held("e"):
+ y_vel = -MAX_LIN_VEL_Y
+
+ if keyboard_reader.is_key_held("a"):
+ yaw_vel = MAX_ANG_VEL
+ elif keyboard_reader.is_key_held("d"):
+ yaw_vel = -MAX_ANG_VEL
+
+ if keyboard_reader.is_key_pressed("space"):
+ x_vel = y_vel = yaw_vel = 0.0
+
+ # Terrain switching
+ for key, terrain_name in [("1", "flat"), ("2", "rough"), ("3", "stairs")]:
+ if keyboard_reader.is_key_pressed(key):
+ current_terrain = terrain_name
+ terrain_changed = True
+ print(f"[TERRAIN] Switch to: {current_terrain}")
+
+ # Reset
+ if keyboard_reader.is_key_pressed("r") or terrain_changed:
+ if terrain_changed and args.terrain == "combined":
+ spawn_xyz = TERRAIN_SPAWN[current_terrain].copy()
+ elif keyboard_reader.is_key_pressed("r"):
+ spawn_xyz = TERRAIN_SPAWN.get(current_terrain, TERRAIN_SPAWN["flat"]).copy()
+
+ data.qpos[0:3] = spawn_xyz
+ data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0])
+ data.qpos[7:19] = DEFAULT_JOINT_ANGLES
+ data.qvel[:] = 0.0
+ data.ctrl[:] = 0.0
+ last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ mujoco.mj_forward(model, data)
+ terrain_changed = False
+ print(f"[RESET] Terrain={current_terrain}, pos={spawn_xyz}")
+
+ if keyboard_reader.is_key_pressed("escape"):
+ break
+
+ # --- Inference ---
+ if inference_step == 0:
+ commands[0] = x_vel
+ commands[1] = y_vel
+ commands[2] = yaw_vel
+
+ base_rot = data.xmat[1].reshape(3, 3)
+ obs = compute_observations(model, data, commands, last_actions, base_rot)
+
+ action = session.run(None, {"observations": obs.reshape(1, -1).astype(np.float32)})[0][0]
+ action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
+ last_actions = action.copy()
+
+ # --- PD control ---
+ joint_targets = DEFAULT_JOINT_ANGLES + action * ACTION_SCALE
+ current_pos = data.qpos[7:19]
+ current_vel = data.qvel[6:18]
+ torques = KP * (joint_targets - current_pos) - KD * current_vel
+ torques = np.clip(torques, -CLIP_ACTIONS, CLIP_ACTIONS)
+ data.ctrl[:] = torques
+
+ mujoco.mj_step(model, data)
+ viewer_handle.sync()
+
+ expected_time = step_count * ctrl_dt
+ elapsed = time.time() - loop_start_time
+ sleep_time = expected_time - elapsed
+ if sleep_time > 0:
+ time.sleep(sleep_time)
+
+ step_count += 1
+ inference_step = (inference_step + 1) % num_steps_per_inference
+
+ if step_count % 500 == 0:
+ trunk_z = data.qpos[2]
+ print(f"[{step_count}] cmd=({x_vel:.1f},{y_vel:.1f},{yaw_vel:.1f}) "
+ f"z={trunk_z:.3f}m terrain={current_terrain}")
+
+ keyboard_reader.restore()
+ viewer_handle.close()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/go1_sim2sim_mujoco.py b/scripts/go1_sim2sim_mujoco.py
new file mode 100644
index 0000000..2e78893
--- /dev/null
+++ b/scripts/go1_sim2sim_mujoco.py
@@ -0,0 +1,371 @@
+#!/usr/bin/env python3
+"""
+ MOTRIXLAB_UNTRIEE_GO1_SIM2SIM
+ source /opt/mujoco/venv/bin/activate
+ cd /opt/unitree_mujoco
+ python demo/go1_sim2sim_mujoco.py
+
+"""
+
+import numpy as np
+import mujoco
+from mujoco import viewer
+import os
+import threading
+import signal
+import queue
+import argparse
+import time
+
+g_exit_requested = False
+
+def signal_handler(signum, frame):
+ global g_exit_requested
+ g_exit_requested = True
+
+signal.signal(signal.SIGINT, signal_handler)
+
+# ============================================================
+# 配置
+# ============================================================
+_PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+DEFAULT_ONNX_PATH = os.path.join(_PROJECT_DIR, "exports_go1_flat", "policy.onnx")
+MOTRIX_XML_DIR = os.path.join(_PROJECT_DIR, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls")
+XML_PATH = f"{MOTRIX_XML_DIR}/go1_motor_actuator.xml"
+
+TERRAIN = "none"
+
+# ============================================================
+# MotrixLab 参数 (来自 cfg.py)
+# ============================================================
+NUM_OBS = 45 # 去掉线速度观测 (原来是48)
+NUM_ACTIONS = 12
+OBS_SCALES = {'lin_vel': 2.0, 'ang_vel': 0.25, 'dof_pos': 1.0, 'dof_vel': 0.05}
+ACTION_SCALE = 0.05
+KP = 80.0
+KD = 1.0
+CLIP_ACTIONS = 23.7
+CLIP_OBSERVATIONS = 100.0
+MAX_LIN_VEL_X = 1.0
+MAX_LIN_VEL_Y = 1.0
+MAX_ANG_VEL = 1.0 # 匹配训练时的角速度命令范围 [-1.0, 1.0]
+
+# ============================================================
+# 关节名称和顺序
+# ============================================================
+POLICY_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",
+]
+
+DEFAULT_JOINT_ANGLES = np.array([
+ -0.0, 0.9, -1.8, # FR_hip, FR_thigh, FR_calf
+ 0.0, 0.9, -1.8, # FL_hip, FL_thigh, FL_calf
+ -0.0, 0.9, -1.8, # RR_hip, RR_thigh, RR_calf
+ 0.0, 0.9, -1.8, # RL_hip, RL_thigh, RL_calf
+], dtype=np.float32)
+
+MUJOCO_TO_POLICY = np.arange(12, dtype=np.int64)
+POLICY_TO_MUJOCO = np.arange(12, dtype=np.int64)
+
+# ============================================================
+# 键盘输入
+# ============================================================
+from pynput import keyboard
+
+class KeyboardReader:
+ def __init__(self):
+ self._event_queue = queue.Queue()
+ self.running = True
+ self.shared_keys_held = set()
+ self.shared_one_shot = set()
+ self._reader_thread = None
+ self._listener = None
+
+ def _normalize_key(self, key):
+ try:
+ if hasattr(key, 'char') and key.char is not None:
+ return key.char.lower()
+ except:
+ pass
+ key_str = str(key)
+ if key_str == 'Key.esc':
+ return 'escape'
+ elif key_str == 'Key.space':
+ return 'space'
+ elif key_str.startswith('Key.'):
+ return key_str.lower()
+ return key_str.lower()
+
+ def _reader_worker(self):
+ while self.running:
+ try:
+ event_type, key = self._event_queue.get(timeout=0.05)
+ k = self._normalize_key(key)
+ if event_type == 'press':
+ self.shared_keys_held.add(k)
+ self.shared_one_shot.discard(k)
+ elif event_type == 'release':
+ self.shared_keys_held.discard(k)
+ self.shared_one_shot.discard(k)
+ except queue.Empty:
+ pass
+
+ def init(self):
+ def on_press(key):
+ self._event_queue.put(('press', key))
+ def on_release(key):
+ self._event_queue.put(('release', key))
+ try:
+ self._listener = keyboard.Listener(on_press=on_press, on_release=on_release)
+ self._listener.start()
+ self._reader_thread = threading.Thread(target=self._reader_worker, daemon=True)
+ self._reader_thread.start()
+ print("[INFO] 键盘监听已启动")
+ except Exception as e:
+ print(f"[WARN] 无法初始化键盘监听: {e}")
+
+ def is_key_pressed(self, key):
+ k = self._normalize_key(key) if isinstance(key, str) else self._normalize_key(key)
+ if k not in self.shared_keys_held or k in self.shared_one_shot:
+ return False
+ self.shared_one_shot.add(k)
+ return True
+
+ def is_key_held(self, key):
+ k = self._normalize_key(key) if isinstance(key, str) else self._normalize_key(key)
+ return k in self.shared_keys_held
+
+ def restore(self):
+ self.running = False
+ if self._listener:
+ self._listener.stop()
+
+
+# ============================================================
+# Sensor 读取
+# ============================================================
+def get_sensor(model, data, name):
+ sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, name)
+ if sid < 0:
+ return None
+ adr = model.sensor_adr[sid]
+ dim = model.sensor_dim[sid]
+ return data.sensordata[adr:adr + dim].copy()
+
+
+def compute_observations_motrix(model, data, commands, last_actions):
+ """计算 45 维观测 (去掉局部线速度,策略仅靠命令+关节信息+陀螺仪来推理)
+
+ 布局 (45 dims):
+ [0:3] 陀螺仪 (ang_vel * 0.25)
+ [3:6] 重力向量 (躯干坐标系,无缩放)
+ [6:18] 关节位置偏差 (dof_pos * 1.0)
+ [18:30] 关节速度 (dof_vel * 0.05)
+ [30:42] 上一步动作 (原始值)
+ [42:45] 命令 [vx*2.0, vy*2.0, wz*0.25]
+ """
+ obs = np.zeros(NUM_OBS, dtype=np.float32)
+
+ # 陀螺仪
+ gyro = get_sensor(model, data, "gyro")
+ if gyro is not None:
+ obs[0:3] = gyro * OBS_SCALES['ang_vel']
+ else:
+ obs[0:3] = data.qvel[3:6] * OBS_SCALES['ang_vel']
+
+ # 重力向量 (躯干坐标系)
+ base_rot = data.xmat[1].reshape(3, 3)
+ gravity_world = np.array([0., 0., -1.], dtype=np.float64)
+ local_gravity = base_rot.T @ gravity_world
+ obs[3:6] = local_gravity.astype(np.float32)
+
+ # 关节位置偏差
+ joint_pos = data.qpos[7:19]
+ dof_pos_rel = (joint_pos - DEFAULT_JOINT_ANGLES) * OBS_SCALES['dof_pos']
+ obs[6:18] = dof_pos_rel
+
+ # 关节速度
+ joint_vel = data.qvel[6:18]
+ obs[18:30] = joint_vel * OBS_SCALES['dof_vel']
+
+ # 上一步动作
+ obs[30:42] = last_actions
+
+ # 命令
+ obs[42:45] = commands * np.array([OBS_SCALES['lin_vel'], OBS_SCALES['lin_vel'], OBS_SCALES['ang_vel']], dtype=np.float32)
+
+ # 限幅
+ obs = np.clip(obs, -CLIP_OBSERVATIONS, CLIP_OBSERVATIONS)
+ return obs
+
+
+def main():
+ import re
+ import onnxruntime as ort
+
+ parser = argparse.ArgumentParser(description="MotrixLab Go1 Policy Inference in MuJoCo")
+ parser.add_argument("--onnx", type=str, default=DEFAULT_ONNX_PATH)
+ parser.add_argument("--terrain", type=str, default=TERRAIN, choices=["none", "rough", "stairs"])
+ args = parser.parse_args()
+
+ os.chdir(MOTRIX_XML_DIR)
+
+ if args.terrain == "rough":
+ xml_file = f"{MOTRIX_XML_DIR}/scene_rough_terrain.xml"
+ elif args.terrain == "stairs":
+ xml_file = f"{MOTRIX_XML_DIR}/scene_stairs_terrain.xml"
+ else:
+ xml_file = f"{MOTRIX_XML_DIR}/scene_motor_actuator.xml" # flat floor
+
+ with open(xml_file, 'r') as f:
+ xml_content = f.read()
+
+ model = mujoco.MjModel.from_xml_string(xml_content)
+ data = mujoco.MjData(model)
+
+ print(f"[INFO] Model: {model.nbody} bodies, {model.nq} DoF, {model.nu} actuators")
+ print(f"[INFO] MuJoCo timestep: {model.opt.timestep}")
+ print(f"[INFO] 关节顺序 (qpos[7:19]): {[mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_JOINT, i) for i in range(1, 13)]}")
+
+ for sensor_name in ["gyro", "local_linvel"]:
+ sid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_SENSOR, sensor_name)
+ print(f"[SENSOR] {sensor_name}: {'存在' if sid >= 0 else '不存在'}")
+
+ # 初始化
+ data.qpos[0:3] = np.array([0.0, 0.0, 0.42])
+ data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0]) # 默认四元数
+ data.qpos[7:19] = DEFAULT_JOINT_ANGLES
+ data.qvel[:] = 0.0
+ data.ctrl[:] = 0.0
+ mujoco.mj_forward(model, data)
+
+ print(f"[INIT] qpos[2]={data.qpos[2]:.3f}")
+ print(f"[INIT] qpos[7:19]={data.qpos[7:19]}")
+ print(f"[INIT] DEFAULT_JOINT_ANGLES={DEFAULT_JOINT_ANGLES}")
+
+ # 加载onnx
+ session = ort.InferenceSession(args.onnx, providers=['CPUExecutionProvider'])
+ print(f"[INFO] loaded")
+
+ # 主循环
+ ctrl_dt = 0.01 # 100Hz
+ num_steps_per_inference = int(ctrl_dt / model.opt.timestep)
+ print(f"[INFO] 每 {num_steps_per_inference} 步推理一次 ")
+
+ step_count = 0
+ inference_step = 0
+
+ x_vel_cmd = 0.0
+ y_vel_cmd = 0.0
+ yaw_vel_cmd = 0.0
+ commands = np.array([x_vel_cmd, y_vel_cmd, yaw_vel_cmd], dtype=np.float32)
+ last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+
+ # 键盘
+ keyboard_reader = KeyboardReader()
+ keyboard_reader.init()
+
+ view = viewer.launch_passive(model, data)
+ print("[INFO] 已启动!")
+
+ loop_start_time = time.time()
+
+ while view.is_running() and not g_exit_requested:
+ # 键盘命令
+ if keyboard_reader.is_key_pressed(' '):
+ x_vel_cmd = 0.0
+ y_vel_cmd = 0.0
+ yaw_vel_cmd = 0.0
+
+ if keyboard_reader.is_key_held('w'):
+ x_vel_cmd = MAX_LIN_VEL_X
+ elif keyboard_reader.is_key_held('s'):
+ x_vel_cmd = -MAX_LIN_VEL_X
+ else:
+ x_vel_cmd = 0.0
+
+ if keyboard_reader.is_key_held('q'):
+ y_vel_cmd = MAX_LIN_VEL_Y
+ elif keyboard_reader.is_key_held('e'):
+ y_vel_cmd = -MAX_LIN_VEL_Y
+ else:
+ y_vel_cmd = 0.0
+
+ if keyboard_reader.is_key_held('a'):
+ yaw_vel_cmd = MAX_ANG_VEL
+ elif keyboard_reader.is_key_held('d'):
+ yaw_vel_cmd = -MAX_ANG_VEL
+ else:
+ yaw_vel_cmd = 0.0
+
+ if keyboard_reader.is_key_pressed('r'):
+ # 重置机器人到初始位置
+ data.qpos[0:3] = np.array([0.0, 0.0, 0.42])
+ data.qpos[3:7] = np.array([1.0, 0.0, 0.0, 0.0])
+ data.qpos[7:19] = DEFAULT_JOINT_ANGLES
+ data.qvel[:] = 0.0
+ data.ctrl[:] = 0.0
+ last_actions = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ mujoco.mj_forward(model, data)
+ print("[RESET] 机器人已重置")
+
+ if keyboard_reader.is_key_pressed('escape'):
+ break
+
+ # 推理 (每 N 步一次)
+ if inference_step == 0:
+ commands[0] = x_vel_cmd
+ commands[1] = y_vel_cmd
+ commands[2] = yaw_vel_cmd
+
+ obs = compute_observations_motrix(model, data, commands, last_actions)
+
+ # 推理
+ action = session.run(None, {'observations': obs.reshape(1, -1).astype(np.float32)})[0][0]
+ action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
+ last_actions = action.copy()
+
+ # PD 控制
+ # joint_targets = action * action_scale + default_angles
+ joint_targets = DEFAULT_JOINT_ANGLES + action * ACTION_SCALE
+
+ current_pos = data.qpos[7:19]
+ current_vel = data.qvel[6:18]
+ torques = KP * (joint_targets - current_pos) - KD * current_vel
+ torques = np.clip(torques, -CLIP_ACTIONS, CLIP_ACTIONS)
+ data.ctrl[:] = torques
+
+ mujoco.mj_step(model, data)
+ view.sync()
+
+ expected_time = step_count * ctrl_dt
+ elapsed = time.time() - loop_start_time
+ sleep_time = expected_time - elapsed
+ if sleep_time > 0:
+ time.sleep(sleep_time)
+
+ step_count += 1
+ inference_step = (inference_step + 1) % num_steps_per_inference
+
+ if step_count % 200 == 0:
+ trunk_z = data.qpos[2]
+ lin_vel = np.linalg.norm(data.qvel[0:3])
+ print(f"\n========== Step {step_count} ==========")
+ print(f"[CMD] x={x_vel_cmd:.2f}, y={y_vel_cmd:.2f}, yaw={yaw_vel_cmd:.2f}")
+ print(f"[OBS] gyro={obs[0:3]}, grav={obs[3:6]}")
+ print(f"[ACTION] raw={action[:4]}... scaled={action[:4]*ACTION_SCALE}...")
+ print(f"[TARGET] {joint_targets[:4]}...")
+ print(f"[TORQUE] {torques[:4]}...")
+ print(f"[STATE] z={trunk_z:.3f}m, vel={lin_vel:.3f}m/s")
+ print(f"==========================================\n")
+
+ keyboard_reader.restore()
+ view.close()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/play.py b/scripts/play.py
index 47a92cc..f732726 100644
--- a/scripts/play.py
+++ b/scripts/play.py
@@ -36,6 +36,7 @@ _RAND_SEED = flags.DEFINE_bool("rand-seed", False, "Generate random seed")
_RLLIB = flags.DEFINE_string(
"rllib", None, "The RL framework (skrl/rslrl). Auto-discovered from latest training if not specified."
)
+_FORCE_PHASE = flags.DEFINE_integer("force-phase", None, "Lock terrain phase (0=flat,1=rough,2=stairs,3=mixed)")
def get_inference_backend(policy_path: Path | str, rllib: str):
@@ -179,13 +180,21 @@ def main(argv):
backend = get_inference_backend(policy_path, rllib)
+ # Build env config overrides
+ env_cfg_override = {}
+ if _FORCE_PHASE.present:
+ env_cfg_override["force_phase"] = _FORCE_PHASE.value
+ if not env_cfg_override:
+ env_cfg_override = None
+
if rllib == "rslrl":
# RSLRL evaluation flow (always uses torch backend)
assert device_supports.torch, "PyTorch is not available on your device"
from motrix_rl.rslrl.torch.train import ppo
config.torch.backend = "torch"
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
trainer.play(policy_path)
elif backend == "jax":
@@ -193,7 +202,8 @@ def main(argv):
from motrix_rl.skrl.jax.train import ppo
config.jax.backend = "jax" # or "numpy"
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
trainer.play(policy_path)
elif backend == "torch":
@@ -201,7 +211,8 @@ def main(argv):
from motrix_rl.skrl.torch.train import ppo
config.torch.backend = "torch"
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
trainer.play(policy_path)
diff --git a/scripts/play_dreamwaq.py b/scripts/play_dreamwaq.py
new file mode 100644
index 0000000..b05614a
--- /dev/null
+++ b/scripts/play_dreamwaq.py
@@ -0,0 +1,147 @@
+#!/usr/bin/env python3
+"""DreamWaQ play — renders env with trained policy in MotrixSim.
+
+Usage:
+ uv run scripts/play_dreamwaq.py
+ uv run scripts/play_dreamwaq.py --num-envs 16
+"""
+import argparse, os, time, sys
+# CRITICAL: disable JAX GPU memory preallocation BEFORE importing jax.
+# Otherwise JAX grabs 75% of GPU memory and starves the MotrixSim (Vulkan)
+# renderer → "Couldn't get swap chain texture" crash. Must be set first.
+os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false")
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+import motrix_envs.locomotion.go1.dreamwaq # noqa
+import motrix_rl.tasks.go1_dreamwaq # noqa
+import numpy as np
+import jax, jax.numpy as jnp
+import pickle, msgpack
+
+from motrix_envs import registry as env_registry
+from motrix_envs.np.renderer import NpRenderer
+from motrix_rl.skrl.jax.train.dreamwaq_ppo import DreamWaQWrapper, CENet
+
+
+def _decode_arr(ext):
+ 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])
+ return np.frombuffer(parts[2], dtype=np.dtype(parts[1])).reshape(shape)
+
+
+def load_params(ckpt_path):
+ with open(ckpt_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_arr(v) for k, v in val.items()}
+ else:
+ params[name] = _decode_arr(val)
+ # State-preprocessor stats for first 64 dims (REQUIRED: policy trained on normalized obs)
+ mean64 = std64 = None
+ if 'state_preprocessor' in ckpt:
+ sp = msgpack.unpackb(ckpt['state_preprocessor'], raw=False)
+ mean64 = _decode_arr(sp['running_mean'])[:64].astype(np.float32)
+ std64 = np.sqrt(_decode_arr(sp['running_variance'])[:64]).astype(np.float32)
+ return params, mean64, std64
+
+
+CLIP_ACT = 23.7
+CLIP_OBS = 100.0
+
+def policy_forward(x, p, mean64=None, std64=None):
+ x = jnp.array(x[:, :64])
+ # Apply state-preprocessor normalization (clip((x-mean)/(std+eps), -5, 5))
+ if mean64 is not None:
+ x = jnp.clip((x - jnp.array(mean64)) / (jnp.array(std64) + 1e-8), -5.0, 5.0)
+ else:
+ x = jnp.clip(x, -CLIP_OBS, CLIP_OBS)
+ x = jax.nn.elu(x @ jnp.array(p['Dense_0']['kernel']) + jnp.array(p['Dense_0']['bias']))
+ x = jax.nn.elu(x @ jnp.array(p['Dense_1']['kernel']) + jnp.array(p['Dense_1']['bias']))
+ x = jax.nn.elu(x @ jnp.array(p['Dense_2']['kernel']) + jnp.array(p['Dense_2']['bias']))
+ return np.clip(np.array(x @ jnp.array(p['Dense_3']['kernel']) + jnp.array(p['Dense_3']['bias'])),
+ -CLIP_ACT, CLIP_ACT)
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--num-envs", type=int, default=9)
+ p.add_argument("--checkpoint", default=None)
+ args = p.parse_args()
+
+ # Auto-find checkpoint
+ if args.checkpoint is None:
+ run_dir = "runs/go1-dreamwaq-walk/skrl"
+ runs = sorted([d for d in os.listdir(run_dir) if os.path.isdir(os.path.join(run_dir, d)) and d.startswith("26-")])
+ args.checkpoint = os.path.join(run_dir, runs[-1], "checkpoints", "best_agent.pickle")
+
+ # Load policy + state-preprocessor normalization
+ policy_params, mean64, std64 = load_params(args.checkpoint)
+ print(f"[Play] Policy: {args.checkpoint}")
+ print(f"[Play] State normalization: {'ON' if mean64 is not None else 'OFF'}")
+
+ # Load VAE (saved in skrl/ base dir, not run subdir)
+ run_dir = os.path.dirname(os.path.dirname(os.path.dirname(args.checkpoint))) # skrl/ base
+ vae_path = os.path.join(run_dir, "cenet_params.pkl")
+ if not os.path.exists(vae_path):
+ vae_files = sorted([f for f in os.listdir(run_dir) if f.startswith("vae_")],
+ key=lambda x: int(x.split("_")[1].split(".")[0]))
+ if vae_files:
+ vae_path = os.path.join(run_dir, vae_files[-1])
+ with open(vae_path, 'rb') as f:
+ vae_params = pickle.load(f)
+ print(f"[Play] VAE: {vae_path}")
+
+ # Create env + renderer (like view.py)
+ raw_env = env_registry.make("go1-dreamwaq-walk", num_envs=args.num_envs)
+ renderer = NpRenderer(raw_env)
+
+ # CENet for policy inference
+ cenet = CENet()
+ rng = jax.random.PRNGKey(42)
+ wrapper = DreamWaQWrapper(raw_env, cenet, vae_params, rng=rng)
+
+ # Init env
+ raw_env.init_state()
+ wrapper._vae_buf = []
+
+ n = raw_env._num_envs
+ print(f"[Play] {n} envs, Ctrl+C to stop")
+
+ from motrixsim.render import RenderClosedError
+ try:
+ while True:
+ # CENet inference (mean mode)
+ hist = jnp.array(raw_env._state.info.get("obs_history",
+ np.zeros((n, 5, 45), dtype=np.float32)))
+ z, vel = cenet.apply(vae_params, hist, method=cenet.inference)
+ code = np.concatenate([np.array(vel), np.array(z)], axis=-1)
+ obs_arr = raw_env._state.obs
+ priv = raw_env._state.info.get("privileged_obs", np.zeros((n, 235), dtype=np.float32))
+ heights = priv[:, 48:] if priv.shape[1] > 48 else np.zeros((n, 187), dtype=np.float32)
+ base_vel = raw_env._state.info.get("base_vel", np.zeros((n, 3), dtype=np.float32))
+ base_vel_n = base_vel * np.array([2.0, 2.0, 1.0], dtype=np.float32)
+ aug_obs = np.concatenate([code, obs_arr, base_vel_n, heights], axis=-1)
+
+ actions = policy_forward(aug_obs, policy_params, mean64, std64)
+ wrapper.step(actions)
+ renderer.render()
+ time.sleep(0.01)
+ except (KeyboardInterrupt, RenderClosedError):
+ pass
+ try: renderer.close()
+ except: pass
+ print("[Play] Done")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/play_dreamwaq_rsl.py b/scripts/play_dreamwaq_rsl.py
new file mode 100644
index 0000000..972a1ac
--- /dev/null
+++ b/scripts/play_dreamwaq_rsl.py
@@ -0,0 +1,141 @@
+#!/usr/bin/env python3
+"""DreamWaQ rsl_rl play — render the trained ActorCritic_DWAQ policy in NATIVE MotrixSim.
+
+Loads the PyTorch checkpoint DIRECTLY (no ONNX). ONNX is only for cross-sim
+deployment (e.g. MuJoCo sim2sim); the native MotrixSim env runs the torch policy.
+Deterministic inference: mean CENet code + actor.
+
+Usage:
+ uv run scripts/play_dreamwaq_rsl.py # auto-find latest, walk forward
+ uv run scripts/play_dreamwaq_rsl.py --checkpoint runs/.../model_1100.pt --vx 0.5
+ uv run scripts/play_dreamwaq_rsl.py --vx 0 --num-envs 1 # stand still, single robot
+"""
+import argparse, glob, os, sys, time
+
+# avoid JAX grabbing GPU memory and starving the MotrixSim (Vulkan) renderer
+os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false")
+os.environ.setdefault("JAX_PLATFORMS", "cpu")
+# --terrain / --level / --flat-stairs / --stairs: pick hfield scene (before import).
+terrain_type = "pyramid"
+if "--flat-stairs" in sys.argv: terrain_type = "flat_stairs"
+elif "--stairs" in sys.argv: terrain_type = "stairs"
+if "--terrain" in sys.argv or "--level" in sys.argv or "--flat-stairs" in sys.argv or "--stairs" in sys.argv:
+ os.environ["DREAMWAQ_TERRAIN"] = terrain_type
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+import numpy as np
+import torch
+
+import motrix_envs.locomotion.go1.dreamwaq # noqa: F401 register env
+from motrix_envs import registry as env_registry
+from motrix_envs.np.renderer import NpRenderer
+from motrix_rl.dwaq_rsl.actor_critic_dwaq import ActorCritic_DWAQ
+
+PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+NUM_OBS, NUM_PRIV, NUM_HIST, NUM_ACT, CENET_OUT = 45, 235, 5, 12, 19
+CLIP_ACT = 23.7
+
+
+def _iter_of(path):
+ try:
+ return int(os.path.basename(path).split("_")[1].split(".")[0])
+ except Exception:
+ return -1
+
+
+def find_latest():
+ models = glob.glob(os.path.join(PROJECT, "runs", "go1-dreamwaq-walk", "rsl_dwaq", "*", "model_*.pt"))
+ if not models:
+ print("[ERROR] no rsl_dwaq checkpoints found"); sys.exit(1)
+ return max(models, key=_iter_of) # highest iteration (flat model_1100 > terrain early models)
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--checkpoint", default=None)
+ p.add_argument("--num-envs", type=int, default=4)
+ p.add_argument("--vx", type=float, default=0.5, help="forward velocity command [m/s]")
+ p.add_argument("--vy", type=float, default=0.0, help="lateral velocity command [m/s]")
+ p.add_argument("--wz", type=float, default=0.0, help="yaw rate command [rad/s]")
+ p.add_argument("--terrain", action="store_true",
+ help="view the training pyramid hfield (default: flat plane)")
+ p.add_argument("--flat-stairs", action="store_true",
+ help="view the 2-level flat+stairs terrain (implies terrain)")
+ p.add_argument("--stairs", action="store_true",
+ help="view the stairs terrain scene (implies terrain)")
+ p.add_argument("--level", type=int, default=None,
+ help="force ALL spawns at this terrain level (implies --terrain)")
+ p.add_argument("--spawn-height", type=float, default=None,
+ help="spawn clearance above terrain in meters (default 0.45; try 1-2 to experiment)")
+ args = p.parse_args()
+
+ ckpt = args.checkpoint or find_latest()
+ ac = ActorCritic_DWAQ(NUM_OBS + CENET_OUT, NUM_PRIV, NUM_ACT, NUM_HIST * NUM_OBS, CENET_OUT)
+ ac.load_state_dict(torch.load(ckpt, map_location="cpu")["model_state_dict"])
+ ac.eval()
+ print(f"[Play-rsl] policy (native torch): {ckpt}")
+
+ env = env_registry.make("go1-dreamwaq-walk", num_envs=args.num_envs)
+ if args.level is not None:
+ env._force_level = args.level # pin all spawns to this level (read in reset)
+ print(f"[Play-rsl] forcing ALL spawns at terrain level {args.level}")
+ if args.spawn_height is not None:
+ env._spawn_absolute = args.spawn_height # absolute world z, no offset
+ print(f"[Play-rsl] spawn absolute z = {args.spawn_height}m")
+ renderer = NpRenderer(env)
+ env.init_state()
+ n = env._num_envs
+ cmd = np.array([args.vx, args.vy, args.wz], dtype=np.float32)
+ print(f"[Play-rsl] {n} envs | cmd=(vx={args.vx}, vy={args.vy}, wz={args.wz}) | Ctrl+C to stop")
+
+ @torch.no_grad()
+ def act_fn(obs, hist):
+ obs_t = torch.from_numpy(obs)
+ h = ac.encoder(torch.from_numpy(hist).reshape(obs.shape[0], -1)) # (n,225)->(n,64)
+ code = torch.cat([ac.encode_mean_vel(h), ac.encode_mean_latent(h)], dim=-1) # (n,19)
+ return ac.actor(torch.cat([code, obs_t], dim=-1)).numpy() # (n,12)
+
+ from motrixsim.render import RenderClosedError
+ show_heights = False
+ try:
+ while True:
+ if renderer._render.input.is_key_just_pressed("r"):
+ env.init_state()
+ print("[R] Reset all envs")
+ if renderer._render.input.is_key_just_pressed("h"):
+ show_heights = not show_heights
+ print(f"[H] Height points: {'ON' if show_heights else 'OFF'}")
+ env._state.info["commands"][:] = cmd
+ obs = env._state.obs.astype(np.float32)
+ hist = env._state.info.get("obs_history",
+ np.zeros((n, NUM_HIST, NUM_OBS), np.float32)).astype(np.float32)
+ act = act_fn(obs, hist)
+ env.step(np.clip(act, -CLIP_ACT, CLIP_ACT).astype(np.float32))
+
+ if show_heights:
+ from motrix_envs.math import quaternion
+ pose = env._body.get_pose(env._state.data)
+ bp = pose[0, :3]
+ yaw = quaternion.get_yaw(pose[0:1, 3:7])[0]
+ cos_y, sin_y = np.cos(yaw), np.sin(yaw)
+ for gy in env._hy:
+ for gx in env._hx:
+ wx = bp[0] + cos_y*gx - sin_y*gy
+ wy = bp[1] + sin_y*gx + cos_y*gy
+ wz = float(env._sample_terrain_height(np.array([[wx,wy]]))[0])
+ g = renderer._render.gizmos
+ g.draw_sphere(0.02, (np.float32(wx), np.float32(wy), np.float32(wz)))
+
+ renderer.render()
+ time.sleep(0.01)
+ except (KeyboardInterrupt, RenderClosedError):
+ pass
+ try:
+ renderer.close()
+ except Exception:
+ pass
+ print("[Play-rsl] done")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/terrain_editor.py b/scripts/terrain_editor.py
new file mode 100644
index 0000000..cc5e3d4
--- /dev/null
+++ b/scripts/terrain_editor.py
@@ -0,0 +1,430 @@
+#!/usr/bin/env python3
+"""Terrain editor GUI — draw pyramids, mark spawn zones, export PNG + coordinates.
+
+Usage:
+ uv run scripts/terrain_editor.py
+"""
+import tkinter as tk
+from tkinter import ttk, messagebox
+import numpy as np, os, cv2
+
+# ═══ defaults ═══
+HS = 0.05; VS = 0.005
+CELL_M = 8.0; BORDER_M = 5.0
+PLATFORM_M = 1.0 # platform 1m
+SPAWN_RADIUS_M = 0.5 # spawn zone ±0.5m around center
+DEFAULT_STEP_H = 0.20; DEFAULT_STEP_D = 0.20; DEFAULT_NUM_STEPS = 10
+DEFAULT_REF_PLANE_CM = 200 # all cells start from same reference height
+CELL_PX = int(CELL_M / HS); BORDER_PX = int(BORDER_M / HS)
+PLATFORM_PX = int(PLATFORM_M / HS); SPAWN_RADIUS_PX = int(SPAWN_RADIUS_M / HS)
+
+
+class TerrainEditor:
+ def __init__(self, root):
+ self.root = root
+ self.root.title("Terrain Editor")
+ self.rows = 2; self.cols = 4
+ self.cell_types = {}
+ self._init_defaults()
+ self.selected = (0, 0)
+ self._dragging = False
+ self._build_ui()
+ self._sync_params()
+ self._redraw_all()
+
+ def _init_defaults(self):
+ for r in range(self.rows):
+ for c in range(self.cols):
+ if r == 0:
+ self.cell_types[(r, c)] = {"type": "flat", "spawn": True, "level": 0,
+ "ref_plane_cm": DEFAULT_REF_PLANE_CM}
+ else:
+ self.cell_types[(r, c)] = {
+ "type": "convex" if c % 2 == 0 else "concave",
+ "step_h": DEFAULT_STEP_H, "step_d": DEFAULT_STEP_D,
+ "num_steps": DEFAULT_NUM_STEPS, "spawn": True, "level": 1,
+ "ref_plane_cm": DEFAULT_REF_PLANE_CM,
+ }
+
+ # ═══ UI ═══
+ def _build_ui(self):
+ paned = ttk.PanedWindow(self.root, orient=tk.HORIZONTAL)
+ paned.pack(fill=tk.BOTH, expand=True)
+ left = ttk.Frame(paned); paned.add(left, weight=2)
+ right = ttk.Frame(paned); paned.add(right, weight=1)
+ self._build_preview_ui(left)
+ self._build_params_ui(right)
+
+ def _build_preview_ui(self, parent):
+ ttk.Label(parent, text="Terrain Preview (click to select cell, right-click toggle spawn)", font=("", 10)).pack(pady=2)
+ self.info_label = ttk.Label(parent, text="")
+ self.info_label.pack()
+ self.preview = tk.Canvas(parent, bg="#333", width=600, height=400)
+ self.preview.pack(fill=tk.BOTH, expand=True, padx=5, pady=5)
+ self.preview.bind("", self._on_click)
+ self.preview.bind("", self._on_drag)
+ self.preview.bind("", self._on_right_click)
+ ttk.Label(parent, text="Left-click: select | Right-click: toggle spawn | Drag: select").pack()
+ ctrl = ttk.Frame(parent)
+ ctrl.pack(pady=5)
+ ttk.Label(ctrl, text="Rows:").pack(side=tk.LEFT)
+ self.rows_var = tk.IntVar(value=self.rows)
+ ttk.Spinbox(ctrl, from_=1, to=10, width=4, textvariable=self.rows_var,
+ command=self._on_grid_size).pack(side=tk.LEFT, padx=2)
+ ttk.Label(ctrl, text="Cols:").pack(side=tk.LEFT, padx=(10,0))
+ self.cols_var = tk.IntVar(value=self.cols)
+ ttk.Spinbox(ctrl, from_=1, to=10, width=4, textvariable=self.cols_var,
+ command=self._on_grid_size).pack(side=tk.LEFT, padx=2)
+ ttk.Button(parent, text="Export PNG + Coords", command=self._export).pack(pady=5)
+
+ def _build_params_ui(self, parent):
+ f = ttk.Frame(parent); f.pack(padx=10, pady=5, fill=tk.X)
+ ttk.Label(f, text="Cell Type:").grid(row=0, column=0, sticky=tk.W)
+ self.type_var = tk.StringVar(value="flat")
+ ttk.Combobox(f, textvariable=self.type_var, values=["flat", "convex", "concave"],
+ state="readonly", width=10).grid(row=0, column=1, padx=5)
+ self.type_var.trace("w", lambda *a: self._on_param_change())
+ ttk.Label(f, text="Level:").grid(row=0, column=2, sticky=tk.W, padx=(20,0))
+ self.level_var = tk.IntVar(value=0)
+ ttk.Spinbox(f, from_=0, to=9, width=3, textvariable=self.level_var,
+ command=self._on_param_change).grid(row=0, column=3)
+ self.spawn_var = tk.BooleanVar(value=True)
+ ttk.Checkbutton(f, text="Spawn", variable=self.spawn_var,
+ command=self._on_param_change).grid(row=0, column=4, padx=10)
+
+ ttk.Separator(parent, orient=tk.HORIZONTAL).pack(fill=tk.X, pady=5, padx=10)
+ ttk.Label(parent, text="Pyramid Params").pack()
+ g = ttk.Frame(parent); g.pack(padx=10, pady=5, fill=tk.X)
+ ttk.Label(g, text="Step height (m):").grid(row=0, column=0, sticky=tk.W)
+ self.sh_var = tk.StringVar(value=str(DEFAULT_STEP_H))
+ ttk.Entry(g, textvariable=self.sh_var, width=7).grid(row=0, column=1, padx=5)
+ self.sh_var.trace("w", lambda *a: self._on_param_change())
+ ttk.Label(g, text="Step tread (m):").grid(row=1, column=0, sticky=tk.W)
+ self.sd_var = tk.StringVar(value=str(DEFAULT_STEP_D))
+ ttk.Entry(g, textvariable=self.sd_var, width=7).grid(row=1, column=1, padx=5)
+ self.sd_var.trace("w", lambda *a: self._on_param_change())
+ ttk.Label(g, text="Num steps:").grid(row=2, column=0, sticky=tk.W)
+ self.ns_var = tk.StringVar(value=str(DEFAULT_NUM_STEPS))
+ ttk.Entry(g, textvariable=self.ns_var, width=7).grid(row=2, column=1, padx=5)
+ ttk.Label(g, text="Ref plane (cm):").grid(row=3, column=0, sticky=tk.W)
+ self.ref_var = tk.StringVar(value=str(DEFAULT_REF_PLANE_CM))
+ ttk.Entry(g, textvariable=self.ref_var, width=7).grid(row=3, column=1, padx=5)
+ self.ns_var.trace("w", lambda *a: self._on_param_change())
+
+ ttk.Separator(parent, orient=tk.HORIZONTAL).pack(fill=tk.X, pady=5, padx=10)
+ ttk.Label(parent, text="Selected Cell").pack()
+ self.cell_label = ttk.Label(parent, text="")
+ self.cell_label.pack()
+
+ # ═══ events ═══
+ def _cell_at(self, ex, ey):
+ m = 5; pw = self.preview.winfo_width(); ph = self.preview.winfo_height()
+ cw = (pw - 2*m) // max(self.cols,1); ch = (ph - 2*m) // max(self.rows,1)
+ col = (ex - m) // cw; row = (ey - m) // ch
+ if 0 <= col < self.cols and 0 <= row < self.rows:
+ return row, col, m + col*cw, m + row*ch, cw, ch
+ return None
+
+ def _on_click(self, event):
+ v = self._cell_at(event.x, event.y)
+ if v:
+ self.selected = (v[0], v[1])
+ self._sync_params(); self._redraw_all()
+
+ def _on_drag(self, event):
+ v = self._cell_at(event.x, event.y)
+ if v:
+ self.selected = (v[0], v[1])
+ self._sync_params(); self._redraw_all()
+
+ def _on_right_click(self, event):
+ v = self._cell_at(event.x, event.y)
+ if v:
+ r, c = v[0], v[1]
+ ct = self.cell_types.get((r, c), {"type": "flat", "spawn": True, "level": r})
+ ct = dict(ct) # copy before modifying
+ ct["spawn"] = not ct.get("spawn", True)
+ self.cell_types[(r, c)] = ct
+ if (r, c) == self.selected:
+ self._sync_params()
+ self._redraw_all()
+
+ def _on_grid_size(self):
+ try: nr = self.rows_var.get()
+ except: nr = self.rows
+ try: nc = self.cols_var.get()
+ except: nc = self.cols
+ if nr == self.rows and nc == self.cols: return
+ old = self.cell_types
+ self.rows, self.cols = nr, nc
+ self.cell_types = {}
+ for r in range(nr):
+ for c in range(nc):
+ self.cell_types[(r,c)] = old.get((r,c), {"type": "flat", "spawn": True, "level": r})
+ self._redraw_all()
+
+ # ═══ sync ═══
+ def _sync_params(self):
+ ct = self.cell_types.get(self.selected, {"type": "flat", "spawn": True, "level": self.selected[0]})
+ self.type_var.set(ct.get("type", "flat"))
+ self.spawn_var.set(ct.get("spawn", True))
+ r = self.selected[0]
+ self.level_var.set(ct.get("level", r))
+ self.ref_var.set(str(ct.get("ref_plane_cm", DEFAULT_REF_PLANE_CM)))
+ self.sh_var.set(str(ct.get("step_h", DEFAULT_STEP_H)))
+ self.sd_var.set(str(ct.get("step_d", DEFAULT_STEP_D)))
+ self.ns_var.set(str(ct.get("num_steps", DEFAULT_NUM_STEPS)))
+ r, c = self.selected
+ half_x = BORDER_M + self.cols * CELL_M / 2
+ half_y = BORDER_M + self.rows * CELL_M / 2
+ cx = -half_x + BORDER_M + c * CELL_M + CELL_M / 2
+ cy = half_y - BORDER_M - r * CELL_M - CELL_M / 2
+ self.cell_label.config(text=f"({r},{c}) center: x={cx:+.1f} y={cy:+.1f} type={ct['type']}")
+
+ def _on_param_change(self):
+ r, c = self.selected
+ try: sh = float(self.sh_var.get()); sd = float(self.sd_var.get())
+ except: return
+ try: ns = int(self.ns_var.get())
+ except: return
+ try: ref_cm = float(self.ref_var.get())
+ except: ref_cm = DEFAULT_REF_PLANE_CM
+ cell = {"type": self.type_var.get(), "spawn": self.spawn_var.get(),
+ "level": self.level_var.get(), "ref_plane_cm": ref_cm}
+ if cell["type"] != "flat":
+ cell.update({"step_h": sh, "step_d": sd, "num_steps": ns})
+ self.cell_types[(r, c)] = cell
+ self._redraw_all()
+
+ # ═══ draw ═══
+ def _redraw_all(self):
+ w = self.preview.winfo_width(); h = self.preview.winfo_height()
+ if w < 10: w = 600
+ if h < 10: h = 400
+ self._draw_preview(w, h)
+ half_x = BORDER_M + self.cols * CELL_M / 2
+ half_y = BORDER_M + self.rows * CELL_M / 2
+ self.info_label.config(
+ text=f"{self.rows}×{self.cols} "
+ f"{self.cols*CELL_M+2*BORDER_M:.0f}×{self.rows*CELL_M+2*BORDER_M:.0f}m "
+ f"spawn_cy = {half_y-BORDER_M-CELL_M/2:.0f} - row*{CELL_M:.0f}")
+
+ def _draw_preview(self, pw, ph):
+ cv = self.preview; cv.delete("all")
+ m = 5; cw = (pw - 2*m) // max(self.cols, 1); ch = (ph - 2*m) // max(self.rows, 1)
+ cw = max(cw, 30); ch = max(ch, 30)
+ colors = {"flat": "#5b8c5a", "convex": "#c0392b", "concave": "#2471a3"}
+
+ for r in range(self.rows):
+ for c in range(self.cols):
+ x1, y1 = m + c*cw, m + r*ch
+ x2, y2 = x1 + cw, y1 + ch
+ ct = self.cell_types.get((r,c), {"type": "flat", "spawn": True, "level": r})
+ cv.create_rectangle(x1, y1, x2, y2, fill=colors.get(ct["type"], "#555"),
+ outline="#888", width=1)
+ # Cell center dot
+ cx = (x1+x2)//2; cy = (y1+y2)//2
+ cv.create_oval(cx-3, cy-3, cx+3, cy+3, fill="white", outline="")
+
+ # Pyramid stairs rings
+ if ct["type"] != "flat":
+ sh = ct.get("step_h", DEFAULT_STEP_H)
+ sd = ct.get("step_d", DEFAULT_STEP_D)
+ ns = ct.get("num_steps", DEFAULT_NUM_STEPS)
+ concave = ct["type"] == "concave"
+ p2 = max(2, cw // 16)
+ step_px = max(1, (cw//2 - p2) // max(ns, 1))
+ h_max = int(sh * ns / VS)
+ for i in range(ns + 1):
+ half = p2 + (ns - i) * step_px
+ if concave:
+ frac = (ns - i) / max(ns, 1)
+ else:
+ frac = i / max(ns, 1)
+ g = int(180 - frac * 100)
+ clr = f"#{g:02x}{g:02x}{g:02x}"
+ cv.create_rectangle(cx - half, cy - half, cx + half, cy + half,
+ fill=clr, outline="")
+
+ # Spawn zone (green rect)
+ if ct.get("spawn", True):
+ sz = max(2, int(SPAWN_RADIUS_M / CELL_M * cw))
+ cv.create_rectangle(cx - sz, cy - sz, cx + sz, cy + sz,
+ outline="#00ff00", width=2)
+
+ # Level + type label
+ lvl = ct.get("level", r)
+ lbl = f"L{lvl} {ct['type'][:3]}"
+ if ct["type"] != "flat":
+ lbl = f"L{lvl} {ct['type'][:3]}-{sh*100:.0f}cm"
+ cv.create_text(x1 + 20, y1 + 10, text=lbl, fill="white",
+ font=("", 8), anchor=tk.NW)
+
+ # Highlight selected cell
+ r, c = self.selected
+ x1, y1 = m + c*cw, m + r*ch
+ x2, y2 = x1 + cw, y1 + ch
+ cv.create_rectangle(x1, y1, x2, y2, outline="yellow", width=3)
+
+ # Level labels on right
+ for r in range(self.rows):
+ y = m + r*ch + ch//2
+ cv.create_text(pw - 15, y, text=f"L{r}", fill="white", font=("", 12, "bold"))
+
+ # ═══ generate + export ═══
+ def _generate_png(self):
+ tot_rows = self.rows * CELL_PX + 2 * BORDER_PX
+ tot_cols = self.cols * CELL_PX + 2 * BORDER_PX
+ canvas = np.zeros((tot_rows, tot_cols), dtype=np.uint16)
+ for r in range(self.rows):
+ for c in range(self.cols):
+ x0 = BORDER_PX + c * CELL_PX; y0 = BORDER_PX + r * CELL_PX
+ ct = self.cell_types.get((r,c), {"type": "flat", "level": r})
+ if ct["type"] == "flat": continue
+ sh = ct.get("step_h", DEFAULT_STEP_H)
+ sd = ct.get("step_d", DEFAULT_STEP_D)
+ ns = ct.get("num_steps", DEFAULT_NUM_STEPS)
+ concave = ct["type"] == "concave"
+ ref_cm = ct.get("ref_plane_cm", DEFAULT_REF_PLANE_CM)
+ ref_vs = int(ref_cm / 100.0 / VS)
+ h_vs = int(sh / VS); d_px = int(sd / HS)
+ p2 = PLATFORM_PX // 2
+ cx = x0 + CELL_PX // 2; cy = y0 + CELL_PX // 2
+ cv2.rectangle(canvas, (x0, y0), (x0+CELL_PX, y0+CELL_PX), int(ref_vs), -1)
+ for i in range(ns + 1):
+ half = p2 + (ns - i) * d_px
+ x1, y1 = cx - half, cy - half; x2, y2 = cx + half, cy + half
+ if concave:
+ h = ref_vs - h_vs * i
+ else:
+ h = ref_vs + h_vs * i
+ cv2.rectangle(canvas, (x1, y1), (x2, y2), int(h), -1)
+ hf_m = canvas.astype(np.float32) * VS
+ z_min, z_max = float(hf_m.min()), float(hf_m.max())
+ z_range = max(z_max - z_min, 0.001)
+ png = ((hf_m - z_min) / z_range * 65535).astype(np.uint16)
+ return png, z_range, z_min
+
+ def _export(self):
+ png, z_range, z_min = self._generate_png()
+ out_dir = os.path.join(os.path.dirname(__file__), "..",
+ "motrix_envs", "src", "motrix_envs", "locomotion",
+ "go1", "xmls", "assets")
+ os.makedirs(out_dir, exist_ok=True)
+ out_path = os.path.join(out_dir, "flat_stairs.png")
+ cv2.imwrite(out_path, png)
+ w_m = (self.cols * CELL_PX + 2 * BORDER_PX) * HS
+ h_m = (self.rows * CELL_PX + 2 * BORDER_PX) * HS
+ half_x = BORDER_M + self.cols * CELL_M / 2
+ half_y = BORDER_M + self.rows * CELL_M / 2
+
+ lines = [
+ f"# Terrain: {self.rows}×{self.cols} {w_m:.0f}×{h_m:.0f}m",
+ f"XML: size=\"{w_m/2:.1f} {h_m/2:.1f} {z_range:.3f} {max(z_min,0.001):.3f}\"",
+ f"dreamwaq.py: terrain_rows={self.rows} terrain_cols={self.cols}",
+ f"",
+ f"# === Cell centers ===",
+ ]
+ for r in range(self.rows):
+ for c in range(self.cols):
+ cx = -half_x + BORDER_M + c * CELL_M + CELL_M / 2
+ cy = half_y - BORDER_M - r * CELL_M - CELL_M / 2
+ ct = self.cell_types.get((r,c), {"type": "flat", "level": r})
+ lines.append(f" ({r},{c}): x={cx:+.1f} y={cy:+.1f} {ct['type']}")
+
+ lines.append("")
+ lines.append("# === Level boundaries (robot out of bounds → reset) ===")
+ level_bounds = {}
+ for r in range(self.rows):
+ for c in range(self.cols):
+ lv = self.cell_types.get((r,c), {"level": r})["level"]
+ if lv not in level_bounds:
+ level_bounds[lv] = {"rmin": r, "rmax": r, "cmin": c, "cmax": c}
+ else:
+ b = level_bounds[lv]
+ b["rmin"] = min(b["rmin"], r)
+ b["rmax"] = max(b["rmax"], r)
+ b["cmin"] = min(b["cmin"], c)
+ b["cmax"] = max(b["cmax"], c)
+ for lv in sorted(level_bounds):
+ b = level_bounds[lv]
+ x_min = -half_x + BORDER_M + b["cmin"] * CELL_M
+ x_max = -half_x + BORDER_M + (b["cmax"] + 1) * CELL_M
+ y_min = half_y - BORDER_M - (b["rmax"] + 1) * CELL_M
+ y_max = half_y - BORDER_M - b["rmin"] * CELL_M
+ lines.append(f" level {lv}: x=[{x_min:+.1f}, {x_max:+.1f}] "
+ f"y=[{y_min:+.1f}, {y_max:+.1f}] "
+ f"({b['rmax']-b['rmin']+1}×{b['cmax']-b['cmin']+1} cells)")
+
+ lines.append("")
+ lines.append("# === Spawn positions ===")
+ for lv in sorted(level_bounds):
+ b = level_bounds[lv]
+ spawn_cells = [(r,c) for r in range(b["rmin"], b["rmax"]+1)
+ for c in range(b["cmin"], b["cmax"]+1)
+ if self.cell_types.get((r,c), {}).get("spawn", True)]
+ if spawn_cells:
+ lines.append(f" level {lv}: {len(spawn_cells)} spawn cells")
+ for (rr, cc) in spawn_cells:
+ cx = -half_x + BORDER_M + cc * CELL_M + CELL_M / 2
+ cy = half_y - BORDER_M - rr * CELL_M - CELL_M / 2
+ ct = self.cell_types.get((rr,cc), {})
+ if ct.get("type") == "flat":
+ z_plat = 0
+ else:
+ ref_cm = ct.get("ref_plane_cm", DEFAULT_REF_PLANE_CM)
+ sh = ct.get("step_h", 0)
+ ns = ct.get("num_steps", 0)
+ concave = ct["type"] == "concave"
+ z_plat = (ref_cm - ns * sh * 100) / 100.0 if concave else (ref_cm + ns * sh * 100) / 100.0
+ lines.append(f" ({rr},{cc}) x={cx:+.1f} y={cy:+.1f} {ct['type']} "
+ f"z_plat={z_plat*100:.0f}cm")
+
+ lines.append("")
+ lines.append("# === Pyramid tread details ===")
+ for r in range(self.rows):
+ for c in range(self.cols):
+ ct = self.cell_types.get((r,c), {"type": "flat"})
+ if ct["type"] == "flat":
+ continue
+ cx = -half_x + BORDER_M + c * CELL_M + CELL_M / 2
+ cy = half_y - BORDER_M - r * CELL_M - CELL_M / 2
+ sh = ct.get("step_h", DEFAULT_STEP_H)
+ sd = ct.get("step_d", DEFAULT_STEP_D)
+ ns = ct.get("num_steps", DEFAULT_NUM_STEPS)
+ concave = ct["type"] == "concave"
+ ref_cm = ct.get("ref_plane_cm", DEFAULT_REF_PLANE_CM)
+ ref_z = ref_cm / 100.0
+ h_vs = int(sh / VS)
+ p2 = PLATFORM_PX // 2
+ d_px = int(sd / HS)
+ plat_z = (ref_cm - ns * sh * 100) / 100.0 if concave else (ref_cm + ns * sh * 100) / 100.0
+ lines.append(f" ({r},{c}) {ct['type']} center=({cx:+.1f}, {cy:+.1f}) "
+ f"ref_plane={ref_z*100:.0f}cm platform={plat_z*100:.0f}cm "
+ f"step_h={sh*100:.0f}cm tread={sd*100:.0f}cm steps={ns}")
+ for i in range(ns + 1):
+ half_m = (p2 + (ns - i) * d_px) * HS
+ if concave:
+ z = ref_z - (h_vs * i) * VS
+ else:
+ z = ref_z + (h_vs * i) * VS
+ ring_type = "platform" if i == ns else "ring"
+ lines.append(f" {ring_type} {i}: z={z*100:5.0f}cm "
+ f"half={half_m:.2f}m "
+ f"x=[{cx-half_m:+.1f},{cx+half_m:+.1f}] "
+ f"y=[{cy-half_m:+.1f},{cy+half_m:+.1f}]")
+
+ lines.append("")
+ lines.append(f"# Training: DREAMWAQ_TERRAIN=flat_stairs "
+ f"uv run scripts/train_dreamwaq_rsl.py --level N")
+
+ info = "\n".join(lines)
+ print(info)
+ messagebox.showinfo("Exported", f"{out_path}\n\n{info}")
+
+
+if __name__ == "__main__":
+ root = tk.Tk()
+ root.geometry("900x550")
+ TerrainEditor(root)
+ root.mainloop()
diff --git a/scripts/train.py b/scripts/train.py
index 0ad21ac..26e559f 100644
--- a/scripts/train.py
+++ b/scripts/train.py
@@ -35,6 +35,9 @@ _TRAIN_BACKEND = flags.DEFINE_string("train-backend", None, "The learning backen
_SEED = flags.DEFINE_integer("seed", None, "Random seed for reproducibility")
_RAND_SEED = flags.DEFINE_bool("rand-seed", False, "Generate random seed")
_RLLIB = flags.DEFINE_string("rllib", "skrl", "The RL framework (skrl/rslrl)")
+_CHECKPOINT = flags.DEFINE_string("checkpoint", None, "Resume training from a checkpoint (.pickle/.pt)")
+_FORCE_PHASE = flags.DEFINE_integer("force-phase", None, "Lock curriculum to a specific phase (0=flat,1=rough,2=stairs,3=mixed)")
+_TRACKING_LINVEL_SCALE = flags.DEFINE_float("tracking-linvel-scale", None, "Override tracking_lin_vel reward scale")
def get_train_backend(supports: utils.DeviceSupports, train_backend_arg: str | None, rllib: str):
@@ -104,6 +107,15 @@ def main(argv):
# Determine the training backend
train_backend = get_train_backend(device_supports, _TRAIN_BACKEND.value, rllib)
+ # Build env config overrides from command-line flags
+ env_cfg_override = {}
+ if _FORCE_PHASE.present:
+ env_cfg_override["force_phase"] = _FORCE_PHASE.value
+ if _TRACKING_LINVEL_SCALE.present:
+ env_cfg_override["tracking_lin_vel_scale"] = _TRACKING_LINVEL_SCALE.value
+ if not env_cfg_override:
+ env_cfg_override = None
+
trainer = None
if rllib == "rslrl":
# RSLRL training flow
@@ -111,22 +123,25 @@ def main(argv):
assert train_backend == "torch", "RSLRL only supports PyTorch backend"
from motrix_rl.rslrl.torch.train import ppo
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
elif train_backend == "jax":
from motrix_rl.skrl.jax.train import ppo
config.jax.backend = "jax" # or "numpy"
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
elif train_backend == "torch":
from motrix_rl.skrl.torch.train import ppo
- trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
+ enable_render=enable_render, env_cfg_override=env_cfg_override)
else:
raise Exception(f"Unknown train backend: {train_backend}")
- trainer.train()
+ trainer.train(checkpoint=_CHECKPOINT.value)
if __name__ == "__main__":
diff --git a/scripts/train_cts.py b/scripts/train_cts.py
new file mode 100644
index 0000000..58ab914
--- /dev/null
+++ b/scripts/train_cts.py
@@ -0,0 +1,42 @@
+#!/usr/bin/env python3
+"""Train CTS (Concurrent Teacher-Student) Go1 locomotion.
+
+Usage:
+ uv run scripts/train_cts.py
+ uv run scripts/train_cts.py --num-envs 512
+"""
+import logging
+import motrix_rl.tasks.go1_go2style # noqa: triggers env + rlcfg registration
+
+from absl import app, flags
+from skrl import config as skrl_config
+
+from motrix_rl import utils
+from motrix_rl.skrl.jax.train.cts_ppo import CTSTrainer
+
+logger = logging.getLogger(__name__)
+
+_ENV = flags.DEFINE_string("env", "go1-cts-flat-walk-go2style", "CTS env to train")
+_NUM_ENVS = flags.DEFINE_integer("num-envs", 1024, "Number of environments")
+_SEED = flags.DEFINE_integer("seed", None, "Random seed")
+
+
+def main(argv):
+ supports = utils.get_device_supports()
+ logger.info(supports)
+
+ env_name = _ENV.value
+ override = {}
+ if _NUM_ENVS.present:
+ override["num_envs"] = _NUM_ENVS.value
+ if _SEED.present:
+ override["runner.seed"] = _SEED.value
+
+ skrl_config.jax.backend = "jax"
+
+ trainer = CTSTrainer(env_name=env_name, cfg_override=override)
+ trainer.train()
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/scripts/train_dreamwaq.py b/scripts/train_dreamwaq.py
new file mode 100644
index 0000000..5e947d2
--- /dev/null
+++ b/scripts/train_dreamwaq.py
@@ -0,0 +1,42 @@
+#!/usr/bin/env python3
+"""DreamWaQ training — Manaro-Alpha aligned.
+
+Usage:
+ uv run scripts/train_dreamwaq.py # default (2048 envs, 100M steps)
+ uv run scripts/train_dreamwaq.py --num-envs 4096 --timesteps 150M
+"""
+import argparse
+
+# Register env + config
+import motrix_envs.locomotion.go1.dreamwaq # noqa
+import motrix_rl.tasks.go1_dreamwaq # noqa
+from motrix_rl.skrl.jax.train.dreamwaq_ppo import DreamWaQTrainer
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--num-envs", type=int, default=2048)
+ p.add_argument("--timesteps", type=str, default="100M")
+ p.add_argument("--seed", type=int, default=42)
+ args = p.parse_args()
+
+ ts = args.timesteps
+ if ts.endswith("M"): ts = int(float(ts[:-1]) * 1_000_000)
+ elif ts.endswith("K"): ts = int(float(ts[:-1]) * 1_000)
+ else: ts = int(ts)
+
+ # SKRL timesteps = env.step() calls, NOT individual env steps
+ skrl_ts = ts // args.num_envs
+
+ override = {
+ "num_envs": args.num_envs,
+ "runner.seed": args.seed,
+ "runner.trainer.timesteps": skrl_ts,
+ }
+
+ trainer = DreamWaQTrainer(env_name="go1-dreamwaq-walk", cfg_override=override)
+ trainer.train()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/train_dreamwaq_rsl.py b/scripts/train_dreamwaq_rsl.py
new file mode 100644
index 0000000..15301da
--- /dev/null
+++ b/scripts/train_dreamwaq_rsl.py
@@ -0,0 +1,122 @@
+#!/usr/bin/env python3
+"""DreamWaQ training via the faithful rsl_rl-1.0.2 port (upstream-aligned).
+
+Uses upstream's ActorCritic_DWAQ + PPO (joint VAE training) + OnPolicyRunner,
+with MotrixLab's DreamWaQ env. Config matches upstream Go1RoughCfgPPO.
+
+Usage:
+ uv run scripts/train_dreamwaq_rsl.py --num-envs 2048 --iterations 3000
+"""
+import argparse
+import os
+import sys
+import torch
+
+import motrix_envs.locomotion.go1.dreamwaq # noqa: F401 (registers env)
+from motrix_envs import registry as env_registry
+from motrix_rl.dwaq_rsl import OnPolicyRunner, DwaqVecEnv
+
+
+# Upstream Go1RoughCfgPPO (legged_robot_config.py LeggedRobotCfgPPO + Go1 overrides)
+TRAIN_CFG = {
+ "runner": {
+ "policy_class_name": "ActorCritic_DWAQ",
+ "algorithm_class_name": "PPO",
+ "num_steps_per_env": 24,
+ "save_interval": 50,
+ },
+ "algorithm": {
+ "value_loss_coef": 1.0,
+ "use_clipped_value_loss": True,
+ "clip_param": 0.2,
+ "entropy_coef": 0.01,
+ "num_learning_epochs": 5,
+ "num_mini_batches": 4,
+ "learning_rate": 1.0e-3,
+ "schedule": "adaptive",
+ "gamma": 0.99,
+ "lam": 0.95,
+ "desired_kl": 0.01,
+ "max_grad_norm": 1.0,
+ },
+ "policy": {
+ "init_noise_std": 1.0,
+ },
+}
+
+
+def main():
+ p = argparse.ArgumentParser()
+ p.add_argument("--num-envs", type=int, default=2048)
+ p.add_argument("--iterations", type=int, default=3000)
+ p.add_argument("--seed", type=int, default=1)
+ p.add_argument("--init-noise-std", type=float, default=1.0)
+ p.add_argument("--force-std", action="store_true",
+ help="force reset action std to --init-noise-std even when resuming")
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
+ p.add_argument("--resume", default=None,
+ help="checkpoint .pt to warm-start from (e.g. the flat policy before "
+ "the pyramid-terrain curriculum). Loads model + optimizer; the "
+ "action std comes from the checkpoint, not --init-noise-std.")
+ p.add_argument("--level", type=int, default=None,
+ help="force ALL envs to this terrain level (0-9), skip curriculum")
+ args = p.parse_args()
+
+ torch.manual_seed(args.seed)
+ import numpy as np
+ np.random.seed(args.seed)
+
+ TRAIN_CFG["policy"]["init_noise_std"] = args.init_noise_std
+
+ raw_env = env_registry.make("go1-dreamwaq-walk", num_envs=args.num_envs)
+ if args.level is not None:
+ raw_env._force_level = args.level
+ print(f"[DreamWaQ-rsl] forcing ALL envs at terrain level {args.level}")
+ env = DwaqVecEnv(raw_env, device=args.device)
+ print(f"[DreamWaQ-rsl] {args.num_envs} envs | obs={env.num_obs} priv={env.num_privileged_obs} "
+ f"hist={env.num_obs_hist} act={env.num_actions} | device={args.device}")
+ print(f"[DreamWaQ-rsl] actor_in={env.num_obs + 19} critic_in={env.num_privileged_obs} "
+ f"cenet_in={env.num_obs_hist * env.num_obs}")
+
+ log_dir = os.path.join("runs", "go1-dreamwaq-walk", "rsl_dwaq",
+ __import__("datetime").datetime.now().strftime("%m-%d_%H-%M-%S"))
+ os.makedirs(log_dir, exist_ok=True)
+
+ # Save config snapshot for later reference
+ import json
+ cfg_snapshot = {
+ "command_line": sys.argv,
+ "kp": raw_env.cfg.control_config.stiffness,
+ "kd": raw_env.cfg.control_config.damping,
+ "action_scale": raw_env.cfg.control_config.action_scale,
+ "rewards": dict(raw_env.cfg.reward_config.scales),
+ "sigma": raw_env.cfg.reward_config.tracking_sigma,
+ "only_positive": raw_env.cfg.reward_config.only_positive_rewards,
+ "force_level": args.level,
+ "init_noise_std": TRAIN_CFG["policy"]["init_noise_std"],
+ "entropy_coef": TRAIN_CFG["algorithm"]["entropy_coef"],
+ "terrain_rows": raw_env.cfg.terrain_rows,
+ "terrain_cols": raw_env.cfg.terrain_cols,
+ "scene": raw_env.cfg.model_file,
+ }
+ with open(os.path.join(log_dir, "config.json"), "w") as f:
+ json.dump(cfg_snapshot, f, indent=2, default=str)
+
+ runner = OnPolicyRunner(env, TRAIN_CFG, log_dir=log_dir, device=args.device)
+ if args.resume:
+ runner.load(args.resume)
+ runner.current_learning_iteration = 0
+ if args.force_std:
+ runner.alg.actor_critic.std.data.fill_(args.init_noise_std)
+ print(f"[DreamWaQ-rsl] warm-start from {args.resume} "
+ f"(std FORCED to {args.init_noise_std})")
+ else:
+ print(f"[DreamWaQ-rsl] warm-start from {args.resume} "
+ f"(model+optimizer; std from checkpoint)")
+ print(f"[DreamWaQ-rsl] log_dir={log_dir} | training {args.iterations} iterations...")
+ runner.learn(args.iterations, init_at_random_ep_len=True)
+ print("[DreamWaQ-rsl] done.")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/train_go2style.py b/scripts/train_go2style.py
new file mode 100644
index 0000000..f70ff4b
--- /dev/null
+++ b/scripts/train_go2style.py
@@ -0,0 +1,95 @@
+#!/usr/bin/env python3
+"""Train Go1 go2style flat-terrain locomotion.
+
+Usage:
+ uv run scripts/train_go2style.py
+ uv run scripts/train_go2style.py --rllib rslrl
+"""
+
+import logging
+
+# IMPORTANT: trigger registration of go2style env + rl config
+import motrix_rl.tasks.go1_go2style # noqa: F401
+
+# Now run the standard training pipeline
+from absl import app, flags
+from skrl import config
+
+from motrix_rl import utils
+
+logger = logging.getLogger(__name__)
+
+_ENV = flags.DEFINE_string("env", "go1-flat-terrain-walk-go2style", "The env to train")
+_SIM_BACKEND = flags.DEFINE_string("sim-backend", None, "Simulation backend")
+_NUM_ENVS = flags.DEFINE_integer("num-envs", 2048, "Number of envs")
+_RENDER = flags.DEFINE_bool("render", False, "Render the env")
+_TRAIN_BACKEND = flags.DEFINE_string("train-backend", None, "learning backend (jax/torch)")
+_SEED = flags.DEFINE_integer("seed", None, "Random seed")
+_RAND_SEED = flags.DEFINE_bool("rand-seed", False, "Generate random seed")
+_RLLIB = flags.DEFINE_string("rllib", "skrl", "RL framework (skrl/rslrl)")
+
+
+def get_train_backend(supports, train_backend_arg, rllib):
+ if rllib == "rslrl":
+ if train_backend_arg is not None and train_backend_arg != "torch":
+ raise Exception("RSLRL only supports PyTorch backend.")
+ if not supports.torch:
+ raise Exception("RSLRL requires PyTorch.")
+ return "torch"
+ if train_backend_arg is not None:
+ backend = train_backend_arg
+ if backend == "jax" and not supports.jax:
+ raise Exception("JAX not available.")
+ if backend == "torch" and not supports.torch:
+ raise Exception("PyTorch not available.")
+ return backend
+ if supports.jax and supports.jax_gpu:
+ return "jax"
+ elif supports.torch and supports.torch_gpu:
+ return "torch"
+ elif supports.jax:
+ return "jax"
+ elif supports.torch:
+ return "torch"
+ else:
+ raise Exception("Neither JAX nor PyTorch available.")
+
+
+def main(argv):
+ device_supports = utils.get_device_supports()
+ logger.info(device_supports)
+ env_name = _ENV.value
+ enable_render = _RENDER.value
+
+ rl_override = {}
+ if _NUM_ENVS.present:
+ rl_override["num_envs"] = _NUM_ENVS.value
+ if _RAND_SEED.value:
+ rl_override["runner.seed"] = None
+ elif _SEED.present:
+ rl_override["runner.seed"] = _SEED.value
+
+ sim_backend = _SIM_BACKEND.value
+ rllib = _RLLIB.value
+ train_backend = get_train_backend(device_supports, _TRAIN_BACKEND.value, rllib)
+
+ if rllib == "rslrl":
+ assert device_supports.torch
+ from motrix_rl.rslrl.torch.train import ppo
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ elif train_backend == "jax":
+ from motrix_rl.skrl.jax.train import ppo
+ config.jax.backend = "jax"
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ elif train_backend == "torch":
+ from motrix_rl.skrl.torch.train import ppo
+ config.torch.backend = "torch"
+ trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
+ else:
+ raise Exception(f"Unknown train backend: {train_backend}")
+
+ trainer.train()
+
+
+if __name__ == "__main__":
+ app.run(main)
diff --git a/scripts/view_go2style.py b/scripts/view_go2style.py
new file mode 100644
index 0000000..37c1d57
--- /dev/null
+++ b/scripts/view_go2style.py
@@ -0,0 +1,135 @@
+#!/usr/bin/env python3
+"""MuJoCo sim2sim visualization for go2style policy (45-dim obs, no linvel).
+
+Controls:
+ W/S: forward/back Q/E: left/right A/D: rotate Space: stop R: reset
+"""
+import numpy as np
+import mujoco
+from mujoco import viewer
+import onnxruntime as ort
+import os, sys, time, threading, queue
+
+PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+ONNX_PATH = os.path.join(PROJECT, "exports_go1_go2style", "policy.onnx")
+XML_DIR = os.path.join(PROJECT, "motrix_envs", "src", "motrix_envs", "locomotion", "go1", "xmls")
+
+# go2style params
+NUM_OBS = 45
+NUM_ACTIONS = 12
+OBS_SCALES = {'ang_vel': 0.25, 'dof_pos': 1.0, 'dof_vel': 0.05}
+ACTION_SCALE = 0.25
+KP = 20.0
+KD = 0.0 # MuJoCo joint自带damping=0.5, PD kd=0避免过阻尼
+CLIP_ACTIONS = 23.7
+CLIP_OBS = 100.0
+MAX_VX, MAX_VY, MAX_WZ = 1.0, 1.0, 1.0
+
+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)
+
+from pynput import keyboard
+
+class KB:
+ def __init__(self):
+ self._q = queue.Queue(); self.running = True
+ self.held = set(); self._t = None; self._l = None
+ def _n(self, k):
+ try:
+ if hasattr(k,'char') and k.char: return k.char.lower()
+ except: pass
+ return str(k).lower()
+ def _w(self):
+ while self.running:
+ try:
+ et, k = self._q.get(timeout=0.05)
+ n = self._n(k)
+ if et == 'press': self.held.add(n)
+ elif et == 'release': self.held.discard(n)
+ except queue.Empty: pass
+ def init(self):
+ def op(k): self._q.put(('press',k))
+ def or_(k): self._q.put(('release',k))
+ self._l = keyboard.Listener(on_press=op, on_release=or_)
+ self._l.start()
+ self._t = threading.Thread(target=self._w, daemon=True); self._t.start()
+ print("[KB] 键盘就绪")
+ def held_keys(self): return self.held.copy()
+ def stop(self): self.running = False; self._l.stop()
+
+def get_sensor(m, d, name):
+ sid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_SENSOR, name)
+ if sid < 0: return None
+ adr = m.sensor_adr[sid]; dim = m.sensor_dim[sid]
+ return d.sensordata[adr:adr+dim].copy()
+
+def compute_obs(model, data, commands, last_action):
+ obs = np.zeros(NUM_OBS, dtype=np.float32)
+ # gyro [0:3]
+ g = get_sensor(model, data, "gyro")
+ obs[0:3] = (g if g is not None else data.qvel[3:6]) * OBS_SCALES['ang_vel']
+ # gravity [3:6]
+ R = data.xmat[1].reshape(3,3)
+ obs[3:6] = (R.T @ np.array([0.,0.,-1.])).astype(np.float32)
+ # joint pos [6:18]
+ obs[6:18] = (data.qpos[7:19] - DEFAULT_ANGLES) * OBS_SCALES['dof_pos']
+ # joint vel [18:30]
+ obs[18:30] = data.qvel[6:18] * OBS_SCALES['dof_vel']
+ # last action [30:42]
+ obs[30:42] = last_action
+ # commands [42:45]
+ obs[42:45] = commands * np.array([2.0, 2.0, 0.25], dtype=np.float32)
+ return np.clip(obs, -CLIP_OBS, CLIP_OBS)
+
+def main():
+ os.chdir(XML_DIR)
+ xml = open("scene_motor_actuator.xml").read()
+ model = mujoco.MjModel.from_xml_string(xml)
+ data = mujoco.MjData(model)
+ data.qpos[0:3] = [0,0,0.42]; data.qpos[3:7] = [1,0,0,0]; data.qpos[7:19] = DEFAULT_ANGLES
+ mujoco.mj_forward(model, data)
+
+ session = ort.InferenceSession(ONNX_PATH, providers=['CPUExecutionProvider'])
+ print(f"[ONNX] {ONNX_PATH}")
+ print(f"[CTRL] W/S前后 Q/E左右 A/D旋转 Space停 R重置 Esc退出")
+
+ kb = KB(); kb.init()
+ view = viewer.launch_passive(model, data)
+
+ step, vx, vy, wz = 0, 0.0, 0.0, 0.0
+ last_action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ action = np.zeros(NUM_ACTIONS, dtype=np.float32)
+ decimation = 2 # MuJoCo dt=0.005, policy dt=0.01 → 2 steps per inference
+
+ while view.is_running():
+ keys = kb.held_keys()
+ if 'escape' in keys: break
+ if 'r' in keys:
+ data.qpos[0:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0]; data.qpos[7:19]=DEFAULT_ANGLES
+ data.qvel[:]=0; last_action[:]=0; mujoco.mj_forward(model,data); print("[R] 重置")
+ if ' ' in keys: vx=vy=wz=0.0
+
+ vx = MAX_VX if 'w' in keys else (-MAX_VX if 's' in keys else 0.0)
+ vy = MAX_VY if 'q' in keys else (-MAX_VY if 'e' in keys else 0.0)
+ wz = MAX_WZ if 'a' in keys else (-MAX_WZ if 'd' in keys else 0.0)
+
+ if step % decimation == 0:
+ cmd = np.array([vx, vy, wz], dtype=np.float32)
+ obs = compute_obs(model, data, cmd, last_action)
+ action = session.run(None, {'observations': obs.reshape(1,-1).astype(np.float32)})[0][0]
+ action = np.clip(action, -CLIP_ACTIONS, CLIP_ACTIONS)
+ last_action = action.copy()
+
+ target = DEFAULT_ANGLES + action * ACTION_SCALE
+ torques = KP*(target - data.qpos[7:19]) - KD*data.qvel[6:18]
+ data.ctrl[:] = np.clip(torques, -CLIP_ACTIONS, CLIP_ACTIONS)
+ mujoco.mj_step(model, data)
+ view.sync()
+ step += 1
+ time.sleep(0.001)
+
+ kb.stop(); view.close()
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/view_orig.py b/scripts/view_orig.py
new file mode 100644
index 0000000..1c97ab4
--- /dev/null
+++ b/scripts/view_orig.py
@@ -0,0 +1,89 @@
+#!/usr/bin/env python3
+"""MuJoCo viewer for original Go1 (45-dim, PD 80/1.0, action_scale=0.05)"""
+import numpy as np, mujoco, onnxruntime as ort, os, time, threading, queue
+from mujoco import viewer
+from pynput import keyboard
+
+PROJECT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+ONNX = os.path.join(PROJECT, "exports_go1_orig", "policy.onnx")
+XML_DIR = os.path.join(PROJECT, "motrix_envs/src/motrix_envs/locomotion/go1/xmls")
+
+# Original params
+NUM_OBS = 57
+KP, KD = 80.0, 0.5 # KD=0.5 + joint_damping(0.5) = 1.0 = training kd
+ACTION_SCALE = 0.05
+CLIP = 23.7
+DEFAULT = 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)
+
+class KB:
+ def __init__(self):
+ self._q=queue.Queue(); self.running=True; self.held=set()
+ def _n(self,k):
+ try:
+ if hasattr(k,'char') and k.char: return k.char.lower()
+ except: pass
+ return str(k).lower()
+ def _w(self):
+ while self.running:
+ try:
+ et,k=self._q.get(timeout=0.05); n=self._n(k)
+ if et=='press': self.held.add(n)
+ elif et=='release': self.held.discard(n)
+ except queue.Empty: pass
+ def init(self):
+ self._l=keyboard.Listener(on_press=lambda k:self._q.put(('press',k)), on_release=lambda k:self._q.put(('release',k)))
+ self._l.start(); self._t=threading.Thread(target=self._w,daemon=True); self._t.start()
+ def held_keys(self): return self.held.copy()
+ def stop(self): self.running=False; self._l.stop()
+
+def get_sensor(m,d,name):
+ sid=mujoco.mj_name2id(m,mujoco.mjtObj.mjOBJ_SENSOR,name)
+ if sid<0: return None
+ adr=m.sensor_adr[sid]; return d.sensordata[adr:adr+m.sensor_dim[sid]].copy()
+
+def compute_obs(model,data,cmd,last_a):
+ obs=np.zeros(NUM_OBS,dtype=np.float32)
+ g=get_sensor(model,data,"gyro")
+ obs[0:3]=(g if g is not None else data.qvel[3:6])*0.25
+ R=data.xmat[1].reshape(3,3)
+ obs[3:6]=(R.T@np.array([0.,0.,-1.])).astype(np.float32)
+ obs[6:18]=(data.qpos[7:19]-DEFAULT)*1.0
+ obs[18:30]=data.qvel[6:18]*0.05
+ obs[30:42]=last_a
+ obs[42:45]=cmd; obs[45:57]=0.0 # contact_force*np.array([2.,2.,0.25],dtype=np.float32)
+ return np.clip(obs,-100.,100.)
+
+def main():
+ os.chdir(XML_DIR)
+ model=mujoco.MjModel.from_xml_string(open("scene_motor_actuator.xml").read())
+ data=mujoco.MjData(model)
+ data.qpos[0:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0]; data.qpos[7:19]=DEFAULT
+ mujoco.mj_forward(model,data)
+ session=ort.InferenceSession(ONNX,providers=['CPUExecutionProvider'])
+ print(f"[ORIG] kp={KP} kd={KD+0.5} action_scale={ACTION_SCALE} | W/S前后 Q/E左右 A/D旋转")
+ kb=KB(); kb.init()
+ view=viewer.launch_passive(model,data)
+ step,vx,vy,wz=0,0.,0.,0.
+ last_a=np.zeros(12,dtype=np.float32); action=np.zeros(12,dtype=np.float32)
+ dec=2
+ while view.is_running():
+ keys=kb.held_keys()
+ if 'escape' in keys: break
+ if 'r' in keys:
+ data.qpos[0:3]=[0,0,0.42]; data.qpos[3:7]=[1,0,0,0]; data.qpos[7:19]=DEFAULT
+ data.qvel[:]=0; last_a[:]=0; mujoco.mj_forward(model,data)
+ if ' ' in keys: vx=vy=wz=0.
+ vx=1.0 if 'w' in keys else (-1.0 if 's' in keys else 0.)
+ vy=1.0 if 'q' in keys else (-1.0 if 'e' in keys else 0.)
+ wz=1.0 if 'a' in keys else (-1.0 if 'd' in keys else 0.)
+ if step%dec==0:
+ obs=compute_obs(model,data,np.array([vx,vy,wz],dtype=np.float32),last_a)
+ action=session.run(None,{'observations':obs.reshape(1,-1).astype(np.float32)})[0][0]
+ action=np.clip(action,-CLIP,CLIP); last_a=action.copy()
+ target=DEFAULT+action*ACTION_SCALE
+ t=KP*(target-data.qpos[7:19])-KD*data.qvel[6:18]
+ data.ctrl[:]=np.clip(t,-CLIP,CLIP)
+ mujoco.mj_step(model,data); view.sync(); step+=1; time.sleep(0.001)
+ kb.stop(); view.close()
+
+if __name__=="__main__": main()
diff --git a/uv.lock b/uv.lock
index f19a5ae..0f0a52d 100644
--- a/uv.lock
+++ b/uv.lock
@@ -2,9 +2,6 @@ version = 1
revision = 3
requires-python = "==3.10.*"
-[options]
-prerelease-mode = "allow"
-
[manifest]
members = [
"motrix-envs",
@@ -94,7 +91,7 @@ name = "certifi"
version = "2022.12.7"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/certifi-2022.12.7-py3-none-any.whl", hash = "sha256:4ad3232f5e926d6718ec31cfc1fcadfde020920e278684144551c91769c7bc18" },
+ { url = "https://download.pytorch.org/whl/certifi-2022.12.7-py3-none-any.whl", hash = "sha256:4ad3232f5e926d6718ec31cfc1fcadfde020920e278684144551c91769c7bc18", upload-time = "2023-10-03T17:33:19Z" },
]
[[package]]
@@ -102,7 +99,7 @@ name = "charset-normalizer"
version = "2.1.1"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/charset_normalizer-2.1.1-py3-none-any.whl", hash = "sha256:83e9a75d1911279afd89352c68b45348559d1fc0506b054b346651b5e7fee29f" },
+ { url = "https://download.pytorch.org/whl/charset_normalizer-2.1.1-py3-none-any.whl", hash = "sha256:83e9a75d1911279afd89352c68b45348559d1fc0506b054b346651b5e7fee29f", upload-time = "2023-10-03T17:33:19Z" },
]
[[package]]
@@ -136,7 +133,7 @@ name = "colorama"
version = "0.4.6"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6" },
+ { url = "https://download.pytorch.org/whl/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6", upload-time = "2023-10-05T23:50:34Z" },
]
[[package]]
@@ -346,7 +343,7 @@ name = "idna"
version = "3.4"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/idna-3.4-py3-none-any.whl", hash = "sha256:90b77e79eaa3eba6de819a0c442c0b4ceefc341a7a2ab77d7562bf49f425c5c2" },
+ { url = "https://download.pytorch.org/whl/idna-3.4-py3-none-any.whl", hash = "sha256:90b77e79eaa3eba6de819a0c442c0b4ceefc341a7a2ab77d7562bf49f425c5c2", upload-time = "2023-10-06T21:35:40Z" },
]
[[package]]
@@ -366,7 +363,7 @@ dependencies = [
{ name = "zipp" },
]
wheels = [
- { url = "https://download.pytorch.org/whl/importlib_metadata-7.1.0-py3-none-any.whl" },
+ { url = "https://download.pytorch.org/whl/importlib_metadata-7.1.0-py3-none-any.whl", upload-time = "2025-08-05T21:34:24Z" },
]
[[package]]
@@ -466,7 +463,7 @@ dependencies = [
{ name = "markupsafe" },
]
wheels = [
- { url = "https://download.pytorch.org/whl/jinja2-3.1.6-py3-none-any.whl" },
+ { url = "https://download.pytorch.org/whl/jinja2-3.1.6-py3-none-any.whl", upload-time = "2025-10-14T18:38:59Z" },
]
[[package]]
@@ -526,10 +523,10 @@ name = "markupsafe"
version = "3.0.3"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-macosx_11_0_arm64.whl" },
- { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl" },
- { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl" },
- { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-win_amd64.whl" },
+ { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-macosx_11_0_arm64.whl", upload-time = "2026-03-27T13:54:27Z" },
+ { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", upload-time = "2026-03-27T13:54:27Z" },
+ { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", upload-time = "2026-03-27T13:54:27Z" },
+ { url = "https://download.pytorch.org/whl/markupsafe-3.0.3-cp310-cp310-win_amd64.whl", upload-time = "2026-03-27T13:54:29Z" },
]
[[package]]
@@ -1058,7 +1055,7 @@ name = "packaging"
version = "24.1"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/packaging-24.1-py3-none-any.whl", hash = "sha256:5b8f2217dbdbd2f7f384c41c628544e6d52f2d0f53c6d0c3ea61aa5d1d7ff124" },
+ { url = "https://download.pytorch.org/whl/packaging-24.1-py3-none-any.whl", hash = "sha256:5b8f2217dbdbd2f7f384c41c628544e6d52f2d0f53c6d0c3ea61aa5d1d7ff124", upload-time = "2024-10-29T23:48:01Z" },
]
[[package]]
@@ -1191,7 +1188,7 @@ dependencies = [
{ name = "urllib3" },
]
wheels = [
- { url = "https://download.pytorch.org/whl/requests-2.28.1-py3-none-any.whl", hash = "sha256:8fefa2a1a1365bf5520aac41836fbee479da67864514bdb821f31ce07ce65349" },
+ { url = "https://download.pytorch.org/whl/requests-2.28.1-py3-none-any.whl", hash = "sha256:8fefa2a1a1365bf5520aac41836fbee479da67864514bdb821f31ce07ce65349", upload-time = "2023-10-06T21:36:51Z" },
]
[[package]]
@@ -1245,7 +1242,7 @@ name = "setuptools"
version = "70.2.0"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/setuptools-70.2.0-py3-none-any.whl", hash = "sha256:b8b8060bb426838fbe942479c90296ce976249451118ef566a5a0b7d8b78fb05" },
+ { url = "https://download.pytorch.org/whl/setuptools-70.2.0-py3-none-any.whl", hash = "sha256:b8b8060bb426838fbe942479c90296ce976249451118ef566a5a0b7d8b78fb05", upload-time = "2025-01-30T19:44:58Z" },
]
[[package]]
@@ -1624,9 +1621,9 @@ dependencies = [
{ name = "typing-extensions" },
]
wheels = [
- { url = "https://download.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_aarch64.whl", hash = "sha256:b1f0cdd0720ad60536deb5baa427b782fd920dd4fcf72e244d32974caafa3b9e" },
- { url = "https://download.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:ac1849553ee673dfafb44c610c60cb60a2890f0e117f43599a526cf777eb8b8c" },
- { url = "https://download.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:c52c4b869742f00b12cb34521d1381be6119fa46244791704b00cc4a3cb06850" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_aarch64.whl", hash = "sha256:b1f0cdd0720ad60536deb5baa427b782fd920dd4fcf72e244d32974caafa3b9e", upload-time = "2025-05-14T03:33:41Z" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:ac1849553ee673dfafb44c610c60cb60a2890f0e117f43599a526cf777eb8b8c", upload-time = "2025-04-22T18:19:25Z" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:c52c4b869742f00b12cb34521d1381be6119fa46244791704b00cc4a3cb06850", upload-time = "2025-04-22T18:19:44Z" },
]
[[package]]
@@ -1637,8 +1634,8 @@ dependencies = [
{ name = "torch" },
]
wheels = [
- { url = "https://download-r2.pytorch.org/whl/cu128/torchaudio-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:d62b6e9b792ad37af6d1289ba283e1029e71b4ff9cd3c6cf7f0e7776f23254b2" },
- { url = "https://download-r2.pytorch.org/whl/cu128/torchaudio-2.7.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:f96c2be8aff6c827e76fd3a85e69a54ba5b9a37090853ed886f056ddfbca09a4" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torchaudio-2.7.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:d62b6e9b792ad37af6d1289ba283e1029e71b4ff9cd3c6cf7f0e7776f23254b2", upload-time = "2025-04-22T18:30:28Z" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torchaudio-2.7.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:f96c2be8aff6c827e76fd3a85e69a54ba5b9a37090853ed886f056ddfbca09a4", upload-time = "2025-04-22T18:30:29Z" },
]
[[package]]
@@ -1651,8 +1648,8 @@ dependencies = [
{ name = "torch" },
]
wheels = [
- { url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:59df5a550113a80ce523047066eaaedb168c69482da88c3ab246716ab45ba092" },
- { url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:cdd90b768b01b0d638cb06a6c211b550b275c0c207b5210b7cbb5cea8dde11db" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0%2Bcu128-cp310-cp310-manylinux_2_28_x86_64.whl", hash = "sha256:59df5a550113a80ce523047066eaaedb168c69482da88c3ab246716ab45ba092", upload-time = "2025-04-22T18:30:20Z" },
+ { url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0%2Bcu128-cp310-cp310-win_amd64.whl", hash = "sha256:cdd90b768b01b0d638cb06a6c211b550b275c0c207b5210b7cbb5cea8dde11db", upload-time = "2025-04-22T18:30:20Z" },
]
[[package]]
@@ -1663,7 +1660,7 @@ dependencies = [
{ name = "colorama", marker = "sys_platform == 'win32'" },
]
wheels = [
- { url = "https://download.pytorch.org/whl/tqdm-4.66.5-py3-none-any.whl", hash = "sha256:90279a3770753eafc9194a0364852159802111925aa30eb3f9d85b0e805ac7cd" },
+ { url = "https://download.pytorch.org/whl/tqdm-4.66.5-py3-none-any.whl", hash = "sha256:90279a3770753eafc9194a0364852159802111925aa30eb3f9d85b0e805ac7cd", upload-time = "2024-10-30T00:09:55Z" },
]
[[package]]
@@ -1686,8 +1683,8 @@ dependencies = [
{ name = "setuptools" },
]
wheels = [
- { url = "https://download-r2.pytorch.org/whl/triton-3.3.0-cp310-cp310-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6cbb1b3c214d17e06e2c12399747f645fb7e8b72cb24271557cdce0893c764c6" },
- { url = "https://download-r2.pytorch.org/whl/triton-3.3.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:4198996c9fa3fd811e3bc007f0fc9853c784be3dae6d30714f579c5106d70616" },
+ { url = "https://download-r2.pytorch.org/whl/triton-3.3.0-cp310-cp310-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6cbb1b3c214d17e06e2c12399747f645fb7e8b72cb24271557cdce0893c764c6", upload-time = "2026-01-22T22:41:17Z" },
+ { url = "https://download-r2.pytorch.org/whl/triton-3.3.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:4198996c9fa3fd811e3bc007f0fc9853c784be3dae6d30714f579c5106d70616", upload-time = "2026-01-22T22:41:22Z" },
]
[[package]]
@@ -1703,7 +1700,7 @@ name = "urllib3"
version = "1.26.13"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/urllib3-1.26.13-py2.py3-none-any.whl", hash = "sha256:47cc05d99aaa09c9e72ed5809b60e7ba354e64b59c9c173ac3018642d8bb41fc" },
+ { url = "https://download.pytorch.org/whl/urllib3-1.26.13-py2.py3-none-any.whl", hash = "sha256:47cc05d99aaa09c9e72ed5809b60e7ba354e64b59c9c173ac3018642d8bb41fc", upload-time = "2023-10-06T21:48:51Z" },
]
[[package]]
@@ -1747,5 +1744,5 @@ name = "zipp"
version = "3.19.2"
source = { registry = "https://download.pytorch.org/whl/cu128" }
wheels = [
- { url = "https://download.pytorch.org/whl/zipp-3.19.2-py3-none-any.whl" },
+ { url = "https://download.pytorch.org/whl/zipp-3.19.2-py3-none-any.whl", upload-time = "2025-08-05T22:17:38Z" },
]