feat: DreamWaQ full replication — env, terrain, CENet, PPO

This commit is contained in:
8x54zj-m
2026-06-30 14:53:12 +08:00
parent c569f4d6d9
commit 422421d263
17 changed files with 2056 additions and 17 deletions

269
scripts/view_dreamwaq.py Normal file
View File

@@ -0,0 +1,269 @@
#!/usr/bin/env python3
"""DreamWaQ 地形可视化。
键盘:
R=重置 H=高度采样点 T=遍历出生点调试 Esc=退出
用法:
uv run scripts/view_dreamwaq.py # 金字塔地形
uv run scripts/view_dreamwaq.py --flat --num-envs 1
"""
import argparse, os, sys, time
import numpy as np
os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false")
os.environ.setdefault("JAX_PLATFORMS", "cpu")
if "--flat" in sys.argv:
os.environ["DREAMWAQ_TERRAIN"] = "flat"
elif "--flat-stairs" in sys.argv:
os.environ["DREAMWAQ_TERRAIN"] = "flat_stairs"
elif "--stairs" in sys.argv:
os.environ["DREAMWAQ_TERRAIN"] = "stairs"
else:
os.environ.setdefault("DREAMWAQ_TERRAIN", "pyramid")
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import motrix_envs.locomotion.go1.dreamwaq # noqa: F401
from motrix_envs import registry as env_registry
from motrix_envs.np.renderer import NpRenderer
from motrix_envs.math import quaternion
from motrixsim.render import RenderClosedError
# 地形类型名称(与 gen_dreamwaq_terrain.py 的 PROPORTIONS 对应)
PROPORTIONS = [0.1, 0.1, 0.35, 0.35, 0.1]
CUM = [sum(PROPORTIONS[:i + 1]) for i in range(len(PROPORTIONS))]
TYPE_NAMES = ["平滑斜坡", "粗糙斜坡", "下行楼梯", "上行楼梯", "离散障碍"]
def _cell_origin(row, col, border_m=5.0, cell_m=8.0, num_rows=10, num_cols=20):
"""计算 cell (row, col) 的中心世界坐标。"""
half_x = border_m + num_cols * cell_m / 2.0
half_y = border_m + num_rows * cell_m / 2.0
cx = -half_x + border_m + col * cell_m + cell_m / 2
cy = half_y - border_m - row * cell_m - cell_m / 2
return cx, cy
def _type_name(col):
"""根据列索引返回地形类型名称。"""
choice = col / 20 + 0.001
for i, cum in enumerate(CUM):
if choice < cum:
return TYPE_NAMES[i]
return TYPE_NAMES[-1]
def _spawn_at(env, cx, cy, spawn_z=None):
"""在指定世界坐标 spawn 单个机器人。"""
state = env._state
data = state.data
init_pos = env._init_dof_pos.copy().reshape(1, -1)
init_pos[0, 0] = cx
init_pos[0, 1] = cy
# 计算地形高度
terrain_z = float(env._sample_terrain_height(
np.array([[cx, cy]], dtype=np.float32), radius=0.35)[0])
if spawn_z is None:
spawn_z = terrain_z + 0.45 # 默认 clearance
init_pos[0, 2] = spawn_z
data.reset(env._model)
data.set_dof_pos(init_pos, env._model)
env._model.forward_kinematic(data)
# 重置 info
state.info["commands"][0] = np.array([0.0, 0.0, 0.0], dtype=np.float32)
state.info["steps"][0] = 0
state.info["obs_history"][0] = 0.0
return terrain_z, spawn_z
def main():
p = argparse.ArgumentParser(description="DreamWaQ 地形可视化")
p.add_argument("--num-envs", type=int, default=1)
p.add_argument("--flat", action="store_true")
p.add_argument("--flat-stairs", action="store_true")
p.add_argument("--stairs", action="store_true")
p.add_argument("--level", type=int, default=None)
p.add_argument("--no-stand", action="store_true")
p.add_argument("--vx", type=float, default=0.5)
args = p.parse_args()
env = env_registry.make("go1-dreamwaq-walk", num_envs=max(args.num_envs, 1))
if args.level is not None:
env._force_level = args.level
env.init_state()
n = env._num_envs
cmd = np.array([args.vx, 0.0, 0.0], dtype=np.float32)
try:
renderer = NpRenderer(env)
except Exception as e:
print(f"[ERROR] 渲染器创建失败: {e}")
renderer = None
terrain_name = os.environ.get("DREAMWAQ_TERRAIN", "pyramid")
print(f"[View] {n} 机器人 | 地形={terrain_name}")
print(f"[View] R=重置 H=高度点 T=遍历出生点 Esc=退出")
show_heights = False
traverse_mode = False
traverse_row = 0
traverse_col = 0
traverse_pending = False # 刚切换 cell, 等待稳定
traverse_settle = 0
step_count = 0
# 地形信息
num_rows = env._num_rows
num_cols = env._num_cols
cell_m = env._cell_size
def enter_traverse():
nonlocal traverse_mode, traverse_row, traverse_col, traverse_pending, traverse_settle
traverse_mode = True
traverse_row = 0
traverse_col = 0
traverse_pending = True
traverse_settle = 0
print(f"\n[T] 遍历模式: {num_rows}× {num_cols}")
print(f"[T] 按 T 前进, R 退出遍历\n")
def exit_traverse():
nonlocal traverse_mode, traverse_pending
traverse_mode = False
traverse_pending = False
env.init_state()
print("[T] 退出遍历模式\n")
def advance_traverse():
nonlocal traverse_row, traverse_col, traverse_pending, traverse_settle
traverse_col += 1
if traverse_col >= num_cols:
traverse_col = 0
traverse_row += 1
if traverse_row >= num_rows:
print("[T] 遍历完成! 按 R 退出")
traverse_row = num_rows - 1
traverse_col = num_cols - 1
return False
traverse_pending = True
traverse_settle = 0
return True
def do_traverse_spawn():
"""在当前位置 spawn 并打印信息。"""
nonlocal traverse_settle
cx, cy = _cell_origin(traverse_row, traverse_col)
tname = _type_name(traverse_col)
terrain_z, spawn_z = _spawn_at(env, cx, cy)
# 标记 spawn 点(绿色球)
if renderer is not None:
g = renderer._render.gizmos
g.draw_sphere(0.15, (np.float32(cx), np.float32(cy),
np.float32(spawn_z)))
# 让机器人稳定几步
for _ in range(30):
env.step(np.zeros((1, 12), dtype=np.float32))
if renderer is not None:
renderer.render()
time.sleep(0.005)
base_z = env._body.get_pose(env._state.data)[0, 2]
contacts = env._state.info.get("contacts", np.zeros(4))
cf = env._state.info.get("privileged_obs",
np.zeros((1, 247)))[0, 45:57]
total_cf = np.sum(np.abs(cf))
print(f" r{traverse_row}c{traverse_col:02d} {tname:6s} "
f"origin=({cx:+.0f},{cy:+.0f}) "
f"terrain_z={terrain_z:.3f} spawn_z={spawn_z:.3f} "
f"base_z={base_z:.3f} cf={total_cf:.0f}N "
f"feet={contacts.astype(int).tolist()}")
traverse_settle = 30
try:
while True:
# ── 键盘 ──
if renderer is not None:
try:
inp = renderer._render.input
if inp.is_key_just_pressed("r"):
if traverse_mode:
exit_traverse()
else:
env.init_state()
step_count = 0
print("[R] 重置")
if inp.is_key_just_pressed("h"):
show_heights = not show_heights
print(f"[H] 高度点: {'' if show_heights else ''}")
if inp.is_key_just_pressed("t"):
if traverse_mode:
advance_traverse()
else:
enter_traverse()
except Exception:
pass
# ── 遍历模式 ──
if traverse_mode and traverse_pending:
do_traverse_spawn()
traverse_pending = False
# ── 正常模式动作 ──
if not traverse_mode or traverse_settle > 0:
if args.no_stand:
env._state.info["commands"][:] = cmd
else:
act = 0.3 * (np.random.rand(n, 12).astype(np.float32) - 0.5)
env.step(act)
if traverse_settle > 0:
traverse_settle -= 1
if step_count % 100 == 0 and step_count > 0:
bz = env._body.get_pose(env._state.data)[0, 2]
print(f"[{step_count}] base_z={bz:.3f}")
# ── 高度点 ──
if show_heights and renderer is not None:
state = env._state
pose = env._body.get_pose(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
try:
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)))
except Exception:
pass
if renderer is not None:
renderer.render()
time.sleep(0.01)
step_count += 1
except (KeyboardInterrupt, RenderClosedError):
pass
try:
if renderer is not None:
renderer.close()
except Exception:
pass
print("[View] 结束")
if __name__ == "__main__":
main()