perf(unitree_g1): cap ONNX Runtime thread pool for locomotion policies

Add make_ort_session_options() to build quiet ORT SessionOptions with an
optional CPU thread-pool cap, and load the GR00T balance/walk ONNX policies with
1 intra + 1 inter thread. These small MLPs run at 50 Hz in a background thread;
letting ORT grab every core starves the control loop / torch policy and stutters
the rollout.
This commit is contained in:
Martino Russi
2026-07-30 11:51:24 +02:00
parent 0d383d09f2
commit f240000e70
2 changed files with 27 additions and 3 deletions
+18
View File
@@ -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
@@ -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")