diff --git a/src/lerobot/robots/unitree_g1/g1_utils.py b/src/lerobot/robots/unitree_g1/g1_utils.py index 91f009b26..df84c8c30 100644 --- a/src/lerobot/robots/unitree_g1/g1_utils.py +++ b/src/lerobot/robots/unitree_g1/g1_utils.py @@ -43,6 +43,24 @@ def get_gravity_orientation(quaternion: list[float] | np.ndarray) -> np.ndarray: return gravity_orientation +def make_ort_session_options(intra_op_num_threads: int | None = None, inter_op_num_threads: int | None = None): + """Build quiet ONNX Runtime SessionOptions, optionally capping the CPU thread pool. + + These tiny MLP policies are latency-bound, not throughput-bound, so letting ORT grab + every core starves the real-time control loop / torch policy and causes stutter. Pass + 1 intra + 1 inter thread for lowest-latency per-step inference. + """ + import onnxruntime as ort + + so = ort.SessionOptions() + so.log_severity_level = 3 + if intra_op_num_threads is not None: + so.intra_op_num_threads = intra_op_num_threads + if inter_op_num_threads is not None: + so.inter_op_num_threads = inter_op_num_threads + return so + + class G1_29_JointArmIndex(IntEnum): # Left arm kLeftShoulderPitch = 15 diff --git a/src/lerobot/robots/unitree_g1/gr00t_locomotion.py b/src/lerobot/robots/unitree_g1/gr00t_locomotion.py index 12fe26073..21b97acea 100644 --- a/src/lerobot/robots/unitree_g1/gr00t_locomotion.py +++ b/src/lerobot/robots/unitree_g1/gr00t_locomotion.py @@ -68,9 +68,15 @@ def load_groot_policies( filename="GR00T-WholeBodyControl-Walk.onnx", ) - # Load ONNX policies - policy_balance = ort.InferenceSession(balance_path) - policy_walk = ort.InferenceSession(walk_path) + # Load ONNX policies with a capped thread pool. GR00T runs at 50 Hz in a + # background thread alongside the (torch) upper-body policy and IK; letting ORT + # grab every core starves those and makes the whole rollout stutter. These are + # small MLPs, so 1 thread is both enough and lowest-latency. + from .g1_utils import make_ort_session_options + + so = make_ort_session_options(intra_op_num_threads=1, inter_op_num_threads=1) + policy_balance = ort.InferenceSession(balance_path, sess_options=so) + policy_walk = ort.InferenceSession(walk_path, sess_options=so) logger.info("GR00T policies loaded successfully")