Files
Motrixlab/scripts/view_dreamwaq.py

270 lines
9.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""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()