#!/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()