feat: DreamWaQ full replication — env, terrain, CENet, PPO
This commit is contained in:
269
scripts/view_dreamwaq.py
Normal file
269
scripts/view_dreamwaq.py
Normal 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()
|
||||
Reference in New Issue
Block a user