diff --git a/config/eval/xarm7_starvla_eval_config.yaml b/config/eval/xarm7_starvla_eval_config.yaml index 5e0a2e6..152542b 100644 --- a/config/eval/xarm7_starvla_eval_config.yaml +++ b/config/eval/xarm7_starvla_eval_config.yaml @@ -53,5 +53,3 @@ single_task: "Pick up the black bottle and place it on the blue bag" n_episodes: 50 # Key of the camera in the robot observation dict (camera name above). camera_key: "camera" -# Enable the action safety guard (thresholds in ActionSafetyConfig). -enable_safety: false diff --git a/src/lerobot_robot_ufactory/scripts/uf_lerobot_eval.py b/src/lerobot_robot_ufactory/scripts/uf_lerobot_eval.py index b3ea4f8..16748cf 100644 --- a/src/lerobot_robot_ufactory/scripts/uf_lerobot_eval.py +++ b/src/lerobot_robot_ufactory/scripts/uf_lerobot_eval.py @@ -37,7 +37,6 @@ from lerobot.configs.policies import PreTrainedConfig from lerobot.scripts.lerobot_record import DatasetRecordConfig from lerobot.datasets.lerobot_dataset import LeRobotDataset from lerobot_robot_ufactory.utils.utils import init_keyboard_listener -from lerobot_robot_ufactory.utils.action_safety import ActionSafetyConfig, ActionSafetyGuard from lerobot_robot_ufactory.devices.umi.vive_tracker.transformations import Transformations @@ -111,11 +110,9 @@ class EvalConfig: return ["policy"] -def eval_loop(cfg: EvalConfig, relative=False, rx_continuous=False, safety_guard: ActionSafetyGuard | None = None): +def eval_loop(cfg: EvalConfig, relative=False, rx_continuous=False): init_logging() logging.info(pformat(asdict(cfg))) - if safety_guard is not None: - safety_guard.log_config() robot = make_robot_from_config(cfg.robot) @@ -234,8 +231,6 @@ def eval_loop(cfg: EvalConfig, relative=False, rx_continuous=False, safety_guard prev_robot_dict[key] = {'type': 1 if is_tcp else 0, 'pose': np.array(pose)} prev_action_dict[key] = {'type': 1 if is_tcp else 0, 'pose': np.array(pose)} - safety_halted = False - while True: start_loop_t = time.perf_counter() @@ -244,12 +239,6 @@ def eval_loop(cfg: EvalConfig, relative=False, rx_continuous=False, safety_guard print("\n********** Policy Eval Episode (Reset) **********") break - # Safety halt: stop inference and stop sending actions, - # wait for the operator to press the right arrow key to resume - if safety_halted: - precise_sleep(sleep_time_s) - continue - # Get robot observation obs = robot.get_observation() @@ -324,16 +313,6 @@ def eval_loop(cfg: EvalConfig, relative=False, rx_continuous=False, safety_guard ]) prev_action_dict[key]['pose'] = curr_action_pose - # Safety check: any violation triggers an e-stop; the action is - # NOT sent and the loop waits for the operator to press right arrow - if safety_guard is not None: - violation = safety_guard.check(robot_action_to_send, curr_robot_dict, keys) - if violation is not None: - safety_halted = True - logging.error(f"*** SAFETY HALT *** {violation}") - print(f"\n*** SAFETY HALT *** {violation}\nAction was NOT sent. Press right arrow (->) to reset and resume, ESC to exit.") - continue - robot.send_action(robot_action_to_send) dt_s = time.perf_counter() - start_loop_t @@ -354,15 +333,11 @@ def main(): parser = argparse.ArgumentParser(description='configuration args') parser.add_argument('--relative', action='store_true', help='is relative motion or not') parser.add_argument('--rx_continuous', action='store_true', help='rx continuous or not') - parser.add_argument('--enable_safety', action='store_true', help='enable action safety guard (thresholds in ActionSafetyConfig)') args, unknown = parser.parse_known_args() sys.argv = [sys.argv[0]] + unknown register_third_party_plugins() cfg = get_cfg() - # Action safety guard: tune thresholds in ActionSafetyConfig directly. - # Any violation triggers an e-stop; press right arrow to resume. - safety_guard = ActionSafetyGuard(ActionSafetyConfig(enabled=args.enable_safety)) - eval_loop(cfg, args.relative, args.rx_continuous, safety_guard) + eval_loop(cfg, args.relative, args.rx_continuous) if __name__ == "__main__": diff --git a/src/lerobot_robot_ufactory/scripts/uf_starvla_eval.py b/src/lerobot_robot_ufactory/scripts/uf_starvla_eval.py index 2204731..4364d7c 100644 --- a/src/lerobot_robot_ufactory/scripts/uf_starvla_eval.py +++ b/src/lerobot_robot_ufactory/scripts/uf_starvla_eval.py @@ -41,7 +41,6 @@ from lerobot.utils.import_utils import register_third_party_plugins from lerobot.utils.robot_utils import precise_sleep from lerobot.utils.utils import init_logging -from lerobot_robot_ufactory.utils.action_safety import ActionSafetyConfig, ActionSafetyGuard from lerobot_robot_ufactory.utils.starvla_ws_client import WebsocketClientPolicy from lerobot_robot_ufactory.utils.utils import init_keyboard_listener @@ -61,8 +60,6 @@ class StarVLAEvalConfig: n_episodes: int = 50 # Key of the camera in the robot observation dict (camera name in robot config). camera_key: str = "camera" - # Enable the action safety guard (thresholds in ActionSafetyConfig). - enable_safety: bool = False def _build_state(obs: dict) -> np.ndarray: @@ -81,11 +78,9 @@ def _build_action_dict(action: np.ndarray) -> dict: return action_dict -def eval_loop(cfg: StarVLAEvalConfig, safety_guard: ActionSafetyGuard | None = None): +def eval_loop(cfg: StarVLAEvalConfig): init_logging() logging.info(pformat(asdict(cfg))) - if safety_guard is not None: - safety_guard.log_config() robot = make_robot_from_config(cfg.robot) robot.connect() @@ -159,20 +154,6 @@ def eval_loop(cfg: StarVLAEvalConfig, safety_guard: ActionSafetyGuard | None = N start_loop_t = time.perf_counter() action_dict = _build_action_dict(action) - # Safety check: joint-space pose limits do not apply (the - # guard only checks TCP-pose actions and gripper NaN/Inf), - # but it still catches non-finite gripper commands. - if safety_guard is not None: - violation = safety_guard.check(action_dict, {}, [""]) - if violation is not None: - logging.error(f"*** SAFETY HALT *** {violation}") - print( - f"\n*** SAFETY HALT *** {violation}\n" - "Action was NOT sent. Press right arrow (->) to reset and resume, ESC to exit." - ) - events["reset"] = True - break - robot.send_action(action_dict) dt_s = time.perf_counter() - start_loop_t @@ -197,10 +178,7 @@ def get_cfg(cfg: StarVLAEvalConfig) -> StarVLAEvalConfig: def main(): register_third_party_plugins() cfg = get_cfg() - # Action safety guard: tune thresholds in ActionSafetyConfig directly. - # Any violation triggers an e-stop; press right arrow to resume. - safety_guard = ActionSafetyGuard(ActionSafetyConfig(enabled=cfg.enable_safety)) - eval_loop(cfg, safety_guard) + eval_loop(cfg) if __name__ == "__main__": diff --git a/src/lerobot_robot_ufactory/utils/action_safety.py b/src/lerobot_robot_ufactory/utils/action_safety.py deleted file mode 100644 index 7fe6b80..0000000 --- a/src/lerobot_robot_ufactory/utils/action_safety.py +++ /dev/null @@ -1,130 +0,0 @@ -"""Model-agnostic action safety checks. - -Validates the final action dict right before it is sent to the robot; any -failed check triggers an emergency halt. Only relies on the action key -naming convention: `[prefix.]pose.x/y/z/rx/ry/rz`, `*.gripper.pos`, so it -works with any policy (ACT / DP / pi0 / custom backbones) and with both -single-arm and multi-arm setups. -""" - -import logging -import math -from dataclasses import dataclass - -import numpy as np - -from lerobot_robot_ufactory.devices.umi.vive_tracker.transformations import Transformations - -POSE_AXES = ("x", "y", "z", "rx", "ry", "rz") - - -@dataclass -class ActionSafetyConfig: - """Configuration for ActionSafetyGuard. - - All limits are compared against the robot's current actual pose per - control step. `workspace_min`/`workspace_max` use the same unit as the - pose commands (e.g. mm) and must be provided together. - """ - - enabled: bool = True - max_step_mm: float = 25.0 - max_rot_step: float = 0.35 - workspace_min: list | None = None - workspace_max: list | None = None - - def __post_init__(self): - if (self.workspace_min is None) != (self.workspace_max is None): - raise ValueError("workspace_min and workspace_max must be provided together") - - -def _rotation_delta_norm(rot_prev, rot_curr) -> float: - """Relative rotation angle (rad) between two axis-angle rotations. - - Uses rotation matrices for the diff to avoid the ±π discontinuity - of subtracting raw rotvecs. - """ - R_prev = Transformations.rxryrz_to_rotation_matrix(*rot_prev) - R_curr = Transformations.rxryrz_to_rotation_matrix(*rot_curr) - R_delta = R_prev.T @ R_curr - delta = Transformations.rotation_matrix_to_rxryrz(R_delta) - return float(np.linalg.norm(delta)) - - -class ActionSafetyGuard: - """Safety checks for the action dict about to be sent to the robot. - - check() returns None when the action is safe; otherwise it returns a - human-readable violation reason (the caller should trigger an e-stop). - """ - - def __init__(self, config: ActionSafetyConfig | None = None): - config = config or ActionSafetyConfig() - self.config = config - self.enabled = config.enabled - self.max_step_mm = config.max_step_mm - self.max_rot_step = config.max_rot_step - self.workspace_min = None if config.workspace_min is None else np.asarray(config.workspace_min, dtype=np.float64) - self.workspace_max = None if config.workspace_max is None else np.asarray(config.workspace_max, dtype=np.float64) - - def check(self, action: dict, curr_robot_dict: dict, keys) -> str | None: - """Check the action dict. `action` holds absolute pose commands and - `curr_robot_dict` holds each arm's current actual pose.""" - if not self.enabled: - return None - - for key in keys: - prefix = f".{key}" if key else "" - pose_keys = [f"{prefix}pose.{axis}" for axis in POSE_AXES] - if not all(k in action for k in pose_keys): - continue # not TCP-pose controlled (e.g. joint-only arm), skip pose checks - - target = np.array([action[k] for k in pose_keys], dtype=np.float64) - - # 1. Numeric validity: NaN / Inf - if not np.all(np.isfinite(target)): - return f"[{key or 'arm'}] action contains NaN/Inf: {target.tolist()}" - - # 2. Single-step delta limit, measured from the robot's current actual pose - curr = curr_robot_dict.get(key) - if curr is not None and curr.get("type") == 1: - curr_pose = np.asarray(curr["pose"], dtype=np.float64) - pos_dist = float(np.linalg.norm(target[:3] - curr_pose[:3])) - if pos_dist > self.max_step_mm: - return ( - f"[{key or 'arm'}] position step {pos_dist:.1f}mm exceeds limit " - f"{self.max_step_mm}mm (current {curr_pose[:3].tolist()} -> target {target[:3].tolist()})" - ) - rot_delta = _rotation_delta_norm(curr_pose[3:6], target[3:6]) - if rot_delta > self.max_rot_step: - return ( - f"[{key or 'arm'}] rotation step {rot_delta:.3f}rad exceeds limit " - f"{self.max_rot_step}rad" - ) - - # 3. Workspace bounding box (optional) - if self.workspace_min is not None: - pos = target[:3] - if np.any(pos < self.workspace_min) or np.any(pos > self.workspace_max): - return ( - f"[{key or 'arm'}] target position {pos.tolist()} outside workspace " - f"[{self.workspace_min.tolist()}, {self.workspace_max.tolist()}]" - ) - - # Gripper numeric check - for k, v in action.items(): - if "gripper" in k and isinstance(v, (int, float)) and not math.isfinite(v): - return f"[{k}] gripper action contains NaN/Inf: {v}" - - return None - - def log_config(self): - ws = ( - f"[{self.workspace_min.tolist()}, {self.workspace_max.tolist()}]" - if self.workspace_min is not None - else "not set" - ) - logging.info( - f"ActionSafetyGuard: max_step={self.max_step_mm}mm, " - f"max_rot_step={self.max_rot_step}rad, workspace={ws}" - )