dreamzero / scripts /trex_track /trex_fk.py
zhicao's picture
Upload dreamzero source code (no model checkpoints)
fbd9366 verified
Raw
History Blame Contribute Delete
5.9 kB
"""T-Rex Pinocchio FK helpers (no dexmate_urdf / coal install required)."""
from __future__ import annotations
import sys
import types
from functools import lru_cache
from pathlib import Path
from typing import Callable
import numpy as np
import pinocchio as pin
from pinocchio.robot_wrapper import RobotWrapper
_TREX_ROOT = Path("/scratch1/home/zhicao/T-Rex")
_QS_SRC = _TREX_ROOT / "dataset_quickstart" / "src"
_DEXMATE_PKG = _TREX_ROOT / "hardware_code" / "third_party" / "dexmate-urdf" / "src"
_VEGA_DIR = _DEXMATE_PKG / "dexmate_urdf" / "robots" / "humanoid" / "vega_1"
VEGA_URDF = str(_VEGA_DIR / "vega_1.urdf")
VEGA_PKG = str(_VEGA_DIR)
VEGA_SRDF = str(_VEGA_DIR / "vega_1.srdf")
DEFAULT_TORSO = np.array([0.9, 1.57, 0.1])
DEFAULT_HEAD = np.array([0.28, 0.0, 0.0])
SHARPA_LEFT_HAND_JOINT_ORDER = [
"left_thumb_CMC_FE",
"left_thumb_CMC_AA",
"left_thumb_MCP_FE",
"left_thumb_MCP_AA",
"left_thumb_IP",
"left_index_MCP_FE",
"left_index_MCP_AA",
"left_index_PIP",
"left_index_DIP",
"left_middle_MCP_FE",
"left_middle_MCP_AA",
"left_middle_PIP",
"left_middle_DIP",
"left_ring_MCP_FE",
"left_ring_MCP_AA",
"left_ring_PIP",
"left_ring_DIP",
"left_pinky_CMC",
"left_pinky_MCP_FE",
"left_pinky_MCP_AA",
"left_pinky_PIP",
"left_pinky_DIP",
]
SHARPA_RIGHT_HAND_JOINT_ORDER = [
"right_thumb_CMC_FE",
"right_thumb_CMC_AA",
"right_thumb_MCP_FE",
"right_thumb_MCP_AA",
"right_thumb_IP",
"right_index_MCP_FE",
"right_index_MCP_AA",
"right_index_PIP",
"right_index_DIP",
"right_middle_MCP_FE",
"right_middle_MCP_AA",
"right_middle_PIP",
"right_middle_DIP",
"right_ring_MCP_FE",
"right_ring_MCP_AA",
"right_ring_PIP",
"right_ring_DIP",
"right_pinky_CMC",
"right_pinky_MCP_FE",
"right_pinky_MCP_AA",
"right_pinky_PIP",
"right_pinky_DIP",
]
def _bootstrap_trex_imports() -> None:
if "coal" not in sys.modules:
coal = types.ModuleType("coal")
coal.Box = lambda *args, **kwargs: None
sys.modules["coal"] = coal
vega_1 = types.SimpleNamespace(
urdf=VEGA_URDF,
_parent_dir=VEGA_PKG,
srdf=VEGA_SRDF,
)
vega_1_mod = types.ModuleType("dexmate_urdf.robots.humanoid.vega_1")
vega_1_mod.vega_1 = vega_1
humanoid = types.ModuleType("dexmate_urdf.robots.humanoid")
humanoid.vega_1 = vega_1_mod
robots = types.ModuleType("dexmate_urdf.robots")
robots.humanoid = humanoid
dexmate_urdf = types.ModuleType("dexmate_urdf")
dexmate_urdf.robots = robots
for name, mod in [
("dexmate_urdf", dexmate_urdf),
("dexmate_urdf.robots", robots),
("dexmate_urdf.robots.humanoid", humanoid),
("dexmate_urdf.robots.humanoid.vega_1", vega_1_mod),
]:
sys.modules.setdefault(name, mod)
qs = str(_QS_SRC)
if qs not in sys.path:
sys.path.insert(0, qs)
@lru_cache(maxsize=1)
def _load_robot_modules():
_bootstrap_trex_imports()
from trex_dataset_quickstart.robot import ( # noqa: WPS433
build_full_robot,
build_reduced_bimanual_robot,
forward_kinematics,
)
return build_full_robot, build_reduced_bimanual_robot, forward_kinematics
@lru_cache(maxsize=1)
def get_full_robot() -> tuple[RobotWrapper, Callable, Callable]:
build_full_robot, _, _ = _load_robot_modules()
return build_full_robot({"head": DEFAULT_HEAD, "torso": DEFAULT_TORSO})
@lru_cache(maxsize=1)
def get_bimanual_robot() -> tuple[RobotWrapper, Callable, Callable]:
_, build_reduced_bimanual_robot, _ = _load_robot_modules()
return build_reduced_bimanual_robot({"head": DEFAULT_HEAD, "torso": DEFAULT_TORSO})
def state_to_components(state58: np.ndarray) -> dict[str, np.ndarray]:
s = np.asarray(state58, dtype=np.float64).reshape(-1)
if s.shape[0] != 58:
raise ValueError(f"expected 58-dim state, got {s.shape[0]}")
return {
"left_arm": s[0:7],
"left_hand": s[7:29],
"right_arm": s[29:36],
"right_hand": s[36:58],
}
def hand_joint_positions_3d(
robot: RobotWrapper,
qpos: np.ndarray,
joint_names: list[str],
) -> np.ndarray:
pin.forwardKinematics(robot.model, robot.data, qpos)
pin.updateFramePlacements(robot.model, robot.data)
pts = np.zeros((len(joint_names), 3), dtype=np.float64)
for i, name in enumerate(joint_names):
fid = robot.model.getFrameId(name, pin.FrameType.JOINT)
pts[i] = robot.data.oMf[fid].translation
return pts
def frame_pose_matrix(robot: RobotWrapper, qpos: np.ndarray, frame_name: str) -> np.ndarray:
pin.forwardKinematics(robot.model, robot.data, qpos)
pin.updateFramePlacements(robot.model, robot.data)
fid = robot.model.getFrameId(frame_name)
se3 = robot.data.oMf[fid]
T = np.eye(4, dtype=np.float64)
T[:3, :3] = se3.rotation
T[:3, 3] = se3.translation
return T
def hand_keypoints_world(
state58: np.ndarray,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Return left/right hand (22,3) keypoints and L_ee/R_ee cam-to-world 4x4."""
comps = state_to_components(state58)
full_robot, assemble_full, _ = get_full_robot()
bimanual_robot, assemble_bimanual, _ = get_bimanual_robot()
q_full = assemble_full(comps)
q_bi = assemble_bimanual(
{"left_arm": comps["left_arm"], "right_arm": comps["right_arm"]},
)
left_pts = hand_joint_positions_3d(full_robot, q_full, SHARPA_LEFT_HAND_JOINT_ORDER)
right_pts = hand_joint_positions_3d(full_robot, q_full, SHARPA_RIGHT_HAND_JOINT_ORDER)
T_left_ee = frame_pose_matrix(full_robot, q_full, "L_ee")
T_right_ee = frame_pose_matrix(full_robot, q_full, "R_ee")
T_head_cam = frame_pose_matrix(bimanual_robot, q_bi, "zed_left_camera")
return left_pts, right_pts, T_head_cam, (T_left_ee, T_right_ee)