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