From 553e3f00dc2bcb275539838913c8a81fdd0cc40a Mon Sep 17 00:00:00 2001 From: skywhite1024 <129768272+skywhite1024@users.noreply.github.com> Date: Fri, 21 Aug 2026 00:28:40 +0800 Subject: [PATCH] feat(action-engine): add runtime grounding and action adapters --- .../gen_sim/action_engine/runtime/actions.py | 1511 +++++++++++ .../action_engine/runtime/atomic_compat.py | 84 + .../gen_sim/action_engine/runtime/frames.py | 165 ++ .../runtime/grasp_collision_cache.py | 330 +++ .../action_engine/runtime/grounding.py | 2356 +++++++++++++++++ .../action_engine/runtime/predicates.py | 746 ++++++ .../action_engine/runtime/robot_parts.py | 34 + .../action_engine/runtime/solver_compat.py | 234 ++ .../action_engine/capabilities/__init__.py | 19 + .../capabilities/test_atomic_v2.py | 233 ++ .../action_engine/planning/test_linker.py | 407 +++ .../gen_sim/action_engine/runtime/__init__.py | 19 + .../action_engine/runtime/test_actions.py | 835 ++++++ .../runtime/test_atomic_compat.py | 87 + .../runtime/test_grasp_collision_cache.py | 354 +++ 15 files changed, 7414 insertions(+) create mode 100644 embodichain/gen_sim/action_engine/runtime/actions.py create mode 100644 embodichain/gen_sim/action_engine/runtime/atomic_compat.py create mode 100644 embodichain/gen_sim/action_engine/runtime/frames.py create mode 100644 embodichain/gen_sim/action_engine/runtime/grasp_collision_cache.py create mode 100644 embodichain/gen_sim/action_engine/runtime/grounding.py create mode 100644 embodichain/gen_sim/action_engine/runtime/predicates.py create mode 100644 embodichain/gen_sim/action_engine/runtime/robot_parts.py create mode 100644 embodichain/gen_sim/action_engine/runtime/solver_compat.py create mode 100644 tests/gen_sim/action_engine/capabilities/__init__.py create mode 100644 tests/gen_sim/action_engine/capabilities/test_atomic_v2.py create mode 100644 tests/gen_sim/action_engine/planning/test_linker.py create mode 100644 tests/gen_sim/action_engine/runtime/__init__.py create mode 100644 tests/gen_sim/action_engine/runtime/test_actions.py create mode 100644 tests/gen_sim/action_engine/runtime/test_atomic_compat.py create mode 100644 tests/gen_sim/action_engine/runtime/test_grasp_collision_cache.py diff --git a/embodichain/gen_sim/action_engine/runtime/actions.py b/embodichain/gen_sim/action_engine/runtime/actions.py new file mode 100644 index 000000000..3df18ef09 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/actions.py @@ -0,0 +1,1511 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Adapt Action Engine requests to the shared typed atomic-action planner.""" + +from __future__ import annotations + +from collections.abc import Mapping +from copy import deepcopy +from dataclasses import replace +import math +from typing import Any + +import torch + +from embodichain.gen_sim.action_engine.capabilities import ( + AtomicCapability, + build_atomic_capability_registry, +) +from embodichain.gen_sim.action_engine.config import default_runtime_policy +from embodichain.lab.sim.atomic_actions import ( + ActionBinding, + ActionInvocation, + ActionPlan, + AntipodalAffordance, + AtomicActionEngine, + ControlPartCommandProfile, + CoordinatedPickGoal, + DynamicCollisionMode, + EndEffectorPoseGoal, + EntityState, + ExecutionSession, + MotionPolicy, + ObjectSemantics, + PlanningContext, + RecoveryPolicy, + RobotObservation, + RigidObjectSceneProvider, + SceneProvider, + SceneSnapshot, + StateDelta, +) +from embodichain.lab.sim.planners import ( + CuroboPlannerCfg, + CuroboWorldCfg, + MotionGenCfg, + MotionGenerator, + ToppraPlannerCfg, +) +from embodichain.toolkits.graspkit.pg_grasp import ( + AntipodalSamplerCfg, + GraspGeneratorCfg, + GripperCollisionCfg, +) +from embodichain.utils.logger import log_info + +from .grasp_collision_cache import ensure_vhacd_grasp_collision_cache +from .models import ActionOutcome, GroundedAction +from .state import ExecutionState + +__all__ = ["AtomicActionAdapter"] + + +_DEFAULT_PLANNER_POLICY: dict[str, Any] = { + "backend": "curobo", + "single_arm_strategy": "motion_gen", + "coordinated_strategy": "ik_interp", + "fallback_strategy": "ik_interp", + "allow_fallback": True, + "dynamic_collision": False, + "static_obstacle_uids": [], + "dynamic_obstacle_uids": [], + "curobo": { + "log_level": "error", + "obstacle_representation": "cuboid", + "multi_env": False, + "use_cuda_graph": True, + "preserve_plan_samples": False, + "max_attempts": 5, + "collision_activation_distance": 0.01, + }, +} + +# Preserve cuRobo's fixed world shape while disabling intentional-contact objects. +_COLLISION_PARKING_Z_OFFSET = -100.0 + + +def _collision_cache_for_world( + representation: str, obstacle_count: int +) -> dict[str, int]: + """Size cuRobo's fixed collision cache for the generated scene.""" + cache = {"cuboid": 8, "mesh": 2} + if representation in cache: + cache[representation] = max(cache[representation], obstacle_count) + return cache + + +def _supported_kwargs(config_type: type, values: Mapping[str, Any]) -> dict[str, Any]: + names: set[str] = set() + for cls in reversed(config_type.__mro__): + names.update(getattr(cls, "__annotations__", {})) + return {key: value for key, value in values.items() if key in names} + + +def _as_hand_qpos(value: Any, dof: int, device: Any) -> torch.Tensor: + if dof == 0: + return torch.empty(0, dtype=torch.float32, device=device) + result = torch.as_tensor(value, dtype=torch.float32, device=device).flatten() + if result.numel() == 0: + return torch.zeros(dof, dtype=torch.float32, device=device) + if result.numel() == 1: + return result.repeat(dof) + if result.numel() >= dof: + return result[:dof] + repeats = (dof + result.numel() - 1) // result.numel() + return result.repeat(repeats)[:dof] + + +def _diagonal_approach_direction( + horizontal: torch.Tensor, + *, + vertical: float = -1.0, +) -> torch.Tensor: + """Combine one normalized horizontal role direction with a vertical component.""" + horizontal = horizontal.to(dtype=torch.float32) + norm = torch.linalg.vector_norm(horizontal) + if float(norm) <= 1.0e-6: + raise ValueError("Handover role direction must be non-zero.") + horizontal = horizontal / norm + direction = torch.stack( + (horizontal[0], horizontal[1], horizontal.new_tensor(float(vertical))) + ) + return direction / torch.linalg.vector_norm(direction) + + +class AtomicActionAdapter: + """Own the shared atomic engine and preserve Action Engine runtime contracts.""" + + def __init__( + self, + env: Any, + *, + grasp_policy: Mapping[str, Any] | None = None, + planner_policy: Mapping[str, Any] | None = None, + capability_registry: Any | None = None, + scene_provider: SceneProvider | None = None, + ) -> None: + self.env = env + self.num_envs = int(env.num_envs) + self.device = env.device + if grasp_policy is None: + profile = str(getattr(env, "agent_robot_profile", "dual_ur10")) + grasp_policy = default_runtime_policy(profile).grasp + grasp_policy = { + **grasp_policy, + **(getattr(env, "agent_grasp_runtime_defaults", {}) or {}), + } + self.grasp_policy = deepcopy(dict(grasp_policy)) + self.planner_policy = deepcopy(_DEFAULT_PLANNER_POLICY) + if planner_policy is not None: + self._merge_planner_policy(self.planner_policy, planner_policy) + if not self.planner_policy.get("static_obstacle_uids"): + configured = getattr(env, "agent_static_obstacle_uids", ()) or () + if configured: + self.planner_policy["static_obstacle_uids"] = [ + str(uid) for uid in configured + ] + else: + get_rigid_object = getattr(env.sim, "get_rigid_object", None) + if callable(get_rigid_object) and get_rigid_object("table") is not None: + self.planner_policy["static_obstacle_uids"] = ["table"] + self.capabilities = capability_registry or build_atomic_capability_registry() + self._motion_generator: MotionGenerator | None = None + self._atomic_engine: AtomicActionEngine | None = None + self._semantics: dict[str, ObjectSemantics] = {} + self._scene_time = 0.0 + if scene_provider is not None and not isinstance(scene_provider, SceneProvider): + raise TypeError("scene_provider must implement SceneProvider.") + self.scene_provider = scene_provider or self._build_scene_provider() + + @staticmethod + def _merge_planner_policy( + target: dict[str, Any], + update: Mapping[str, Any], + ) -> None: + for key, value in update.items(): + if isinstance(value, Mapping) and isinstance(target.get(key), dict): + AtomicActionAdapter._merge_planner_policy(target[key], value) + else: + target[key] = deepcopy(value) + + def initial_state(self) -> ExecutionState: + """Capture the initial full-robot planning seed.""" + return ExecutionState(last_qpos=self.env.robot.get_qpos().clone()) + + def start_session( + self, + grounded: GroundedAction, + state: ExecutionState | None = None, + ) -> ExecutionSession: + """Start one closed-loop AtomicAction session from live scene state. + + ProgramExecutor may continue using its compatibility scheduler for + compound and per-arm merged trajectories. New callers can use this + boundary to adopt feedback-driven execution without constructing + private planning contexts. + """ + capability = self.capabilities.require_executable(grounded.action_class) + state = state or self.initial_state() + grounded = self._select_upright_transport_yaw(grounded, state) + context = self._planning_context(state, grounded) + invocation = self._invocation(grounded, capability) + return self._engine().start((invocation,), context) + + def _build_scene_provider(self) -> SceneProvider | None: + """Create the shared live rigid-object provider when entities are available.""" + sim = getattr(self.env, "sim", None) + if sim is None: + return None + dynamic_uids = tuple( + str(uid) for uid in self.planner_policy.get("dynamic_obstacle_uids", ()) + ) + list_uids = getattr(sim, "get_rigid_object_uid_list", None) + uids = tuple(str(uid) for uid in list_uids()) if callable(list_uids) else () + if not uids: + uids = dynamic_uids + get_rigid_object = getattr(sim, "get_rigid_object", None) + if not callable(get_rigid_object): + return None + entities = { + uid: entity for uid in uids if (entity := get_rigid_object(uid)) is not None + } + if not entities: + return None + collision_uids = ( + dynamic_uids + if bool(self.planner_policy.get("dynamic_collision", False)) + else () + ) + return RigidObjectSceneProvider( + entities, + collision_entity_ids=collision_uids, + ) + + def semantics(self, uid: str) -> ObjectSemantics: + """Build object semantics once while retaining the live entity handle.""" + cached = self._semantics.get(uid) + if cached is not None: + return cached + entity = self.env.sim.get_rigid_object(uid) + if entity is None: + raise ValueError(f"Unknown grasp target {uid!r}.") + vertices = entity.get_vertices(env_ids=[0], scale=True) + triangles = entity.get_triangles(env_ids=[0]) + if isinstance(vertices, (tuple, list)): + vertices = vertices[0] + if isinstance(triangles, (tuple, list)): + triangles = triangles[0] + vertices = torch.as_tensor(vertices, dtype=torch.float32) + triangles = torch.as_tensor(triangles, dtype=torch.int64) + if vertices.ndim == 3 and vertices.shape[0] == 1: + vertices = vertices[0] + if triangles.ndim == 3 and triangles.shape[0] == 1: + triangles = triangles[0] + if vertices.ndim != 2 or vertices.shape[-1] != 3 or vertices.numel() == 0: + raise ValueError(f"Object {uid!r} has invalid mesh vertices.") + if triangles.ndim != 2 or triangles.shape[-1] != 3 or triangles.numel() == 0: + raise ValueError(f"Object {uid!r} has invalid mesh triangles.") + + grasp_options = self.grasp_policy + sampler = AntipodalSamplerCfg( + n_sample=int(grasp_options["antipodal_n_sample"]), + max_angle=float(grasp_options["antipodal_max_angle"]), + max_length=float(grasp_options["max_open_length"]), + min_length=float(grasp_options["min_open_length"]), + ) + generator = GraspGeneratorCfg( + viser_port=int(grasp_options["viser_port"]), + antipodal_sampler_cfg=sampler, + max_deviation_angle=float(grasp_options["max_deviation_angle"]), + n_deviated_approach_directions=int( + grasp_options["n_deviated_approach_directions"] + ), + ) + max_hulls = int(grasp_options["max_decomposition_hulls"]) + collision = GripperCollisionCfg( + max_open_length=float(grasp_options["max_open_length"]), + finger_length=float(grasp_options["finger_length"]), + point_sample_dense=float(grasp_options["point_sample_dense"]), + max_decomposition_hulls=max_hulls, + ) + cache_result = ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=max_hulls, + ) + if cache_result.status != "hit": + log_info(f"Prepared V-HACD grasp cache for {uid!r}: {cache_result.status}.") + + semantics = ObjectSemantics( + label=uid, + entity=entity, + geometry={"mesh_vertices": vertices, "mesh_triangles": triangles}, + affordance=AntipodalAffordance( + object_label=uid, + mesh_vertices=vertices, + mesh_triangles=triangles, + generator_cfg=generator, + gripper_collision_cfg=collision, + force_reannotate=bool(grasp_options["force_grasp_reannotate"]), + ), + ) + self._semantics[uid] = semantics + return semantics + + def plan( + self, + grounded: GroundedAction, + state: ExecutionState | None = None, + ) -> ActionOutcome: + """Plan one grounded primitive through the mainline typed contract.""" + capability = self.capabilities.require_executable(grounded.action_class) + state = state or self.initial_state() + grounded = self._select_upright_transport_yaw(grounded, state) + context = self._planning_context(state, grounded) + invocation = self._invocation(grounded, capability) + plan = self._engine().plan(invocation, context) + selected_positions = self._positions_with_agent_holds( + plan, + grounded, + capability, + ) + primary_success = plan.plan_success.to(self.device) + reachability_search = None + if bool(grounded.motion_policy.get("retreat_reachability_search", False)): + ( + grounded, + selected_positions, + primary_success, + reachability_search, + ) = self._search_reachable_retreat( + grounded=grounded, + capability=capability, + state=state, + context=context, + invocation=invocation, + initial_positions=selected_positions, + initial_success=primary_success, + ) + invocation = replace(invocation, goal=grounded.target) + combined_success = primary_success.clone() + fallback_plan: ActionPlan | None = None + use_fallback = torch.zeros_like(combined_success) + fallback_attempted = torch.zeros_like(combined_success) + fallback_success = torch.zeros_like(combined_success) + + fallback_strategy = self.planner_policy.get("fallback_strategy") + collision_safety = str(grounded.motion_policy.get("collision_safety", "auto")) + fallback_allowed = bool(self.planner_policy.get("allow_fallback", True)) and ( + collision_safety != "required" + ) + if ( + fallback_allowed + and invocation.motion_policy.strategy == "motion_gen" + and fallback_strategy in {"ik_interp"} + and not bool(combined_success.all()) + ): + fallback_attempted = ~primary_success + fallback_policy = replace( + invocation.motion_policy, + strategy=str(fallback_strategy), + dynamic_collision_mode=DynamicCollisionMode.OFF, + plan_opts=None, + ) + fallback_plan = self._engine().plan( + replace(invocation, motion_policy=fallback_policy), + context, + ) + fallback_positions = self._positions_with_agent_holds( + fallback_plan, + grounded, + capability, + ) + fallback_success = fallback_plan.plan_success.to(self.device) + use_fallback = fallback_attempted & fallback_success + selected_positions = self._merge_plan_rows( + selected_positions, + fallback_positions, + use_fallback, + state.last_qpos, + ) + combined_success |= fallback_plan.plan_success.to(self.device) + + options = invocation.skill_options + if capability.config_materializer == "handover": + combined_success &= self._handover_receiver_hold_mask( + selected_positions, + grounded, + options, + tolerance=float( + grounded.motion_policy.get( + "receiver_hold_joint_tolerance", + 2.0e-3, + ) + ), + ) + + terminal_qpos = ( + selected_positions[:, -1] + if selected_positions.shape[1] + else state.last_qpos + ) + primary_rows = combined_success & primary_success + projected_task = plan.expected_effects.apply( + context.task, + primary_rows, + ) + held_keys = set(plan.expected_effects.held_object_updates) + if fallback_plan is not None: + fallback_rows = combined_success & use_fallback + projected_task = fallback_plan.expected_effects.apply( + projected_task, + fallback_rows, + ) + held_keys.update(fallback_plan.expected_effects.held_object_updates) + committed_effects = StateDelta( + held_object_updates={ + key: projected_task.held_objects.get(key) for key in held_keys + }, + ) + next_state = ExecutionState.from_task_state( + projected_task, + last_qpos=torch.where( + combined_success[:, None], terminal_qpos, state.last_qpos + ), + ) + return ActionOutcome( + trajectory=selected_positions, + success=combined_success, + next_state=next_state, + grounded=grounded, + prior_state=state, + expected_effects=committed_effects, + planner_trace={ + **self._planner_trace( + grounded=grounded, + invocation=invocation, + context=context, + state=state, + primary_success=primary_success, + fallback_allowed=fallback_allowed, + fallback_strategy=( + str(fallback_strategy) + if invocation.motion_policy.strategy == "motion_gen" + and fallback_strategy in {"ik_interp"} + else None + ), + fallback_attempted=fallback_attempted, + fallback_success=fallback_success, + fallback_used=use_fallback, + reachability_search=reachability_search, + ), + # Auditability takes precedence over compactness here: every + # selected planner route retains its complete joint path. + "planned_trajectory": selected_positions.detach().clone(), + }, + ) + + def _search_reachable_retreat( + self, + *, + grounded: GroundedAction, + capability: AtomicCapability, + state: ExecutionState, + context: PlanningContext, + invocation: ActionInvocation, + initial_positions: torch.Tensor, + initial_success: torch.Tensor, + ) -> tuple[GroundedAction, torch.Tensor, torch.Tensor, dict[str, Any]]: + """Select the highest row-local retreat accepted by the live planner.""" + candidates = self._retreat_search_targets(grounded) + target = getattr(grounded.target, "xpos", None) + if not isinstance(target, torch.Tensor) or len(candidates) <= 1: + return ( + grounded, + initial_positions, + initial_success, + { + "strategy": "bounded_motion_planner", + "attempts": [], + "selected_target_z": ( + None + if not isinstance(target, torch.Tensor) + else target[:, 2, 3] + ), + }, + ) + + selected_target = candidates[0][1].clone() + selected_positions = initial_positions + success = initial_success.clone() + attempts: list[dict[str, Any]] = [ + { + "candidate": candidates[0][0], + "target_z": candidates[0][1][:, 2, 3].detach().clone(), + "success": initial_success.detach().clone(), + } + ] + for label, candidate_target in candidates[1:]: + unresolved = ~success + if not bool(unresolved.any()): + break + row_target = torch.where( + unresolved[:, None, None], + candidate_target, + selected_target, + ) + candidate_grounded = replace( + grounded, + target=EndEffectorPoseGoal(xpos=row_target), + ) + candidate_invocation = replace( + invocation, + goal=candidate_grounded.target, + ) + candidate_plan = self._engine().plan(candidate_invocation, context) + candidate_positions = self._positions_with_agent_holds( + candidate_plan, + candidate_grounded, + capability, + ) + candidate_success = candidate_plan.plan_success.to(self.device) + selected_rows = unresolved & candidate_success + selected_positions = self._merge_plan_rows( + selected_positions, + candidate_positions, + selected_rows, + state.last_qpos, + ) + selected_target = torch.where( + selected_rows[:, None, None], + candidate_target, + selected_target, + ) + success |= candidate_success + attempts.append( + { + "candidate": label, + "target_z": candidate_target[:, 2, 3].detach().clone(), + "success": candidate_success.detach().clone(), + } + ) + + metadata = { + "retreat_selected_target_z": selected_target[:, 2, 3].detach().clone(), + "retreat_reachability_found": success.detach().clone(), + } + selected_grounded = replace( + grounded, + target=EndEffectorPoseGoal(xpos=selected_target), + cfg={**grounded.cfg, **metadata}, + motion_policy={**grounded.motion_policy, **metadata}, + ) + return ( + selected_grounded, + selected_positions, + success, + { + "strategy": "bounded_motion_planner", + "attempts": attempts, + "selected_target_z": selected_target[:, 2, 3].detach().clone(), + }, + ) + + def _retreat_search_targets( + self, + grounded: GroundedAction, + ) -> list[tuple[str, torch.Tensor]]: + """Build bounded height and baseward retreat candidates from live poses.""" + target = getattr(grounded.target, "xpos", None) + reference = grounded.motion_policy.get("retreat_reference_pose") + if not isinstance(target, torch.Tensor) or not isinstance( + reference, torch.Tensor + ): + return [] + target = target.to(device=self.device, dtype=torch.float32) + reference = reference.to(device=self.device, dtype=torch.float32) + if target.shape == (4, 4): + target = target.unsqueeze(0).repeat(self.num_envs, 1, 1) + if reference.shape == (4, 4): + reference = reference.unsqueeze(0).repeat(self.num_envs, 1, 1) + expected = (self.num_envs, 4, 4) + if target.shape != expected or reference.shape != expected: + return [] + + sample_count = int(grounded.cfg.get("retreat_search_samples", 6)) + if not 2 <= sample_count <= 16: + raise ValueError("retreat_search_samples must be in [2, 16].") + minimum_height = float(grounded.cfg.get("minimum_retreat_height", 0.05)) + if not math.isfinite(minimum_height) or minimum_height < 0.0: + raise ValueError("minimum_retreat_height must be finite and non-negative.") + desired_height = torch.clamp( + target[:, 2, 3] - reference[:, 2, 3], + min=0.0, + ) + minimum = torch.minimum( + desired_height, + torch.full_like(desired_height, minimum_height), + ) + fractions = torch.linspace( + 1.0, + 0.0, + sample_count, + dtype=target.dtype, + device=target.device, + ) + heights = ( + minimum[:, None] + (desired_height - minimum)[:, None] * fractions[None] + ) + candidates: list[tuple[str, torch.Tensor]] = [("requested", target.clone())] + for index in range(1, sample_count): + candidate = target.clone() + candidate[:, 2, 3] = reference[:, 2, 3] + heights[:, index] + candidates.append((f"height_{index}", candidate)) + + from .frames import arm_base_poses + + left_base, right_base = arm_base_poses(self.env) + base = left_base if grounded.arm == "left_arm" else right_base + direction = base[:, :2, 3] - reference[:, :2, 3] + norm = torch.linalg.vector_norm(direction, dim=1, keepdim=True) + direction = torch.where( + norm > 1.0e-6, + direction / torch.clamp(norm, min=1.0e-6), + torch.zeros_like(direction), + ) + distance = float(grounded.cfg.get("retreat_distance", 0.10)) + if not math.isfinite(distance) or distance < 0.0: + raise ValueError("retreat_distance must be finite and non-negative.") + for index in range(sample_count): + candidate = target.clone() + candidate[:, :2, 3] = reference[:, :2, 3] + direction * distance + candidate[:, 2, 3] = reference[:, 2, 3] + heights[:, index] + candidates.append((f"baseward_{index}", candidate)) + return candidates + + def _planner_trace( + self, + *, + grounded: GroundedAction, + invocation: ActionInvocation, + context: PlanningContext, + state: ExecutionState, + primary_success: torch.Tensor, + fallback_allowed: bool, + fallback_strategy: str | None, + fallback_attempted: torch.Tensor, + fallback_success: torch.Tensor, + fallback_used: torch.Tensor, + reachability_search: Mapping[str, Any] | None = None, + ) -> dict[str, Any]: + """Build compact per-row evidence for the planner route actually used.""" + exclusions = self._collision_exclusion_masks(grounded, state) + obstacle_positions = { + uid: context.scene.entities[uid].pose[:, :3, 3].detach().clone() + for uid in context.scene.collision_entity_ids + } + revisions = torch.as_tensor( + context.scene.collision_world_revisions(self.num_envs), + dtype=torch.int64, + device=self.device, + ) + trace = { + "action_class": grounded.action_class, + "arm": grounded.arm, + "planner": str(self.planner_policy["backend"]), + "primary_strategy": invocation.motion_policy.strategy, + "dynamic_collision_mode": invocation.motion_policy.dynamic_collision_mode.value, + "primary_success": primary_success.detach().clone(), + "fallback_allowed": fallback_allowed, + "fallback_strategy": fallback_strategy, + "fallback_attempted": fallback_attempted.detach().clone(), + "fallback_success": fallback_success.detach().clone(), + "fallback_used": fallback_used.detach().clone(), + "search_budget": { + "primary_max_attempts": int( + self.planner_policy.get("curobo", {}).get("max_attempts", 1) + ), + "fallback_enabled": bool(fallback_allowed), + }, + "collision_world_revision": revisions, + "collision_obstacle_positions": obstacle_positions, + "collision_exclusions": { + uid: mask.detach().clone() for uid, mask in exclusions.items() + }, + } + if reachability_search is not None: + trace["reachability_search"] = deepcopy(dict(reachability_search)) + options = invocation.skill_options + object_part = getattr(options, "pick_object_part", None) + approach_direction = getattr(options, "approach_direction", None) + if object_part is None: + object_part = getattr(options, "receive_pick_object_part", None) + approach_direction = getattr( + options, + "receive_approach_direction", + approach_direction, + ) + if object_part is not None: + grasp_policy: dict[str, Any] = {"object_part": str(object_part)} + if isinstance(approach_direction, torch.Tensor): + direction = approach_direction.to(dtype=torch.float32) + norm = torch.linalg.vector_norm(direction) + if bool(torch.isfinite(norm)) and float(norm) > 0.0: + grasp_policy["approach_direction"] = ( + (direction / norm).detach().cpu().tolist() + ) + trace["grasp_policy"] = grasp_policy + return trace + + def _select_upright_transport_yaw( + self, + grounded: GroundedAction, + state: ExecutionState, + ) -> GroundedAction: + """Choose the closest IK-feasible yaw for an upright object target.""" + sample_count = int(grounded.cfg.get("upright_yaw_samples", 1)) + capability = self.capabilities.get(grounded.action_class) + if ( + capability.target_materializer != "semantic_held_object" + or sample_count <= 1 + ): + return grounded + target_pose = getattr(grounded.target, "object_target_pose", None) + if not isinstance(target_pose, torch.Tensor): + return grounded + target_pose = target_pose.to(device=self.device, dtype=torch.float32) + if target_pose.shape == (4, 4): + target_pose = target_pose.unsqueeze(0).repeat(self.num_envs, 1, 1) + if target_pose.shape != (self.num_envs, 4, 4): + raise ValueError( + "Upright transport target must have shape (4, 4) or (N, 4, 4)." + ) + + arm_part, _, _ = self._parts(grounded.arm) + held = state.get_held_object(arm_part) + if held is None: + return grounded + object_to_eef = held.object_to_eef.to( + device=self.device, + dtype=target_pose.dtype, + ) + if object_to_eef.shape == (4, 4): + object_to_eef = object_to_eef.unsqueeze(0).repeat(self.num_envs, 1, 1) + variants = self._upright_yaw_variants(target_pose, sample_count) + eef_variants = torch.matmul(variants, object_to_eef[:, None]) + joint_ids = list(self.env.robot.get_joint_ids(name=arm_part)) + start_qpos = state.last_qpos[:, joint_ids] + seeds = start_qpos[:, None].expand(-1, sample_count, -1) + success, qpos = self.env.robot.compute_batch_ik( + pose=eef_variants, + name=arm_part, + joint_seed=seeds, + ) + success = torch.as_tensor( + success, + dtype=torch.bool, + device=self.device, + ).reshape(self.num_envs, sample_count) + qpos = torch.as_tensor(qpos, dtype=torch.float32, device=self.device) + success &= torch.isfinite(qpos).all(dim=-1) + distance = torch.linalg.vector_norm(qpos - seeds, dim=-1) + distance = torch.where( + success, + distance, + torch.full_like(distance, torch.inf), + ) + best = distance.argmin(dim=1) + env_ids = torch.arange(self.num_envs, device=self.device) + selected = variants[env_ids, best] + selected = torch.where( + success.any(dim=1)[:, None, None], + selected, + target_pose, + ) + return replace( + grounded, + target=replace(grounded.target, object_target_pose=selected), + target_object_pose=selected, + ) + + @staticmethod + def _upright_yaw_variants( + target_pose: torch.Tensor, + sample_count: int, + ) -> torch.Tensor: + signed_steps = [0] + for step in range(1, (sample_count + 1) // 2): + signed_steps.extend((step, -step)) + if sample_count % 2 == 0: + signed_steps.append(sample_count // 2) + angles = target_pose.new_tensor(signed_steps) * (2.0 * math.pi / sample_count) + yaw = target_pose.new_zeros((sample_count, 3, 3)) + yaw[:, 0, 0] = torch.cos(angles) + yaw[:, 0, 1] = -torch.sin(angles) + yaw[:, 1, 0] = torch.sin(angles) + yaw[:, 1, 1] = torch.cos(angles) + yaw[:, 2, 2] = 1.0 + variants = target_pose[:, None].repeat(1, sample_count, 1, 1) + variants[:, :, :3, :3] = torch.matmul(yaw[None], target_pose[:, None, :3, :3]) + return variants + + def _planning_context( + self, + state: ExecutionState, + grounded: GroundedAction, + ) -> PlanningContext: + qpos = state.last_qpos.to(device=self.device, dtype=torch.float32) + get_qvel = getattr(self.env.robot, "get_qvel", None) + qvel = get_qvel() if callable(get_qvel) else None + if not isinstance(qvel, torch.Tensor) or qvel.shape != qpos.shape: + qvel = torch.zeros_like(qpos) + else: + qvel = qvel.to(device=self.device, dtype=qpos.dtype) + return PlanningContext( + robot=RobotObservation(timestamp=self._scene_time, qpos=qpos, qvel=qvel), + task=state.to_task_state(), + scene=self._scene_snapshot(grounded, state), + env_ids=torch.arange( + self.num_envs, + dtype=torch.long, + device=self.device, + ), + control_dt=float(getattr(self.env, "step_dt", 1.0 / 60.0)), + ) + + def _scene_snapshot( + self, + grounded: GroundedAction, + state: ExecutionState, + ) -> SceneSnapshot: + dynamic_uids = tuple( + str(uid) for uid in self.planner_policy.get("dynamic_obstacle_uids", ()) + ) + env_ids = torch.arange( + self.num_envs, + dtype=torch.long, + device=self.device, + ) + if self.scene_provider is None: + base = SceneSnapshot(timestamp=self._scene_time, version=0) + else: + base = self.scene_provider.snapshot( + timestamp=self._scene_time, + env_ids=env_ids, + ) + if not bool(self.planner_policy.get("dynamic_collision", False)): + return base + exclusion_masks = self._collision_exclusion_masks(grounded, state) + entities = dict(base.entities) + for uid in dynamic_uids: + entity_state = entities.get(uid) + if entity_state is None: + raise ValueError( + f"SceneProvider omitted cuRobo dynamic obstacle {uid!r}." + ) + pose = entity_state.pose.to(dtype=torch.float32, device=self.device) + if pose.shape == (4, 4): + pose = pose.unsqueeze(0).repeat(self.num_envs, 1, 1) + if pose.shape != (self.num_envs, 4, 4): + raise ValueError( + f"Dynamic obstacle {uid!r} pose must have shape (4, 4) or " + f"({self.num_envs}, 4, 4), got {tuple(pose.shape)}." + ) + excluded = exclusion_masks.get(uid) + if excluded is not None and bool(excluded.any()): + pose = pose.clone() + pose[excluded, 2, 3] += _COLLISION_PARKING_Z_OFFSET + entities[uid] = EntityState( + pose=pose, + confidence=entity_state.confidence, + ) + return SceneSnapshot( + timestamp=base.timestamp, + version=base.version, + entities=entities, + collision_world_revision=base.collision_world_revision, + collision_entity_ids=dynamic_uids, + ) + + def _collision_exclusion_masks( + self, + grounded: GroundedAction, + state: ExecutionState, + ) -> dict[str, torch.Tensor]: + """Return per-environment masks for obstacles intentionally in contact.""" + dynamic_uids = { + str(uid) for uid in self.planner_policy.get("dynamic_obstacle_uids", ()) + } + masks: dict[str, torch.Tensor] = {} + + def include(uid: str | None, env_mask: torch.Tensor | None = None) -> None: + if uid is None or uid not in dynamic_uids: + return + mask = ( + torch.ones(self.num_envs, dtype=torch.bool, device=self.device) + if env_mask is None + else torch.as_tensor( + env_mask, + dtype=torch.bool, + device=self.device, + ).reshape(-1) + ) + if mask.shape != (self.num_envs,): + raise ValueError( + f"Collision exclusion mask for {uid!r} must have shape " + f"({self.num_envs},), got {tuple(mask.shape)}." + ) + masks[uid] = masks.get(uid, torch.zeros_like(mask)) | mask + + if self.capabilities.get(grounded.action_class).allows_target_contact: + target_uid = grounded.object_uid + if target_uid is None: + target_uid = getattr( + getattr(grounded.target, "semantics", None), + "label", + None, + ) + include(target_uid) + + for held in state.held_objects.values(): + include(held.semantics.label, held.env_mask) + collision_exclusion_uids = grounded.motion_policy.get( + "collision_exclusion_uids", () + ) + if isinstance(collision_exclusion_uids, str): + collision_exclusion_uids = (collision_exclusion_uids,) + for uid in collision_exclusion_uids: + include(str(uid)) + return masks + + def _invocation( + self, + grounded: GroundedAction, + capability: AtomicCapability, + ) -> ActionInvocation: + if capability.resource_mode == "coordinated_object": + strategy = str(self.planner_policy["coordinated_strategy"]) + elif grounded.control == "hand": + strategy = "ik_interp" + else: + strategy = str(self.planner_policy["single_arm_strategy"]) + sample_count = max(2, int(grounded.cfg.get("sample_interval", 50))) + dynamic_collision = bool(self.planner_policy.get("dynamic_collision", False)) + collision_required = ( + grounded.motion_policy.get("collision_safety") == "required" + ) + if dynamic_collision and strategy == "motion_gen": + dynamic_mode = ( + DynamicCollisionMode.REQUIRED + if collision_required + else DynamicCollisionMode.AUTO + ) + else: + dynamic_mode = DynamicCollisionMode.OFF + goal = ( + self._coordinated_pickment_goal(grounded) + if capability.config_materializer == "coordinated_pickment" + else grounded.target + ) + return ActionInvocation( + skill_id=str(capability.action_type.skill_id), + goal=goal, + binding=self._binding(grounded, capability), + motion_policy=MotionPolicy( + strategy=strategy, + sample_count=sample_count, + dynamic_collision_mode=dynamic_mode, + ), + recovery_policy=RecoveryPolicy(), + skill_options=self._build_config(grounded, capability), + ) + + @staticmethod + def _coordinated_pickment_goal(grounded: GroundedAction) -> CoordinatedPickGoal: + """Apply GenSim-only coordinated grasp filtering to an owned goal copy.""" + target = grounded.target + if not isinstance(target, CoordinatedPickGoal): + raise TypeError("CoordinatedPickment requires a CoordinatedPickGoal.") + requested = grounded.cfg.get("is_filter_ground_collision") + if requested is None: + return target + if not isinstance(requested, bool): + raise TypeError("is_filter_ground_collision must be a boolean.") + semantics = target.semantics + affordance = semantics.affordance + if not isinstance(affordance, AntipodalAffordance): + raise TypeError( + "CoordinatedPickment requires an AntipodalAffordance for GenSim " + "grasp filtering." + ) + generator_cfg = deepcopy(affordance.generator_cfg or GraspGeneratorCfg()) + generator_cfg.is_filter_ground_collision = requested + scoped_affordance = replace(affordance, generator_cfg=generator_cfg) + scoped_semantics = replace(semantics, affordance=scoped_affordance) + return replace(target, semantics=scoped_semantics) + + def _binding( + self, + action: GroundedAction, + capability: AtomicCapability, + ) -> ActionBinding: + engine = self._engine() + contract = getattr(capability.action_type, "binding_contract", None) + if contract is None: + return ActionBinding(owner_id=engine.binding_owner_id) + + slot_parts: dict[str, tuple[str, str | None]] = {} + if capability.config_materializer == "handover": + transfer_side = str(action.cfg.get("transfer_arm", "left_arm")) + receive_side = "right_arm" if transfer_side == "left_arm" else "left_arm" + transfer_arm, transfer_hand, _ = self._parts(transfer_side) + receive_arm, receive_hand, _ = self._parts(receive_side) + if transfer_hand is None or receive_hand is None: + raise ValueError("HandOver requires two configured end effectors.") + slot_parts = { + "source": (transfer_arm, transfer_hand), + "destination": (receive_arm, receive_hand), + } + elif capability.config_materializer == "coordinated_pickment": + left_arm, left_hand, _ = self._parts("left_arm") + right_arm, right_hand, _ = self._parts("right_arm") + if left_hand is None or right_hand is None: + raise ValueError("Coordinated pickup requires two end effectors.") + slot_parts = { + "left": (left_arm, left_hand), + "right": (right_arm, right_hand), + } + elif capability.config_materializer == "coordinated_placement": + placing_arm, placing_hand, _ = self._parts("left_arm") + support_arm, support_hand, _ = self._parts("right_arm") + if placing_hand is None or support_hand is None: + raise ValueError("Coordinated placement requires two end effectors.") + slot_parts = { + "placing": (placing_arm, placing_hand), + "support": (support_arm, support_hand), + } + else: + arm_part, hand_part, _ = self._parts(action.arm) + motion_part = hand_part if action.control == "hand" else arm_part + if motion_part is None: + raise ValueError( + f"{action.arm} has no configured {action.control} part." + ) + slot_parts = {"primary": (motion_part, hand_part)} + + endpoints: dict[str, dict[str, str]] = {} + for slot in contract.slots: + try: + motion_part, hand_part = slot_parts[slot.slot_id] + except KeyError as exc: + raise ValueError( + f"No GenSim binding is available for slot {slot.slot_id!r}." + ) from exc + selected: dict[str, str] = {} + for requirement in slot.endpoints: + if requirement.endpoint_id == "motion": + selected["motion"] = motion_part + elif requirement.endpoint_id == "grasp": + if hand_part is None: + raise ValueError( + f"{capability.name} requires a grasp endpoint for " + f"slot {slot.slot_id!r}." + ) + selected["grasp"] = hand_part + else: + raise ValueError( + f"Unsupported GenSim endpoint {slot.slot_id}." + f"{requirement.endpoint_id}." + ) + endpoints[slot.slot_id] = selected + return engine.bind_control_parts( + str(capability.action_type.skill_id), + endpoints, + ) + + def _build_config( + self, + action: GroundedAction, + capability: AtomicCapability | type, + ) -> Any: + """Build the mainline immutable ``ActionOptions`` value. + + The method name is retained as a narrow compatibility hook for existing + Action Engine tests and extensions; it no longer constructs legacy + hardware-bound ``ActionCfg`` objects. + """ + if isinstance(capability, type): + registered = self.capabilities.require_executable(action.action_class) + if registered.config_type is not capability: + raise ValueError( + f"Options type {capability.__name__!r} does not match " + f"AtomicAction {action.action_class!r}." + ) + capability = registered + if capability.config_materializer_hook is not None: + return capability.config_materializer_hook( + adapter=self, + action=action, + capability=capability, + ) + builder = getattr( + self, + f"_build_{capability.config_materializer}_config", + self._build_single_arm_config, + ) + return builder(action, capability) + + def _config_policy(self, action: GroundedAction) -> dict[str, Any]: + policy = dict(action.cfg) + for key in ( + "postcondition_tolerance", + "relation_distance", + "hover_height", + "staging_lift_height", + "transport_clearance", + "surface_clearance", + "receiver_hold_joint_tolerance", + "post_hold_steps", + ): + policy.pop(key, None) + return policy + + def _build_single_arm_config( + self, + action: GroundedAction, + capability: AtomicCapability, + ) -> Any: + policy = self._config_policy(action) + config_type = capability.config_type + if capability.target_materializer == "semantic_held_object": + from .atomic_compat import ExactTargetMoveHeldObjectOptions + + config_type = ExactTargetMoveHeldObjectOptions + if int(action.cfg.get("upright_yaw_samples", 1)) > 1: + policy["allow_automatic_transport_rotation"] = False + if capability.target_materializer == "press": + press_depth = policy.pop("press_depth", None) + if press_depth is not None and "press_distance" not in policy: + policy["press_distance"] = press_depth + approach_mode = policy.pop("approach_direction_mode", None) + if approach_mode == "handover_transfer": + from .frames import robot_frame_axes + + _, lateral = robot_frame_axes(self.env) + outward = lateral[0] if action.arm == "left_arm" else -lateral[0] + policy["approach_direction"] = _diagonal_approach_direction( + -outward.to(device=self.device) + ) + elif approach_mode is not None: + raise ValueError(f"Unknown approach_direction_mode {approach_mode!r}.") + for name in ("approach_direction", "obj_upright_direction"): + if name in policy and not isinstance(policy[name], torch.Tensor): + policy[name] = torch.as_tensor( + policy[name], dtype=torch.float32, device=self.device + ) + return config_type(**_supported_kwargs(config_type, policy)) + + def _build_coordinated_pickment_config( + self, + action: GroundedAction, + capability: AtomicCapability, + ) -> Any: + from .frames import arm_base_poses + + policy = self._config_policy(action) + left_base, right_base = arm_base_poses(self.env) + direction = right_base[0, :3, 3] - left_base[0, :3, 3] + norm = torch.linalg.vector_norm(direction) + if not torch.isfinite(direction).all() or norm <= 1.0e-6: + raise ValueError( + "Coordinated pickup requires distinct finite left/right arm bases." + ) + policy["left_to_right_arm_direction"] = direction / norm + return capability.config_type( + **_supported_kwargs(capability.config_type, policy) + ) + + def _build_coordinated_placement_config( + self, + action: GroundedAction, + capability: AtomicCapability, + ) -> Any: + return self._build_single_arm_config(action, capability) + + def _build_handover_config( + self, + action: GroundedAction, + capability: AtomicCapability, + ) -> Any: + policy = self._config_policy(action) + middle = action.cfg.get("middle_object_pose") + final = action.cfg.get("final_object_pose") + if middle is None or final is None: + raise ValueError("HandOver grounding must provide middle and final poses.") + transfer_side = str(action.cfg.get("transfer_arm", "left_arm")) + receive_side = "right_arm" if transfer_side == "left_arm" else "left_arm" + from .frames import robot_frame_axes + + _, lateral = robot_frame_axes(self.env) + receiver_outward = ( + lateral[0] if receive_side == "left_arm" else -lateral[0] + ).to(device=self.device) + receiver_inward_approach = -receiver_outward + policy.update( + { + "middle_object_pose": middle, + # Delivery is represented by a following MoveHeldObject node. + # Keep the receiver fixed while the source retreats here. + "final_object_pose": middle, + "receive_approach_direction": _diagonal_approach_direction( + receiver_inward_approach + ), + } + ) + return capability.config_type( + **_supported_kwargs(capability.config_type, policy) + ) + + def _positions_with_agent_holds( + self, + plan: ActionPlan, + grounded: GroundedAction, + capability: AtomicCapability, + ) -> torch.Tensor: + trajectory = plan.joint_trajectory + if trajectory is None: + raise ValueError( + f"AtomicAction {plan.skill_id!r} did not retain a joint trajectory." + ) + positions = trajectory.positions.to( + device=self.device, + dtype=torch.float32, + ) + hold_steps = int(grounded.cfg.get("post_hold_steps", 0)) + if capability.state_effect != "release" or hold_steps <= 0: + return positions + release = next((item for item in plan.segments if item.name == "release"), None) + if release is None or release.stop <= 0 or release.stop > positions.shape[1]: + return positions + hold = positions[:, release.stop - 1 : release.stop].repeat(1, hold_steps, 1) + return torch.cat( + (positions[:, : release.stop], hold, positions[:, release.stop :]), + dim=1, + ) + + @staticmethod + def _merge_plan_rows( + primary: torch.Tensor, + fallback: torch.Tensor, + use_fallback: torch.Tensor, + hold_qpos: torch.Tensor, + ) -> torch.Tensor: + steps = max(primary.shape[1], fallback.shape[1], 1) + + def padded(value: torch.Tensor) -> torch.Tensor: + if value.shape[1] == 0: + return hold_qpos[:, None].repeat(1, steps, 1) + if value.shape[1] < steps: + value = torch.cat( + (value, value[:, -1:].repeat(1, steps - value.shape[1], 1)), + dim=1, + ) + return value + + primary = padded(primary) + fallback = padded(fallback) + return torch.where(use_fallback[:, None, None], fallback, primary) + + def _handover_receiver_hold_mask( + self, + trajectory: torch.Tensor, + grounded: GroundedAction, + options: Any, + *, + tolerance: float, + ) -> torch.Tensor: + if tolerance < 0.0: + raise ValueError("receiver_hold_joint_tolerance must be non-negative.") + retreat_steps = max(2, int(options.retreat_steps)) + if trajectory.shape[1] < retreat_steps: + return torch.zeros( + self.num_envs, dtype=torch.bool, device=trajectory.device + ) + transfer_side = str(grounded.cfg.get("transfer_arm", "left_arm")) + receive_side = "right_arm" if transfer_side == "left_arm" else "left_arm" + receive_arm, _, _ = self._parts(receive_side) + receiver_ids = self.env.robot.get_joint_ids(name=receive_arm) + receiver = trajectory[:, -retreat_steps:, receiver_ids] + drift = torch.amax(torch.abs(receiver - receiver[:, :1]), dim=(1, 2)) + return torch.isfinite(drift) & (drift <= tolerance) + + def execute_trajectory( + self, + trajectory: torch.Tensor, + *, + active: torch.Tensor, + ) -> list[torch.Tensor]: + """Advance the environment while holding inactive vectorized rows.""" + if trajectory.ndim != 3 or trajectory.shape[0] != self.num_envs: + raise ValueError("Execution trajectory must have shape (N, T, robot_dof).") + active = active.to(device=trajectory.device, dtype=torch.bool) + current = self.env.robot.get_qpos().to( + device=trajectory.device, + dtype=trajectory.dtype, + ) + commands: list[torch.Tensor] = [] + for waypoint in trajectory.unbind(dim=1): + command = torch.where(active[:, None], waypoint, current) + self.env.step(command) + self._scene_time += self._scene_step_duration() + update = getattr(self.env, "update_obj_info", None) + if callable(update): + update() + commands.append(command.detach()) + current = command + sync = getattr(self.env, "sync_agent_state_from_qpos", None) + if callable(sync) and commands: + sync(commands[-1]) + return commands + + def _scene_step_duration(self) -> float: + """Return one positive logical waypoint duration for scene timestamps.""" + sim_config = getattr(getattr(self.env, "sim", None), "sim_config", None) + candidates = ( + getattr(self.env, "physics_dt", None), + getattr(sim_config, "physics_dt", None), + ) + for value in candidates: + if isinstance(value, (int, float)) and not isinstance(value, bool): + duration = float(value) + if math.isfinite(duration) and duration > 0.0: + return duration + return 1.0 + + def combine( + self, + outcomes: Mapping[str, ActionOutcome | None], + masks: Mapping[str, torch.Tensor], + ) -> tuple[torch.Tensor, torch.Tensor]: + """Merge independently planned arm paths into one synchronized stream.""" + present = [item for item in outcomes.values() if item is not None] + if not present: + raise ValueError("At least one arm outcome is required.") + steps = max(int(item.trajectory.shape[1]) for item in present) + current = self.env.robot.get_qpos().to(self.device, dtype=torch.float32) + merged = current[:, None, :].repeat(1, max(steps, 1), 1) + success = torch.ones( + self.num_envs, + dtype=torch.bool, + device=self.device, + ) + for arm, outcome in outcomes.items(): + if outcome is None: + continue + mask = masks[arm].to(self.device, dtype=torch.bool) + success &= ~mask | outcome.success + trajectory = outcome.trajectory + if trajectory.shape[1] == 0: + continue + if trajectory.shape[1] < steps: + padding = trajectory[:, -1:].repeat(1, steps - trajectory.shape[1], 1) + trajectory = torch.cat((trajectory, padding), dim=1) + joint_ids = self.joint_ids(arm, include_hand=True) + if not joint_ids: + continue + selected = merged[:, :, joint_ids] + merged[:, :, joint_ids] = torch.where( + mask[:, None, None], trajectory[:, :, joint_ids], selected + ) + return merged, success + + def joint_ids(self, arm: str, *, include_hand: bool) -> list[int]: + if arm == "coordinated": + return list(range(int(self.env.robot.dof))) + side = "left" if arm == "left_arm" else "right" + result = list(getattr(self.env, f"{side}_arm_joints", ())) + if include_hand: + result.extend(getattr(self.env, f"{side}_eef_joints", ())) + return result + + def _engine(self) -> AtomicActionEngine: + if self._atomic_engine is None: + from .atomic_compat import ExactTargetMoveHeldObject + + engine = AtomicActionEngine( + self._generator(), + control_profiles=self._control_profiles(), + ) + engine.register(ExactTargetMoveHeldObject(), replace=True) + self._atomic_engine = engine + return self._atomic_engine + + def _generator(self) -> MotionGenerator: + if self._motion_generator is None: + backend = str(self.planner_policy.get("backend", "curobo")) + if backend == "curobo": + options = dict(self.planner_policy.get("curobo", {})) + obstacle_uids = tuple( + dict.fromkeys( + [ + *self.planner_policy.get("static_obstacle_uids", ()), + *self.planner_policy.get("dynamic_obstacle_uids", ()), + ] + ) + ) + rigid_objects: dict[str, Any] = {} + for uid in obstacle_uids: + obstacle_uid = str(uid) + entity = self.env.sim.get_rigid_object(obstacle_uid) + if entity is None: + raise ValueError(f"Unknown cuRobo obstacle {uid!r}.") + rigid_objects[obstacle_uid] = entity + obstacle_representation = str( + options.get("obstacle_representation", "cuboid") + ) + world = CuroboWorldCfg( + rigid_objects=rigid_objects or None, + obstacle_representation=obstacle_representation, + collision_cache=_collision_cache_for_world( + obstacle_representation, + len(rigid_objects), + ), + dynamic_obstacle_names=[ + str(uid) + for uid in self.planner_policy.get("dynamic_obstacle_uids", ()) + ], + multi_env=bool(options.get("multi_env", False)), + ) + planner_cfg = CuroboPlannerCfg( + robot_uid=self.env.robot.uid, + log_level=str(options.get("log_level", "error")), + world=world, + use_cuda_graph=bool(options.get("use_cuda_graph", True)), + preserve_plan_samples=bool( + options.get("preserve_plan_samples", False) + ), + max_attempts=int(options.get("max_attempts", 5)), + collision_activation_distance=float( + options.get("collision_activation_distance", 0.01) + ), + ) + elif backend == "toppra": + planner_cfg = ToppraPlannerCfg(robot_uid=self.env.robot.uid) + else: + raise ValueError( + f"Unsupported Action Engine planner backend {backend!r}." + ) + self._motion_generator = MotionGenerator( + cfg=MotionGenCfg(planner_cfg=planner_cfg) + ) + return self._motion_generator + + def _control_profiles(self) -> dict[str, ControlPartCommandProfile]: + profiles: dict[str, ControlPartCommandProfile] = {} + for side in ("left_arm", "right_arm"): + try: + _, hand_part, hand_dof = self._parts(side) + except ValueError: + continue + if hand_part is None or hand_dof == 0 or hand_part in profiles: + continue + profiles[hand_part] = ControlPartCommandProfile.joint_positions( + open=_as_hand_qpos(self.env.open_state, hand_dof, self.device), + grasp=_as_hand_qpos(self.env.close_state, hand_dof, self.device), + ) + return profiles + + def _parts(self, arm: str) -> tuple[str, str | None, int]: + if arm not in {"left_arm", "right_arm"}: + raise ValueError(f"Expected a physical arm, got {arm!r}.") + is_left = arm == "left_arm" + if hasattr(self.env, "get_agent_arm_control_part"): + arm_part = self.env.get_agent_arm_control_part(is_left) + hand_part = self.env.get_agent_eef_control_part(is_left) + else: + arm_part = arm + hand_part = "left_eef" if is_left else "right_eef" + hand_ids = ( + [] + if hand_part is None + else list(self.env.robot.get_joint_ids(name=hand_part)) + ) + return ( + str(arm_part), + None if hand_part is None else str(hand_part), + len(hand_ids), + ) diff --git a/embodichain/gen_sim/action_engine/runtime/atomic_compat.py b/embodichain/gen_sim/action_engine/runtime/atomic_compat.py new file mode 100644 index 000000000..d076e2b12 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/atomic_compat.py @@ -0,0 +1,84 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Action Engine-specific adapters for mainline atomic actions.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from embodichain.lab.sim.atomic_actions import ( + ActionPlan, + HeldObjectPoseGoal, + MoveHeldObject, + MoveHeldObjectOptions, + PlanningContext, + ResolvedActionRequest, +) + +__all__ = ["ExactTargetMoveHeldObject", "ExactTargetMoveHeldObjectOptions"] + + +@dataclass(frozen=True, slots=True, eq=False) +class ExactTargetMoveHeldObjectOptions(MoveHeldObjectOptions): + """Action Engine transport options with an exact-orientation switch.""" + + allow_automatic_transport_rotation: bool = True + """Whether the mainline transport heuristic may replace target rotation.""" + + +class ExactTargetMoveHeldObject(MoveHeldObject): + """Preserve a selected semantic orientation when explicitly requested.""" + + OptionsType = ExactTargetMoveHeldObjectOptions + + def __init__( + self, + default_options: ExactTargetMoveHeldObjectOptions | None = None, + ) -> None: + super().__init__(default_options) + self._allow_automatic_transport_rotation = True + + def _plan( + self, + request: ResolvedActionRequest[ + HeldObjectPoseGoal, + ExactTargetMoveHeldObjectOptions, + ], + context: PlanningContext, + ) -> ActionPlan: + previous = self._allow_automatic_transport_rotation + self._allow_automatic_transport_rotation = ( + request.skill_options.allow_automatic_transport_rotation + ) + try: + return super()._plan(request, context) + finally: + self._allow_automatic_transport_rotation = previous + + def _apply_automatic_transport_rotation( + self, + move_eef_xpos: torch.Tensor, + end_arm_xpos: torch.Tensor, + ) -> None: + """Apply the heuristic unless semantic grounding selected exact yaw.""" + if self._allow_automatic_transport_rotation: + super()._apply_automatic_transport_rotation( + move_eef_xpos, + end_arm_xpos, + ) diff --git a/embodichain/gen_sim/action_engine/runtime/frames.py b/embodichain/gen_sim/action_engine/runtime/frames.py new file mode 100644 index 000000000..513308484 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/frames.py @@ -0,0 +1,165 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Resolve directional relations in live robot and world frames.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from .robot_parts import arm_control_part + +__all__ = [ + "DIRECTIONAL_RELATIONS", + "arm_base_poses", + "relation_axes", + "relation_offset", + "robot_frame_axes", +] + + +_RELATION_COMPONENTS = { + "left": ("left",), + "left_of": ("left",), + "right": ("right",), + "right_of": ("right",), + "front": ("front",), + "front_of": ("front",), + "in_front_of": ("front",), + "behind": ("back",), + "back": ("back",), + "front_left": ("front", "left"), + "front_left_of": ("front", "left"), + "front_right": ("front", "right"), + "front_right_of": ("front", "right"), + "back_left": ("back", "left"), + "back_left_of": ("back", "left"), + "back_right": ("back", "right"), + "back_right_of": ("back", "right"), +} +DIRECTIONAL_RELATIONS = frozenset(_RELATION_COMPONENTS) + + +def arm_base_poses(env: Any) -> tuple[torch.Tensor, torch.Tensor]: + """Return live world poses of the left and right arm bases.""" + left_part = arm_control_part(env, "left_arm") + right_part = arm_control_part(env, "right_arm") + robot = env.robot + if hasattr(robot, "get_solver") and hasattr(robot, "get_link_pose"): + left_solver = robot.get_solver(name=left_part) + right_solver = robot.get_solver(name=right_part) + left_root = getattr(left_solver, "root_link_name", None) + right_root = getattr(right_solver, "root_link_name", None) + if left_root is None or right_root is None: + raise ValueError("Directional grounding requires both arm root links.") + left = robot.get_link_pose(link_name=left_root, to_matrix=True) + right = robot.get_link_pose(link_name=right_root, to_matrix=True) + elif hasattr(robot, "get_control_part_base_pose"): + left = robot.get_control_part_base_pose(name=left_part, to_matrix=True) + right = robot.get_control_part_base_pose(name=right_part, to_matrix=True) + elif hasattr(env, "get_current_xpos_agent"): + left, right = env.get_current_xpos_agent() + else: + raise ValueError( + "Directional grounding requires live left/right arm-base or TCP poses." + ) + + left = _batched_pose(left, env) + right = _batched_pose(right, env) + return left, right + + +def robot_frame_axes(env: Any) -> tuple[torch.Tensor, torch.Tensor]: + """Return normalized world-space forward and left axes for a dual-arm robot.""" + left, right = arm_base_poses(env) + lateral = left[:, :2, 3] - right[:, :2, 3] + norm = torch.linalg.vector_norm(lateral, dim=1, keepdim=True) + if bool((norm <= 1.0e-6).any()): + raise ValueError("Left and right arm bases must have distinct XY positions.") + lateral = lateral / norm + forward = torch.stack((lateral[:, 1], -lateral[:, 0]), dim=1) + return forward, lateral + + +def relation_axes( + env: Any, + relation: str, + *, + frame: str, +) -> tuple[torch.Tensor, ...]: + """Return signed world-space axes whose projections define a relation.""" + relation = str(relation) + if relation not in DIRECTIONAL_RELATIONS: + return () + if frame == "robot": + forward, lateral = robot_frame_axes(env) + elif frame == "world": + count = int(env.num_envs) + forward = torch.tensor( + [1.0, 0.0], dtype=torch.float32, device=env.device + ).repeat(count, 1) + lateral = torch.tensor( + [0.0, 1.0], dtype=torch.float32, device=env.device + ).repeat(count, 1) + else: + raise ValueError(f"Unsupported directional relation frame {frame!r}.") + + component_axes = { + "front": forward, + "back": -forward, + "left": lateral, + "right": -lateral, + } + return tuple(component_axes[item] for item in _RELATION_COMPONENTS[relation]) + + +def relation_offset( + env: Any, + relation: str, + *, + frame: str, + forward_distance: float, + lateral_distance: float, + dtype: torch.dtype, + device: torch.device, +) -> torch.Tensor | None: + """Resolve one directional relation into a batched world-space offset.""" + axes = relation_axes(env, relation, frame=frame) + if not axes: + return None + offset = torch.zeros((int(env.num_envs), 3), dtype=dtype, device=device) + components = _RELATION_COMPONENTS[relation] + for component, axis in zip(components, axes): + axis = axis.to(dtype=dtype, device=device) + distance = ( + forward_distance if component in {"front", "back"} else lateral_distance + ) + offset[:, :2] += axis * float(distance) + return offset + + +def _batched_pose(value: Any, env: Any) -> torch.Tensor: + pose = torch.as_tensor(value, dtype=torch.float32, device=env.device) + if pose.shape == (4, 4): + pose = pose.unsqueeze(0).repeat(int(env.num_envs), 1, 1) + if pose.shape != (int(env.num_envs), 4, 4): + raise ValueError( + "Frame pose must have shape (4, 4) or " + f"({int(env.num_envs)}, 4, 4), got {tuple(pose.shape)}." + ) + return pose diff --git a/embodichain/gen_sim/action_engine/runtime/grasp_collision_cache.py b/embodichain/gen_sim/action_engine/runtime/grasp_collision_cache.py new file mode 100644 index 000000000..252c388d6 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/grasp_collision_cache.py @@ -0,0 +1,330 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Prepare checksummed V-HACD caches for the shared grasp collision checker. + +The sidecar identifies the backend without changing Main's cache key or pickle +payload, so an unlabelled CoACD cache is never silently reused as V-HACD. +""" + +from __future__ import annotations + +from dataclasses import dataclass +import hashlib +import io +import json +import operator +import os +from pathlib import Path +import pickle +import stat +import tempfile +from typing import Literal + +import numpy as np +import torch + +__all__ = [ + "GraspCollisionCacheError", + "GraspCollisionCacheResult", + "ensure_vhacd_grasp_collision_cache", + "grasp_collision_cache_path", +] + +_CACHE_SCHEMA_VERSION = 1 +_METADATA_SUFFIX = ".action_engine.json" +_DEFAULT_CACHE_DIR = ( + Path.home() / ".cache" / "embodichain_cache" / "convex_decomposition" +) + +CacheStatus = Literal["hit", "generated", "replaced"] + + +class GraspCollisionCacheError(RuntimeError): + """Raised when a safe, Main-compatible V-HACD cache cannot be prepared.""" + + +@dataclass(frozen=True) +class GraspCollisionCacheResult: + """Describe the prepared cache files and whether decomposition ran.""" + + status: CacheStatus + cache_path: Path + metadata_path: Path + + +def grasp_collision_cache_path( + mesh_vertices: torch.Tensor | np.ndarray, + mesh_triangles: torch.Tensor | np.ndarray, + max_decomposition_hulls: int, + *, + cache_dir: str | Path | None = None, +) -> Path: + """Return Main's exact ``_.pkl`` cache path.""" + vertices, triangles = _validate_mesh(mesh_vertices, mesh_triangles) + hull_limit = _validate_hull_limit(max_decomposition_hulls) + mesh_hash = hashlib.md5(vertices.tobytes() + triangles.tobytes()).hexdigest() + return _resolve_cache_dir(cache_dir) / f"{mesh_hash}_{hull_limit}.pkl" + + +def ensure_vhacd_grasp_collision_cache( + *, + mesh_vertices: torch.Tensor | np.ndarray, + mesh_triangles: torch.Tensor | np.ndarray, + max_decomposition_hulls: int, + cache_dir: str | Path | None = None, +) -> GraspCollisionCacheResult: + """Create or validate a V-HACD cache and its checksummed backend sidecar.""" + vertices, triangles = _validate_mesh(mesh_vertices, mesh_triangles) + hull_limit = _validate_hull_limit(max_decomposition_hulls) + mesh_hash = hashlib.md5(vertices.tobytes() + triangles.tobytes()).hexdigest() + cache_path = _resolve_cache_dir(cache_dir) / f"{mesh_hash}_{hull_limit}.pkl" + metadata_path = cache_path.with_name(f"{cache_path.name}{_METADATA_SUFFIX}") + expected_metadata: dict[str, object] = { + "schema_version": _CACHE_SCHEMA_VERSION, + "backend": "vhacd", + "mesh_hash": mesh_hash, + "max_decomposition_hulls": hull_limit, + } + + _prepare_private_directory(cache_path.parent) + _refuse_symlink(cache_path) + _refuse_symlink(metadata_path) + if _cache_matches_metadata(cache_path, metadata_path, expected_metadata): + return GraspCollisionCacheResult("hit", cache_path, metadata_path) + + exists = cache_path.exists() or metadata_path.exists() + status: CacheStatus = "replaced" if exists else "generated" + try: + plane_equations = _compute_vhacd_plane_equations( + vertices, + triangles, + hull_limit, + ) + cache_bytes = _serialize_checker_payload(plane_equations) + metadata = { + **expected_metadata, + "cache_sha256": hashlib.sha256(cache_bytes).hexdigest(), + } + + # Publish the complete pickle before its sidecar. A crash between the + # two replaces leaves a cache miss on retry, never a partial pickle. + _write_bytes_atomic(cache_path, cache_bytes) + metadata_bytes = (json.dumps(metadata, sort_keys=True) + "\n").encode() + _write_bytes_atomic(metadata_path, metadata_bytes) + except GraspCollisionCacheError: + raise + except Exception as exc: + raise GraspCollisionCacheError( + f"Failed to prepare V-HACD grasp collision cache {cache_path}: {exc}" + ) from exc + + return GraspCollisionCacheResult(status, cache_path, metadata_path) + + +def _compute_vhacd_plane_equations( + vertices: np.ndarray, + triangles: np.ndarray, + max_decomposition_hulls: int, +) -> list[tuple[np.ndarray, np.ndarray]]: + """Run DexSim V-HACD and convert its hulls to checker plane equations.""" + import open3d as o3d + from dexsim.kit.meshproc import convex_decomposition_vhacd + + from embodichain.toolkits.graspkit.pg_grasp.collision_checker import ( + extract_plane_equations, + ) + + mesh = o3d.t.geometry.TriangleMesh() + mesh.vertex.positions = o3d.core.Tensor(vertices.astype(np.float32, copy=False)) + mesh.triangle.indices = o3d.core.Tensor(triangles.astype(np.int32, copy=False)) + is_success, hull_meshes = convex_decomposition_vhacd( + mesh, + max_convex_hull_num=max_decomposition_hulls, + ) + if not is_success or not hull_meshes: + raise GraspCollisionCacheError( + "V-HACD returned no convex hulls for the grasp collision mesh." + ) + + convex_parts = [ + ( + np.asarray(hull.vertex.positions.numpy()), + np.asarray(hull.triangle.indices.numpy()), + ) + for hull in hull_meshes + ] + plane_equations = extract_plane_equations(convex_parts) + if not plane_equations: + raise GraspCollisionCacheError( + "V-HACD hulls produced no grasp collision plane equations." + ) + return plane_equations + + +def _serialize_checker_payload( + plane_equations: list[tuple[np.ndarray, np.ndarray]], +) -> bytes: + """Pack plane equations in the exact tensor dictionary Main unpickles.""" + if not plane_equations: + raise ValueError("V-HACD must produce at least one convex hull.") + + normalized: list[tuple[np.ndarray, np.ndarray]] = [] + for normals_value, offsets_value in plane_equations: + normals = np.asarray(normals_value, dtype=np.float32) + offsets = np.asarray(offsets_value, dtype=np.float32) + if normals.ndim != 2 or normals.shape[1:] != (3,) or not len(normals): + raise ValueError("Each V-HACD hull must have normals shaped [K, 3].") + if offsets.shape != (len(normals),): + raise ValueError("Each hull needs one offset per plane normal.") + if not np.isfinite(normals).all() or not np.isfinite(offsets).all(): + raise ValueError("V-HACD plane equations must contain finite values.") + normalized.append((normals, offsets)) + + max_plane_count = max(normals.shape[0] for normals, _ in normalized) + equations = torch.zeros((len(normalized), max_plane_count, 4)) + counts = torch.zeros(len(normalized), dtype=torch.int32) + for index, (normals, offsets) in enumerate(normalized): + plane_count = normals.shape[0] + equations[index, :plane_count, :3] = torch.from_numpy(normals) + equations[index, :plane_count, 3] = torch.from_numpy(offsets) + counts[index] = plane_count + + stream = io.BytesIO() + payload = {"plane_equations": equations, "plane_equation_counts": counts} + pickle.dump(payload, stream, protocol=pickle.HIGHEST_PROTOCOL) + return stream.getvalue() + + +def _cache_matches_metadata( + cache_path: Path, + metadata_path: Path, + expected_metadata: dict[str, object], +) -> bool: + if not cache_path.is_file() or not metadata_path.is_file(): + return False + try: + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + checksum = metadata.get("cache_sha256") + expected_checksum = hashlib.sha256(cache_path.read_bytes()).hexdigest() + return ( + isinstance(metadata, dict) + and all( + metadata.get(key) == value for key, value in expected_metadata.items() + ) + and isinstance(checksum, str) + and checksum == expected_checksum + ) + except (AttributeError, OSError, UnicodeDecodeError, json.JSONDecodeError): + return False + + +def _validate_mesh( + mesh_vertices: torch.Tensor | np.ndarray, + mesh_triangles: torch.Tensor | np.ndarray, +) -> tuple[np.ndarray, np.ndarray]: + if isinstance(mesh_vertices, torch.Tensor): + mesh_vertices = mesh_vertices.detach().cpu().numpy() + if isinstance(mesh_triangles, torch.Tensor): + mesh_triangles = mesh_triangles.detach().cpu().numpy() + if not isinstance(mesh_vertices, np.ndarray): + raise TypeError("mesh_vertices must be a torch.Tensor or numpy.ndarray.") + if not isinstance(mesh_triangles, np.ndarray): + raise TypeError("mesh_triangles must be a torch.Tensor or numpy.ndarray.") + vertices = np.ascontiguousarray(mesh_vertices) + triangles = np.ascontiguousarray(mesh_triangles) + if vertices.ndim != 2 or vertices.shape[1:] != (3,) or len(vertices) == 0: + raise ValueError("mesh_vertices must have non-empty shape [N, 3].") + if triangles.ndim != 2 or triangles.shape[1:] != (3,) or len(triangles) == 0: + raise ValueError("mesh_triangles must have non-empty shape [M, 3].") + if not np.issubdtype(vertices.dtype, np.number): + raise TypeError("mesh_vertices must contain numeric values.") + if not np.isfinite(vertices).all(): + raise ValueError("mesh_vertices must contain only finite values.") + if not np.issubdtype(triangles.dtype, np.integer): + raise TypeError("mesh_triangles must contain integer indices.") + if triangles.min() < 0 or triangles.max() >= len(vertices): + raise ValueError("mesh_triangles contains out-of-range vertex indices.") + return vertices, triangles + + +def _validate_hull_limit(value: int) -> int: + if isinstance(value, (bool, np.bool_)): + raise TypeError("max_decomposition_hulls must be an integer.") + try: + hull_limit = operator.index(value) + except TypeError as exc: + raise TypeError("max_decomposition_hulls must be an integer.") from exc + if hull_limit <= 0: + raise ValueError("max_decomposition_hulls must be positive.") + return hull_limit + + +def _resolve_cache_dir(cache_dir: str | Path | None) -> Path: + if cache_dir is not None: + return Path(cache_dir).expanduser().resolve() + try: + from embodichain.lab.sim import CONVEX_DECOMP_DIR + except Exception: + return _DEFAULT_CACHE_DIR + return Path(CONVEX_DECOMP_DIR).expanduser().resolve() + + +def _prepare_private_directory(path: Path) -> None: + try: + path.mkdir(parents=True, exist_ok=True, mode=0o700) + path.chmod(0o700) + except OSError as exc: + raise GraspCollisionCacheError( + f"Cannot secure grasp collision cache directory: {path}" + ) from exc + if path.stat().st_mode & (stat.S_IWGRP | stat.S_IWOTH): + raise GraspCollisionCacheError(f"Refusing writable cache directory: {path}") + + +def _refuse_symlink(path: Path) -> None: + if path.is_symlink(): + raise GraspCollisionCacheError( + f"Refusing symlinked grasp collision cache path: {path}" + ) + + +def _write_bytes_atomic(path: Path, payload: bytes) -> None: + """Publish one complete file with a same-directory atomic replacement.""" + _refuse_symlink(path) + file_descriptor, temporary_name = tempfile.mkstemp( + prefix=f".{path.name}.", + suffix=".tmp", + dir=path.parent, + ) + temporary_path = Path(temporary_name) + try: + os.fchmod(file_descriptor, 0o600) + with os.fdopen(file_descriptor, "wb") as output: + file_descriptor = -1 + output.write(payload) + output.flush() + os.fsync(output.fileno()) + _refuse_symlink(path) + os.replace(temporary_path, path) + path.chmod(0o600) + finally: + if file_descriptor >= 0: + os.close(file_descriptor) + try: + temporary_path.unlink(missing_ok=True) + except OSError: + pass diff --git a/embodichain/gen_sim/action_engine/runtime/grounding.py b/embodichain/gen_sim/action_engine/runtime/grounding.py new file mode 100644 index 000000000..e7a63337c --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/grounding.py @@ -0,0 +1,2356 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Resolve symbolic bindings from live simulator state.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, replace +from typing import Any + +import torch + +from embodichain.gen_sim.action_engine.capabilities import ( + build_atomic_capability_registry, +) +from embodichain.gen_sim.action_engine.config import ( + RuntimePolicyCfg, + default_runtime_policy, +) +from embodichain.gen_sim.action_engine.domain import normalize_placement_relation +from embodichain.gen_sim.action_engine.orientation import ( + AlignAxisConstraint, + MatchRotationConstraint, + OrientationConstraint, + compile_orientation_constraint, +) +from embodichain.lab.sim.atomic_actions import ( + CoordinatedPickGoal, + CoordinatedPlacementGoal, + EndEffectorPoseGoal, + GraspGoal, + HeldObjectPoseGoal, + JointPositionGoal, + ObjectSemantics, + PlaceGoal, + PressAffordance, + PressGoal, +) +from .frames import arm_base_poses, relation_offset, robot_frame_axes +from .models import ExecutionProgram, GroundedAction, SemanticStep +from .motion_policy import resolve_motion_policy, with_motion_modifiers +from .robot_parts import arm_control_part +from .state import ExecutionState + +__all__ = ["ActionGrounder", "LiveArrangementPlan", "LivePlacementPlan"] + + +def _batched_pose(value: Any, env: Any) -> torch.Tensor: + pose = torch.as_tensor(value, dtype=torch.float32, device=env.device) + if pose.shape == (4, 4): + pose = pose.unsqueeze(0).repeat(int(env.num_envs), 1, 1) + if pose.shape != (int(env.num_envs), 4, 4): + raise ValueError( + "Live pose must have shape (4, 4) or " + f"({int(env.num_envs)}, 4, 4), got {tuple(pose.shape)}." + ) + return pose + + +def _object(env: Any, uid: str) -> Any: + entity = env.sim.get_rigid_object(uid) + if entity is None: + raise ValueError(f"Unknown rigid object {uid!r}.") + return entity + + +def _live_pose(env: Any, uid: str) -> torch.Tensor: + return _batched_pose(_object(env, uid).get_local_pose(to_matrix=True), env) + + +def _local_vertices(entity: Any, env: Any, env_id: int = 0) -> torch.Tensor: + value = entity.get_vertices(env_ids=[env_id], scale=True) + if isinstance(value, (list, tuple)): + value = value[0] + vertices = torch.as_tensor(value, dtype=torch.float32, device=env.device) + if vertices.ndim == 3 and vertices.shape[0] == 1: + vertices = vertices[0] + if vertices.ndim != 2 or vertices.shape[-1] != 3 or vertices.numel() == 0: + raise ValueError("Rigid-object mesh vertices must have shape (N, 3).") + return vertices + + +def _world_vertices(entity: Any, env: Any, env_id: int) -> torch.Tensor: + vertices = _local_vertices(entity, env, env_id) + pose = _batched_pose(entity.get_local_pose(to_matrix=True), env)[env_id] + return vertices @ pose[:3, :3].transpose(0, 1) + pose[:3, 3] + + +@dataclass(frozen=True) +class _Geometry: + radius: torch.Tensor + half_height: torch.Tensor + + +class LiveArrangementPlan: + """Materialize collision-aware line slots independently in every env.""" + + def __init__( + self, + env: Any, + steps: Sequence[SemanticStep], + *, + slot_margin: float | None = None, + minimum_spacing: float | None = None, + clearance: float | None = None, + row_search_step: float | None = None, + row_search_radius: float | None = None, + ) -> None: + if not steps: + raise ValueError("An arrangement plan requires at least one step.") + self.env = env + self.steps = tuple(steps) + self.step_by_id = {step.id: step for step in steps} + self.num_envs = int(env.num_envs) + self.device = env.device + self.slot_count = len(steps) + self.axis = str(steps[0].goal.get("axis", "world_x")) + profile = str(getattr(env, "agent_robot_profile", "dual_ur10")) + defaults = default_runtime_policy(profile).grounding["arrangement"] + slot_margin = defaults["slot_margin"] if slot_margin is None else slot_margin + minimum_spacing = ( + defaults["minimum_spacing"] if minimum_spacing is None else minimum_spacing + ) + self.clearance = float( + defaults["layout_clearance"] if clearance is None else clearance + ) + self.row_search_step = float( + defaults["row_search_step"] if row_search_step is None else row_search_step + ) + self.row_search_radius = float( + defaults["row_search_radius"] + if row_search_radius is None + else row_search_radius + ) + + table = _object(env, "table") + bounds = [] + for env_id in range(self.num_envs): + vertices = _world_vertices(table, env, env_id) + bounds.append( + torch.stack((vertices.min(dim=0).values, vertices.max(dim=0).values)) + ) + self.table_bounds = torch.stack(bounds) + self.table_center = self.table_bounds.mean(dim=1) + self.table_top = self.table_bounds[:, 1, 2] + if self.axis == "table_long_axis": + mean_extent = ( + self.table_bounds[:, 1, :2] - self.table_bounds[:, 0, :2] + ).mean(dim=0) + self.axis_index = int(torch.argmax(mean_extent).item()) + else: + self.axis_index = 0 if self.axis in {"x", "world_x"} else 1 + self.perpendicular_index = 1 - self.axis_index + self.geometry = {step.id: self._geometry(step) for step in self.steps} + diameters = torch.stack( + [self.geometry[step.id].radius * 2.0 for step in self.steps], + dim=1, + ) + self.spacing = torch.maximum( + diameters.max(dim=1).values + float(slot_margin), + torch.full( + (self.num_envs,), + float(minimum_spacing), + dtype=torch.float32, + device=self.device, + ), + ) + self.positions = self._make_slots() + self.reassignment_reason: list[str | None] = [None] * self.num_envs + self.reassignment_cost = torch.full( + (self.num_envs,), + float("nan"), + dtype=torch.float32, + device=self.device, + ) + self.assignments = self._initial_slot_assignments() + order_by = str(self.steps[0].goal.get("order_by", "explicit")) + direction = str(self.steps[0].goal.get("order_direction", "given")) + if order_by == "size" and not any( + step.goal.get("slot_constraint") == "free_reassignable" + for step in self.steps + ): + for env_id in range(self.num_envs): + ordered = sorted( + self.steps, + key=lambda step: float(self.geometry[step.id].radius[env_id]), + reverse=direction != "ascending", + ) + for slot_id, step in enumerate(ordered): + self.assignments[step.id][env_id] = slot_id + self.completed = { + step.id: torch.zeros( + self.num_envs, + dtype=torch.bool, + device=self.device, + ) + for step in self.steps + } + + def _initial_slot_assignments(self) -> dict[str, torch.Tensor]: + """Match free-order objects to slots in their current spatial order.""" + assignments = { + step.id: torch.full( + (self.num_envs,), + int(step.goal.get("nominal_slot_index", index)), + dtype=torch.long, + device=self.device, + ) + for index, step in enumerate(self.steps) + } + free_steps = [ + step + for step in self.steps + if step.goal.get("slot_constraint") == "free_reassignable" + ] + if not free_steps: + return assignments + required_slots = { + int(step.goal.get("nominal_slot_index", index)) + for index, step in enumerate(self.steps) + if step.goal.get("slot_constraint") != "free_reassignable" + } + available_slots = [ + slot_id + for slot_id in range(self.slot_count) + if slot_id not in required_slots + ] + if len(available_slots) != len(free_steps): + raise ValueError( + "Arrangement slot constraints do not define a one-to-one assignment." + ) + axis_positions = { + step.id: _live_pose(self.env, step.object_uid)[:, self.axis_index, 3] + for step in free_steps + } + for env_id in range(self.num_envs): + ordered_steps = sorted( + free_steps, + key=lambda step: ( + float(axis_positions[step.id][env_id]), + int(step.goal.get("nominal_slot_index", 0)), + step.id, + ), + ) + ordered_slots = sorted( + available_slots, + key=lambda slot_id: ( + float(self.positions[env_id, slot_id, self.axis_index]), + slot_id, + ), + ) + matching_cost = 0.0 + changed = False + for step, slot_id in zip(ordered_steps, ordered_slots): + nominal = int(step.goal.get("nominal_slot_index", 0)) + assignments[step.id][env_id] = slot_id + changed |= slot_id != nominal + matching_cost += abs( + float(axis_positions[step.id][env_id]) + - float(self.positions[env_id, slot_id, self.axis_index]) + ) + if changed: + self.reassignment_reason[env_id] = ( + "free arrangement initialized from live spatial order" + ) + self.reassignment_cost[env_id] = matching_cost + return assignments + + def _geometry(self, step: SemanticStep) -> _Geometry: + entity = _object(self.env, step.object_uid) + radii = [] + heights = [] + for env_id in range(self.num_envs): + vertices = _local_vertices(entity, self.env, env_id) + half_extent = ( + vertices.max(dim=0).values - vertices.min(dim=0).values + ) * 0.5 + if step.goal.get("orientation_goal", "none") in {"none", "preserve"}: + rotation = _live_pose(self.env, step.object_uid)[env_id, :3, :3] + rotated = vertices @ rotation.transpose(0, 1) + radii.append(torch.linalg.vector_norm(rotated[:, :2], dim=-1).max()) + else: + # A non-preserve target may rotate the longest local dimension + # into the table plane, so retain the conservative bound. + radii.append( + torch.linalg.vector_norm(torch.topk(half_extent, k=2).values) + ) + heights.append((vertices[:, 2].max() - vertices[:, 2].min()) * 0.5) + return _Geometry(torch.stack(radii), torch.stack(heights)) + + def _make_slots(self) -> torch.Tensor: + offsets = ( + torch.arange(self.slot_count, device=self.device, dtype=torch.float32) + - (self.slot_count - 1) / 2.0 + ) + slots = torch.empty( + self.num_envs, + self.slot_count, + 3, + dtype=torch.float32, + device=self.device, + ) + radii = torch.stack( + [self.geometry[step.id].radius for step in self.steps], + dim=1, + ) + # Free slot rematching allows any remaining object to occupy any slot. + # Size every slot for the largest member in that environment rather + # than accidentally baking the nominal object order into geometry. + slot_radii = radii.max(dim=1).values[:, None].repeat(1, self.slot_count) + obstacles = self._obstacle_bounds() + search_offsets = [0.0] + steps = int(self.row_search_radius / self.row_search_step) + for index in range(1, steps + 1): + offset = self.row_search_step * index + search_offsets.extend((offset, -offset)) + for env_id in range(self.num_envs): + chosen = None + for perpendicular in search_offsets: + candidate = self.table_center[env_id].repeat(self.slot_count, 1) + candidate[:, self.axis_index] += self.spacing[env_id] * offsets + candidate[:, self.perpendicular_index] += perpendicular + candidate[:, 2] = self.table_top[env_id] + if self._safe( + candidate, + slot_radii[env_id], + self.table_bounds[env_id], + obstacles[env_id], + ): + chosen = candidate + break + if chosen is None: + raise ValueError( + f"Environment {env_id} has no collision-free arrangement row." + ) + slots[env_id] = chosen + return slots + + def _obstacle_bounds( + self, + ) -> list[list[tuple[torch.Tensor, torch.Tensor]]]: + result: list[list[tuple[torch.Tensor, torch.Tensor]]] = [ + [] for _ in range(self.num_envs) + ] + getter = getattr(self.env.sim, "get_rigid_object_uid_list", None) + if not callable(getter): + return result + movable = {step.object_uid for step in self.steps} + for uid in getter(): + if uid == "table" or uid in movable: + continue + entity = self.env.sim.get_rigid_object(uid) + if entity is None: + continue + for env_id in range(self.num_envs): + vertices = _world_vertices(entity, self.env, env_id) + if float(vertices[:, 2].max()) < float( + self.table_top[env_id] - self.clearance + ): + continue + result[env_id].append( + ( + vertices[:, :2].min(dim=0).values, + vertices[:, :2].max(dim=0).values, + ) + ) + return result + + def _safe( + self, + slots: torch.Tensor, + radii: torch.Tensor, + table_bounds: torch.Tensor, + obstacles: Sequence[tuple[torch.Tensor, torch.Tensor]], + ) -> bool: + lower = table_bounds[0, :2] + radii[:, None] + self.clearance + upper = table_bounds[1, :2] - radii[:, None] - self.clearance + if bool(((slots[:, :2] < lower) | (slots[:, :2] > upper)).any()): + return False + for center, radius in zip(slots[:, :2], radii): + for obstacle_lower, obstacle_upper in obstacles: + closest = torch.maximum( + obstacle_lower, + torch.minimum(center, obstacle_upper), + ) + if float(torch.linalg.vector_norm(center - closest)) <= float( + radius + self.clearance + ): + return False + return True + + def target( + self, + step: SemanticStep, + object_pose: torch.Tensor, + *, + phase: str, + policy: Mapping[str, Any], + ) -> torch.Tensor: + """Return a live final or collision-clear staging object pose.""" + if phase not in {"staging", "final"}: + raise ValueError(f"Unsupported arrangement phase {phase!r}.") + target = object_pose.clone() + env_ids = torch.arange(self.num_envs, device=self.device) + slot_ids = self.assignments[step.id] + target[:, :2, 3] = self.positions[env_ids, slot_ids, :2] + final_z = ( + self.table_top + + self.geometry[step.id].half_height + + float(policy["surface_clearance"]) + ) + target[:, 2, 3] = final_z + if phase == "staging": + target[:, 2, 3] = final_z + float(policy["transport_clearance"]) + return target + + def mark_completed(self, step_id: str, success: torch.Tensor) -> None: + self.completed[step_id] |= success.to(self.device, dtype=torch.bool) + + def remaining(self, env_id: int) -> list[str]: + return [ + step.id for step in self.steps if not bool(self.completed[step.id][env_id]) + ] + + def available_slots(self, env_id: int) -> list[int]: + occupied = { + int(self.assignments[step.id][env_id]) + for step in self.steps + if bool(self.completed[step.id][env_id]) + } + return [index for index in range(self.slot_count) if index not in occupied] + + def assign(self, env_id: int, assignment: Mapping[str, int]) -> None: + for step_id, slot_id in assignment.items(): + self.assignments[step_id][env_id] = int(slot_id) + + def metadata(self, step: SemanticStep, env_id: int) -> dict[str, Any]: + """Describe the live slot resolution used by one environment.""" + nominal = int(step.goal.get("nominal_slot_index", 0)) + resolved = int(self.assignments[step.id][env_id]) + return { + "nominal_slot_index": nominal, + "resolved_slot_index": resolved, + "slot_constraint": str(step.goal.get("slot_constraint", "required")), + "slot_reassigned": resolved != nominal, + "reassignment_reason": self.reassignment_reason[env_id], + "matching_cost": ( + float(self.reassignment_cost[env_id]) + if torch.isfinite(self.reassignment_cost[env_id]) + else None + ), + "spacing": float(self.spacing[env_id]), + "resolved_slot_position": self.positions[env_id, resolved].tolist(), + } + + +class LivePlacementPlan: + """Allocate non-overlapping live slots for one shared container.""" + + def __init__( + self, + env: Any, + steps: Sequence[SemanticStep], + *, + clearance: float | None = None, + ) -> None: + if not steps: + raise ValueError("A placement plan requires at least one step.") + references = {step.goal.get("reference_object") for step in steps} + if len(references) != 1 or not isinstance(next(iter(references)), str): + raise ValueError("Placement-plan steps must share one reference object.") + self.env = env + self.steps = tuple(steps) + self.reference_uid = str(next(iter(references))) + self.num_envs = int(env.num_envs) + profile = str(getattr(env, "agent_robot_profile", "dual_ur10")) + default_clearance = default_runtime_policy(profile).grounding["placement"][ + "clearance" + ] + self.clearance = float(default_clearance if clearance is None else clearance) + self.positions = self._make_slots() + + def _make_slots(self) -> dict[str, torch.Tensor]: + container = _object(self.env, self.reference_uid) + positions = { + step.id: torch.empty( + self.num_envs, + 3, + dtype=torch.float32, + device=self.env.device, + ) + for step in self.steps + } + named_slots = [str(step.goal.get("slot", "auto")) for step in self.steps] + for slot in named_slots: + if slot not in {"auto", "left", "center", "right"}: + raise ValueError(f"Unsupported container slot {slot!r}.") + + for env_id in range(self.num_envs): + vertices = _world_vertices(container, self.env, env_id) + lower = vertices.min(dim=0).values + upper = vertices.max(dim=0).values + center = (lower + upper) * 0.5 + extent = upper[:2] - lower[:2] + axis = int(torch.argmax(extent).item()) + radii = [] + for step in self.steps: + moved_vertices = _local_vertices( + _object(self.env, step.object_uid), + self.env, + env_id, + ) + half = ( + moved_vertices.max(dim=0).values - moved_vertices.min(dim=0).values + )[:2] * 0.5 + radii.append(float(torch.linalg.vector_norm(half))) + radius = max(radii) + usable_span = float(extent[axis]) - 2.0 * (radius + self.clearance) + required_span = 2.0 * radius * max(len(self.steps) - 1, 0) + if usable_span + 1.0e-6 < required_span: + raise ValueError( + f"Environment {env_id} container {self.reference_uid!r} " + "has no non-overlapping slot plan." + ) + offsets = torch.linspace( + -required_span * 0.5, + required_span * 0.5, + len(self.steps), + device=self.env.device, + ) + named_offsets = { + "left": required_span * 0.5, + "center": 0.0, + "right": -required_span * 0.5, + } + used: list[float] = [] + for index, step in enumerate(self.steps): + slot = named_slots[index] + offset = ( + float(offsets[index]) if slot == "auto" else named_offsets[slot] + ) + if any(abs(offset - item) < 2.0 * radius for item in used): + raise ValueError( + f"Container slot {slot!r} overlaps another requested slot." + ) + used.append(offset) + target = center.clone() + target[axis] += offset + target[2] = lower[2] + positions[step.id][env_id] = target + return positions + + def target( + self, + step: SemanticStep, + object_pose: torch.Tensor, + rotation: torch.Tensor, + *, + surface_clearance: float, + ) -> torch.Tensor: + """Return a slot pose corrected for the rotated object mesh bottom.""" + target = object_pose.clone() + target[:, :3, :3] = rotation + target[:, :2, 3] = self.positions[step.id][:, :2] + entity = _object(self.env, step.object_uid) + for env_id in range(self.num_envs): + bottom = ( + _local_vertices(entity, self.env, env_id) + @ rotation[env_id].transpose(0, 1) + )[:, 2].min() + target[env_id, 2, 3] = ( + self.positions[step.id][env_id, 2] + surface_clearance - bottom + ) + return target + + +class ActionGrounder: + """Translate one symbolic action into a public typed atomic-action target.""" + + def __init__( + self, + program: ExecutionProgram, + env: Any, + semantics_factory: Callable[[str], ObjectSemantics], + arrangement: ( + LiveArrangementPlan | Mapping[str, LiveArrangementPlan] | None + ) = None, + placements: Mapping[str, LivePlacementPlan] | None = None, + runtime_policy: RuntimePolicyCfg | None = None, + capability_registry: Any | None = None, + ) -> None: + self.program = program + self.env = env + self.semantics_factory = semantics_factory + self.capabilities = capability_registry or build_atomic_capability_registry() + self.robot_profile = str(getattr(env, "agent_robot_profile", "dual_ur10")) + self.runtime_policy = runtime_policy or default_runtime_policy( + self.robot_profile + ) + if isinstance(arrangement, Mapping): + self.arrangements = dict(arrangement) + elif arrangement is None: + self.arrangements = {} + else: + self.arrangements = { + step.id: arrangement + for step in program.semantic_steps + if step.operator in {"arrange_line", "place_in_line"} + } + self.placements = dict(placements or {}) + + def policy( + self, + action: Mapping[str, Any], + *, + extra_modifiers: tuple[tuple[str, str], ...] = (), + ) -> dict[str, Any]: + action_class = str(action.get("atomic_action_class", "")) + capability = self.capabilities.get(action_class) + motion_base = capability.motion_base or capability.name + policy_spec = action.get("motion_policy", {"modifiers": []}) + if extra_modifiers: + policy_spec = with_motion_modifiers(policy_spec, *extra_modifiers) + inline = action.get("motion_policy_config", action.get("cfg")) + return resolve_motion_policy( + self.robot_profile, + motion_base, + policy_spec, + motion_defaults=self.runtime_policy.motion_defaults, + motion_modifiers=self.runtime_policy.motion_modifiers, + inline_overrides=inline if isinstance(inline, Mapping) else None, + ) + + def _policy_value(self, policy: Mapping[str, Any], key: str) -> Any: + defaults = self.runtime_policy.grounding["semantic_defaults"] + return policy[key] if key in policy else defaults[key] + + def ground( + self, + action: Mapping[str, Any], + step: SemanticStep, + *, + arm: str, + state: ExecutionState, + reference_eef_pose: torch.Tensor | None = None, + orientation_reference_pose: torch.Tensor | None = None, + _handover_workspace: tuple[torch.Tensor, torch.Tensor] | None = None, + ) -> GroundedAction: + action_class = str(action["atomic_action_class"]) + capability = self.capabilities.require_executable(action_class) + self.capabilities.validate_binding(action) + control = str(action.get("control", "arm")) + binding = action.get("target_binding", {}) + if not isinstance(binding, Mapping): + raise ValueError("target_binding must be a mapping.") + kind = str(binding.get("kind", "")) + orientation = compile_orientation_constraint(step.goal) + is_handover_continuation = self._is_handover_continuation(step) + uses_handover_staging = ( + kind == "handover_staging" + and capability.target_materializer == "semantic_held_object" + ) + use_upright_yaw_search = ( + is_handover_continuation or uses_handover_staging + ) and self._uses_upright_yaw_search( + step, + orientation, + ) + extra_modifiers: tuple[tuple[str, str], ...] = () + if ( + is_handover_continuation + and use_upright_yaw_search + and capability.target_materializer + in { + "semantic_held_object", + "current_held_pose", + "eef_pose", + } + ): + extra_modifiers = (("orientation", "upright"),) + policy = self.policy(action, extra_modifiers=extra_modifiers) + if kind == "joint_state": + joint_defaults = self.runtime_policy.grounding["joint_state"] + source = binding.get("source") + if source == "gripper_closed": + policy["sample_interval"] = int( + joint_defaults["hand_close_sample_interval"] + ) + elif source == "gripper_open": + policy["sample_interval"] = int( + joint_defaults["hand_open_sample_interval"] + ) + elif source == "initial" and control == "arm": + # Returning home after release is a safety motion. If the + # collision-aware planner cannot find a route, do not silently + # replace it with collision-unaware joint interpolation. + policy["collision_safety"] = "required" + if uses_handover_staging and use_upright_yaw_search: + # Handover consumes the live payload pose immediately after this + # move. Use the existing upright-yaw feasibility search instead + # of the generic transport orientation heuristic, which can tilt + # a payload while moving it to the exchange point. + policy["upright_yaw_samples"] = max( + int(policy.get("upright_yaw_samples", 1)), + 8, + ) + object_pose = _live_pose(self.env, step.object_uid) + if step.operator == "orient_object": + policy["upright_local_axis"] = self._upright_local_axis(step) + if capability.target_materializer == "object_grasp": + policy["obj_upright_direction"] = self._upright_local_direction(step) + reference_pose = self._reference_pose(step) + target_object_pose = None + + if capability.target_materializer_hook is not None: + grounded = capability.target_materializer_hook( + grounder=self, + action=action, + step=step, + arm=arm, + state=state, + binding=binding, + policy=policy, + object_pose=object_pose, + reference_pose=reference_pose, + reference_eef_pose=reference_eef_pose, + orientation_reference_pose=orientation_reference_pose, + ) + if not isinstance(grounded, GroundedAction): + raise TypeError( + f"AtomicAction {action_class!r} target materializer must " + "return GroundedAction." + ) + return grounded + + if kind == "object": + semantics = self.semantics_factory( + str(binding.get("object", step.object_uid)) + ) + if capability.target_materializer == "object_grasp": + target: Any = GraspGoal(semantics=semantics) + elif capability.target_materializer == "coordinated_pickment": + target_object_pose = self._semantic_target( + step, + object_pose, + reference_pose, + policy, + phase="final", + orientation_reference_pose=orientation_reference_pose, + ) + target = CoordinatedPickGoal( + object_target_pose=target_object_pose, + semantics=semantics, + object_initial_pose=object_pose, + ) + elif capability.target_materializer == "press": + target_object_pose = object_pose.clone() + target = self._press_goal( + step.object_uid, + object_pose, + semantics=semantics, + ) + else: + raise ValueError( + f"{action_class} does not support object target bindings." + ) + elif kind in {"semantic_goal", "coordinated_goal"}: + phase = str(binding.get("phase", "final")) + target_object_pose = self._semantic_target( + step, + object_pose, + reference_pose, + policy, + phase=phase, + orientation_reference_pose=orientation_reference_pose, + ) + if capability.target_materializer == "coordinated_pickment": + semantics = self.semantics_factory(step.object_uid) + target = CoordinatedPickGoal( + object_target_pose=target_object_pose, + semantics=semantics, + object_initial_pose=object_pose, + ) + elif capability.target_materializer == "press": + # Press moves the TCP, not the target object. Keep the object's + # live pose as the postcondition reference while grounding a + # downward contact point from its current surface geometry. + target_object_pose = object_pose.clone() + target = self._press_goal( + step.object_uid, + object_pose, + ) + elif capability.target_materializer == "semantic_held_object": + target = HeldObjectPoseGoal(object_target_pose=target_object_pose) + else: + raise ValueError( + f"Target materializer {capability.target_materializer!r} cannot " + f"resolve {kind!r}." + ) + elif kind == "coordinated_placement_goal": + support_uid = binding.get( + "support_object", + step.goal.get("support_object"), + ) + placing_uid = binding.get("placing_object", step.object_uid) + if not isinstance(placing_uid, str) or not placing_uid: + raise ValueError("coordinated_placement_goal requires placing_object.") + if not isinstance(support_uid, str) or not support_uid: + raise ValueError("coordinated_placement_goal requires support_object.") + support_pose = _live_pose(self.env, support_uid) + target_object_pose = self._semantic_target( + step, + object_pose, + support_pose, + policy, + phase="final", + orientation_reference_pose=orientation_reference_pose, + ) + target = CoordinatedPlacementGoal( + placing_object_target_pose=target_object_pose, + support_object_target_pose=support_pose, + release=bool(step.goal.get("release", True)), + ) + elif kind == "current_held_pose": + if state.get_held_object(arm_control_part(self.env, arm)) is None: + raise ValueError("Place requires a held object from a prior PickUp.") + target = PlaceGoal( + xpos=( + reference_eef_pose + if reference_eef_pose is not None + else self._current_eef_pose(arm) + ) + ) + elif kind == "policy_pose": + source = binding.get("source") + retreat_reference = self._retreat_reference_pose( + arm, + reference_eef_pose, + ) + if binding.get("operation") == "retreat": + policy["retreat_reachability_search"] = True + policy["retreat_reference_pose"] = retreat_reference.clone() + if source in {"release", "handover"}: + policy["clearance_object_uid"] = step.object_uid + policy["collision_safety"] = "required" + contact_uids = [step.object_uid] + reference_uid = step.goal.get("reference_object") + if isinstance(reference_uid, str) and reference_uid: + contact_uids.append(reference_uid) + policy["collision_exclusion_uids"] = list(dict.fromkeys(contact_uids)) + if source == "handover": + policy.update(self.runtime_policy.grounding["handover"]) + policy["transfer_arm"] = arm + policy["transfer_role_axis"] = self._handover_role_axis( + arm, + dtype=object_pose.dtype, + device=object_pose.device, + ) + target = EndEffectorPoseGoal( + xpos=self._retreat_pose( + arm, + policy, + retreat_reference, + clear_exchange=source == "handover", + ) + ) + elif kind == "visual_constraint": + visual_pose = self._visual_target(binding, arm) + if capability.target_materializer == "semantic_held_object": + target_object_pose = object_pose.clone() + target_object_pose[:, :3, 3] = visual_pose[:, :3, 3] + target = HeldObjectPoseGoal(object_target_pose=target_object_pose) + elif capability.target_materializer == "eef_pose": + target = EndEffectorPoseGoal(xpos=visual_pose) + else: + raise ValueError( + f"Target materializer {capability.target_materializer!r} " + "cannot resolve a visual_constraint." + ) + elif kind == "joint_state": + target = JointPositionGoal( + target=self._joint_target( + arm, + control, + str(binding.get("source", "initial")), + binding, + ) + ) + elif kind in {"eef_pose", "pose"}: + target = EndEffectorPoseGoal(xpos=self._explicit_pose(binding, object_pose)) + elif kind == "handover_goal": + target, target_object_pose, policy = self._handover_target( + step, + binding, + object_pose, + reference_pose, + policy, + state, + orientation_reference_pose=orientation_reference_pose, + workspace=_handover_workspace, + ) + elif kind == "handover_staging": + transfer_arm = str(binding.get("transfer_arm", "left_arm")) + receive_arm = str(binding.get("receive_arm", "right_arm")) + middle, _ = self._handover_workspace_poses( + object_pose, + transfer_arm=transfer_arm, + receive_arm=receive_arm, + policy=policy, + step=step, + orientation_reference_pose=orientation_reference_pose, + ) + middle[:, :3, :3] = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + target_object_pose = middle + target = HeldObjectPoseGoal(object_target_pose=middle) + else: + raise ValueError(f"Unsupported target binding kind {kind!r}.") + return GroundedAction( + action_class=action_class, + arm=arm, + control=control, + target=target, + cfg=policy, + object_pose=object_pose, + reference_pose=reference_pose, + target_object_pose=target_object_pose, + motion_policy=policy, + object_uid=step.object_uid, + ) + + def _handover_role_axis( + self, + transfer_arm: str, + *, + dtype: torch.dtype, + device: torch.device, + ) -> torch.Tensor: + """Return the world-space axis from the receiver base to transfer base.""" + if transfer_arm not in {"left_arm", "right_arm"}: + raise ValueError(f"Unknown handover arm {transfer_arm!r}.") + _, lateral = robot_frame_axes(self.env) + horizontal = lateral if transfer_arm == "left_arm" else -lateral + return torch.cat( + ( + horizontal.to(dtype=dtype, device=device), + torch.zeros( + (int(self.env.num_envs), 1), + dtype=dtype, + device=device, + ), + ), + dim=1, + ) + + def ground_candidates( + self, + action: Mapping[str, Any], + step: SemanticStep, + *, + arm: str, + state: ExecutionState, + reference_eef_pose: torch.Tensor | None = None, + orientation_reference_pose: torch.Tensor | None = None, + ) -> tuple[GroundedAction, ...]: + """Return deterministic grounding candidates for an opt-in capability.""" + binding = action.get("target_binding", {}) + if not isinstance(binding, Mapping): + return ( + self.ground( + action, + step, + arm=arm, + state=state, + reference_eef_pose=reference_eef_pose, + orientation_reference_pose=orientation_reference_pose, + ), + ) + placement_support_uid = self._placement_support_uid(step) + placement_relation = ( + normalize_placement_relation(step.goal.get("relation", "on")) + if step.operator == "place_relative" + else str(step.goal.get("relation", "none")) + ) + is_on_placement = ( + binding.get("kind") == "semantic_goal" + and binding.get("phase", "final") != "staging" + and placement_relation in {"on", "on_top", "on_top_of"} + and placement_support_uid is not None + ) + if is_on_placement: + base = self.ground( + action, + step, + arm=arm, + state=state, + reference_eef_pose=reference_eef_pose, + orientation_reference_pose=orientation_reference_pose, + ) + return self._placement_grounding_candidates( + base, + step, + support_uid=placement_support_uid, + ) + if binding.get("kind") != "handover_goal": + return ( + self.ground( + action, + step, + arm=arm, + state=state, + reference_eef_pose=reference_eef_pose, + orientation_reference_pose=orientation_reference_pose, + ), + ) + policy = self.policy(action) + object_pose = _live_pose(self.env, step.object_uid) + rotation = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + workspaces = self._handover_workspace_candidates( + step, + object_pose, + transfer_arm=str(binding.get("transfer_arm", "left_arm")), + receive_arm=str(binding.get("receive_arm", "right_arm")), + policy=policy, + rotation=rotation, + ) + return tuple( + self.ground( + action, + step, + arm=arm, + state=state, + reference_eef_pose=reference_eef_pose, + orientation_reference_pose=orientation_reference_pose, + _handover_workspace=workspace, + ) + for workspace in workspaces + ) + + def _placement_grounding_candidates( + self, + base: GroundedAction, + step: SemanticStep, + *, + support_uid: str, + ) -> tuple[GroundedAction, ...]: + """Sample bounded support-relative poses from live object geometry.""" + if base.target_object_pose is None or not isinstance( + base.target, HeldObjectPoseGoal + ): + return (base,) + support = _object(self.env, support_uid) + moved = _object(self.env, step.object_uid) + placement = self.runtime_policy.grounding["placement"] + count = int(placement["candidate_count"]) + fraction = float(placement["candidate_offset_fraction"]) + margin = float(placement["support_margin"]) + patterns = ( + (0.0, 0.0), + (1.0, 0.0), + (-1.0, 0.0), + (0.0, 1.0), + (0.0, -1.0), + (1.0, 1.0), + (1.0, -1.0), + (-1.0, 1.0), + (-1.0, -1.0), + )[:count] + candidates: list[GroundedAction] = [] + seen_offsets: list[torch.Tensor] = [] + for candidate_index, pattern in enumerate(patterns): + target_pose = base.target_object_pose.clone() + offsets = target_pose.new_zeros((int(self.env.num_envs), 2)) + for env_id in range(int(self.env.num_envs)): + support_vertices = _world_vertices(support, self.env, env_id) + moved_local = _local_vertices(moved, self.env, env_id) + rotated = moved_local @ target_pose[env_id, :3, :3].transpose(0, 1) + support_lower = support_vertices[:, :2].min(dim=0).values + support_upper = support_vertices[:, :2].max(dim=0).values + moved_lower = rotated[:, :2].min(dim=0).values + moved_upper = rotated[:, :2].max(dim=0).values + allowed_lower = support_lower + margin - moved_lower + allowed_upper = support_upper - margin - moved_upper + if bool(torch.all(allowed_lower <= allowed_upper)): + base_xy = target_pose[env_id, :2, 3].clone() + center = torch.minimum( + torch.maximum(base_xy, allowed_lower), + allowed_upper, + ) + direction = target_pose.new_tensor(pattern) + room = torch.where( + direction >= 0.0, + allowed_upper - center, + center - allowed_lower, + ) + candidate_xy = center + direction * room * fraction + offsets[env_id] = candidate_xy - base_xy + target_pose[env_id, :2, 3] = candidate_xy + + footprint_lower = target_pose[env_id, :2, 3] + moved_lower + footprint_upper = target_pose[env_id, :2, 3] + moved_upper + local_mask = torch.all( + (support_vertices[:, :2] >= footprint_lower - margin) + & (support_vertices[:, :2] <= footprint_upper + margin), + dim=1, + ) + if bool(local_mask.any()): + support_height = support_vertices[local_mask, 2].max() + else: + distances = torch.linalg.vector_norm( + support_vertices[:, :2] - target_pose[env_id, :2, 3], + dim=1, + ) + nearest_count = min(8, int(support_vertices.shape[0])) + nearest = torch.topk( + distances, + nearest_count, + largest=False, + ).indices + support_height = support_vertices[nearest, 2].max() + target_pose[env_id, 2, 3] = ( + support_height + + float(self._policy_value(base.motion_policy, "surface_clearance")) + - rotated[:, 2].min() + ) + if any(torch.allclose(offsets, prior) for prior in seen_offsets): + continue + seen_offsets.append(offsets) + candidates.append( + replace( + base, + target=replace(base.target, object_target_pose=target_pose), + target_object_pose=target_pose, + motion_policy={ + **base.motion_policy, + "placement_candidate_index": candidate_index, + "placement_xy_offset": offsets, + }, + ) + ) + return tuple(candidates) or (base,) + + @staticmethod + def _placement_support_uid(step: SemanticStep) -> str | None: + value = step.goal.get("reference_object", step.goal.get("support_object")) + if isinstance(value, str) and value: + return value + if ( + step.postcondition.get("type") == "stack_layer_supported" + and int(step.goal.get("layer_index", -1)) == 0 + ): + return "table" + return None + + def _is_handover_continuation(self, step: SemanticStep) -> bool: + if step.operator != "place_relative": + return False + predecessors = { + candidate.id: candidate for candidate in self.program.semantic_steps + } + return any( + (predecessor := predecessors.get(dependency)) is not None + and predecessor.operator == "handover" + and predecessor.object_uid == step.object_uid + for dependency in step.depends_on + ) + + def _visual_target( + self, + binding: Mapping[str, Any], + arm: str, + ) -> torch.Tensor: + """Unproject one normalized image keypoint using live camera depth.""" + camera_uid = str(binding.get("camera_uid", "")) + sensor = self.env.sim.get_sensor(camera_uid) + if sensor is None: + raise ValueError(f"Unknown visual-constraint camera {camera_uid!r}.") + keypoint_value = binding.get("normalized_keypoint") + if keypoint_value is None: + bbox = binding.get("normalized_bbox") + if isinstance(bbox, Sequence) and len(bbox) == 4: + keypoint_value = [ + (float(bbox[0]) + float(bbox[2])) * 0.5, + (float(bbox[1]) + float(bbox[3])) * 0.5, + ] + if keypoint_value is None: + raise ValueError( + "visual_constraint requires a normalized keypoint or bbox in [0, 1]." + ) + keypoint = torch.as_tensor( + keypoint_value, + dtype=torch.float32, + device=self.env.device, + ).flatten() + if keypoint.numel() != 2 or bool( + ((~torch.isfinite(keypoint)) | (keypoint < 0.0) | (keypoint > 1.0)).any() + ): + raise ValueError( + "visual_constraint requires a normalized keypoint or bbox in [0, 1]." + ) + data = sensor.get_data() + if "depth" not in data: + raise ValueError( + f"Camera {camera_uid!r} must provide depth for visual Grounding." + ) + depth = torch.as_tensor(data["depth"], device=self.env.device).squeeze(-1) + if depth.ndim == 2: + depth = depth.unsqueeze(0).repeat(int(self.env.num_envs), 1, 1) + if depth.ndim != 3 or depth.shape[0] != int(self.env.num_envs): + raise ValueError("Camera depth must have shape (N, H, W) or (N, H, W, 1).") + height, width = depth.shape[-2:] + pixel_x = min(max(int(round(float(keypoint[0]) * (width - 1))), 0), width - 1) + pixel_y = min(max(int(round(float(keypoint[1]) * (height - 1))), 0), height - 1) + distance = depth[:, pixel_y, pixel_x].to(torch.float32) + if bool((~torch.isfinite(distance) | (distance <= 0.0)).any()): + raise ValueError("visual_constraint keypoint has no valid live depth.") + intrinsics = torch.as_tensor( + sensor.get_intrinsics(), + dtype=torch.float32, + device=self.env.device, + ) + if intrinsics.ndim == 2: + intrinsics = intrinsics.unsqueeze(0).repeat(int(self.env.num_envs), 1, 1) + camera_pose = torch.as_tensor( + sensor.get_arena_pose(to_matrix=True), + dtype=torch.float32, + device=self.env.device, + ) + if camera_pose.ndim == 2: + camera_pose = camera_pose.unsqueeze(0).repeat(int(self.env.num_envs), 1, 1) + fx = intrinsics[:, 0, 0] + fy = intrinsics[:, 1, 1] + cx = intrinsics[:, 0, 2] + cy = intrinsics[:, 1, 2] + point = torch.stack( + ( + (float(pixel_x) - cx) * distance / fx, + (float(pixel_y) - cy) * distance / fy, + distance, + torch.ones_like(distance), + ), + dim=1, + ) + world = torch.bmm(camera_pose, point.unsqueeze(-1)).squeeze(-1) + target = self._current_eef_pose(arm).clone() + target[:, :3, 3] = world[:, :3] + return target + + def _handover_target( + self, + step: SemanticStep, + binding: Mapping[str, Any], + object_pose: torch.Tensor, + reference_pose: torch.Tensor | None, + policy: Mapping[str, Any], + state: ExecutionState, + *, + orientation_reference_pose: torch.Tensor | None, + workspace: tuple[torch.Tensor, torch.Tensor] | None = None, + ) -> tuple[GraspGoal, torch.Tensor, dict[str, Any]]: + transfer_arm = str(binding.get("transfer_arm", "left_arm")) + receive_arm = str( + binding.get( + "receive_arm", + "right_arm" if transfer_arm == "left_arm" else "left_arm", + ) + ) + if transfer_arm == receive_arm or {transfer_arm, receive_arm} != { + "left_arm", + "right_arm", + }: + raise ValueError("HandOver requires distinct left_arm/right_arm roles.") + transfer_part = arm_control_part(self.env, transfer_arm) + held = state.get_held_object(transfer_part) + if held is None: + raise ValueError( + f"HandOver requires {transfer_arm} to hold {step.object_uid!r}." + ) + + del reference_pose + if workspace is None: + middle, final = self._handover_workspace_poses( + object_pose, + transfer_arm=transfer_arm, + receive_arm=receive_arm, + policy=policy, + step=step, + orientation_reference_pose=orientation_reference_pose, + ) + else: + middle, final = (item.clone() for item in workspace) + rotation = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + middle[:, :3, :3] = rotation + final[:, :3, :3] = rotation + semantics = self.semantics_factory(step.object_uid) + grounded_policy = dict(policy) + grounded_policy.update( + { + "transfer_arm": transfer_arm, + "receive_arm": receive_arm, + "middle_object_pose": middle, + "final_object_pose": final, + } + ) + return ( + GraspGoal(semantics=semantics), + middle, + grounded_policy, + ) + + def _handover_workspace_poses( + self, + object_pose: torch.Tensor, + *, + transfer_arm: str, + receive_arm: str, + policy: Mapping[str, Any], + step: SemanticStep, + orientation_reference_pose: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Choose the highest-ranked collision-aware handover workspace.""" + rotation = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + candidates = self._handover_workspace_candidates( + step, + object_pose, + transfer_arm=transfer_arm, + receive_arm=receive_arm, + policy=policy, + rotation=rotation, + ) + return candidates[0] + + def _handover_workspace_candidates( + self, + step: SemanticStep, + object_pose: torch.Tensor, + *, + transfer_arm: str, + receive_arm: str, + policy: Mapping[str, Any], + rotation: torch.Tensor, + ) -> tuple[tuple[torch.Tensor, torch.Tensor], ...]: + """Rank exchange poses inside the two arm workspaces and above obstacles.""" + if transfer_arm == receive_arm or {transfer_arm, receive_arm} != { + "left_arm", + "right_arm", + }: + raise ValueError("Handover workspace requires distinct arm roles.") + table = self.env.sim.get_rigid_object("table") + if table is not None and hasattr(table, "get_vertices"): + centers = [] + tops = [] + bounds = [] + for env_id in range(int(self.env.num_envs)): + vertices = _world_vertices(table, self.env, env_id) + lower = vertices[:, :2].min(dim=0).values + upper = vertices[:, :2].max(dim=0).values + centers.append((lower + upper) * 0.5) + tops.append(vertices[:, 2].max()) + bounds.append(torch.stack((lower, upper))) + center = torch.stack(centers) + table_top = torch.stack(tops) + table_bounds = torch.stack(bounds) + else: + left = self._current_eef_pose("left_arm") + right = self._current_eef_pose("right_arm") + center = (left[:, :2, 3] + right[:, :2, 3]) * 0.5 + table_top = object_pose[:, 2, 3] + extent = float(policy.get("exchange_candidate_offset", 0.16)) * 2.0 + table_bounds = torch.stack((center - extent, center + extent), dim=1) + + forward, lateral = robot_frame_axes(self.env) + left_base, right_base = arm_base_poses(self.env) + transfer_base = left_base if transfer_arm == "left_arm" else right_base + receive_base = right_base if receive_arm == "right_arm" else left_base + base_midpoint = (transfer_base[:, :2, 3] + receive_base[:, :2, 3]) * 0.5 + table_forward = torch.sum((center - base_midpoint) * forward, dim=1) + shared_center = base_midpoint + forward * table_forward[:, None] + offset = float(policy.get("exchange_candidate_offset", 0.16)) + obstacle_clearance = float(policy.get("exchange_obstacle_clearance", 0.04)) + tool_horizontal_envelope = float( + policy.get("exchange_gripper_horizontal_envelope", 0.035) + ) + float(policy.get("exchange_wrist_horizontal_envelope", 0.055)) + tool_vertical_envelope = float( + policy.get("exchange_gripper_vertical_envelope", 0.025) + ) + float(policy.get("exchange_wrist_vertical_envelope", 0.04)) + minimum_reach = float(policy.get("exchange_minimum_reach", 0.10)) + maximum_reach = float(policy.get("exchange_maximum_reach", 1.00)) + if not 0.0 <= minimum_reach < maximum_reach: + raise ValueError("Handover reach bounds require 0 <= minimum < maximum.") + requested_count = max(1, int(policy.get("exchange_candidate_count", 4))) + object_clearance = float(policy.get("exchange_clearance", 0.06)) + if ( + min( + obstacle_clearance, + tool_horizontal_envelope, + tool_vertical_envelope, + object_clearance, + ) + < 0.0 + ): + raise ValueError("Handover geometry clearances must be non-negative.") + xy_coefficients = ( + (0.0, 0.0), + (1.0, 0.0), + (-1.0, 0.0), + (2.0, 0.0), + (-2.0, 0.0), + (0.0, 0.5), + (0.0, -0.5), + ) + ranked_by_env: list[list[tuple[float, torch.Tensor]]] = [] + moved = _object(self.env, step.object_uid) + obstacle_uids = ( + self.env.sim.get_rigid_object_uid_list() + if hasattr(self.env.sim, "get_rigid_object_uid_list") + else [] + ) + for env_id in range(int(self.env.num_envs)): + local_vertices = _local_vertices(moved, self.env, env_id) + rotated = local_vertices @ rotation[env_id].transpose(0, 1) + half_xy = ( + rotated[:, :2].max(dim=0).values - rotated[:, :2].min(dim=0).values + ) * 0.5 + bottom = rotated[:, 2].min() + margin = half_xy + obstacle_clearance + tool_horizontal_envelope + lower_limit = table_bounds[env_id, 0] + margin + upper_limit = table_bounds[env_id, 1] - margin + options: list[tuple[float, torch.Tensor]] = [] + for forward_scale, lateral_scale in xy_coefficients: + xy = ( + shared_center[env_id] + + forward[env_id] * (offset * forward_scale) + + lateral[env_id] * (offset * lateral_scale) + ) + if bool(((xy < lower_limit) | (xy > upper_limit)).any()): + continue + transfer_distance = torch.linalg.vector_norm( + xy - transfer_base[env_id, :2, 3] + ) + receive_distance = torch.linalg.vector_norm( + xy - receive_base[env_id, :2, 3] + ) + if not ( + minimum_reach <= float(transfer_distance) <= maximum_reach + and minimum_reach <= float(receive_distance) <= maximum_reach + ): + continue + obstacle_score, nearby_obstacle_top = self._handover_obstacle_metrics( + xy, + env_id=env_id, + object_uid=step.object_uid, + obstacle_uids=obstacle_uids, + half_xy=half_xy, + clearance=obstacle_clearance + tool_horizontal_envelope, + ) + center_cost = float( + torch.linalg.vector_norm(xy - shared_center[env_id]) + ) + pose = object_pose[env_id].clone() + pose[:3, :3] = rotation[env_id] + pose[:2, 3] = xy + safety_floor = torch.maximum(table_top[env_id], nearby_obstacle_top) + safe_z = ( + safety_floor + object_clearance + tool_vertical_envelope - bottom + ) + pose[2, 3] = torch.maximum(object_pose[env_id, 2, 3], safe_z) + lift_cost = max( + 0.0, + float(pose[2, 3] - object_pose[env_id, 2, 3]), + ) + options.append( + (obstacle_score + center_cost * 0.25 + lift_cost * 0.1, pose) + ) + if not options: + raise ValueError( + "No handover exchange pose lies inside the table bounds and " + "the reachable intersection of both arm bases." + ) + options.sort(key=lambda item: item[0]) + ranked_by_env.append(options[:requested_count]) + + candidate_count = min( + requested_count, + max(len(options) for options in ranked_by_env), + ) + candidates = [] + for candidate_index in range(candidate_count): + middle = object_pose.clone() + for env_id, options in enumerate(ranked_by_env): + middle[env_id] = options[min(candidate_index, len(options) - 1)][1] + # The built-in HandOver primitive plans its final transfer/receiver + # phase concurrently. An exchange-to-exchange target makes that + # receiver path stationary; graph-level retreat/home nodes then + # clear the transfer arm before any receiver-side continuation. + final = middle.clone() + candidates.append((middle, final)) + return tuple(candidates) + + def _handover_obstacle_metrics( + self, + xy: torch.Tensor, + *, + env_id: int, + object_uid: str, + obstacle_uids: Sequence[str], + half_xy: torch.Tensor, + clearance: float, + ) -> tuple[float, torch.Tensor]: + score = 0.0 + highest_top = torch.tensor( + -torch.inf, + dtype=xy.dtype, + device=xy.device, + ) + for uid in obstacle_uids: + if uid in {"table", object_uid}: + continue + obstacle = self.env.sim.get_rigid_object(uid) + if obstacle is None or not hasattr(obstacle, "get_vertices"): + continue + vertices = _world_vertices(obstacle, self.env, env_id) + lower = vertices[:, :2].min(dim=0).values - half_xy - clearance + upper = vertices[:, :2].max(dim=0).values + half_xy + clearance + outside = torch.maximum( + torch.maximum(lower - xy, xy - upper), + torch.zeros_like(xy), + ) + if bool((outside > 0.0).any()): + distance = float(torch.linalg.vector_norm(outside)) + score += 1.0 / max(distance, 1.0e-3) + else: + score += 1.0e3 + highest_top = torch.maximum(highest_top, vertices[:, 2].max()) + return score, highest_top + + def _handover_receiver_exit( + self, + middle: torch.Tensor, + receive_arm: str, + policy: Mapping[str, Any], + ) -> torch.Tensor: + final = middle.clone() + receive_pose = self._current_eef_pose(receive_arm) + direction = receive_pose[:, :2, 3] - middle[:, :2, 3] + norm = torch.linalg.vector_norm(direction, dim=1, keepdim=True) + fallback = direction.new_zeros(direction.shape) + fallback[:, 1] = -1.0 if receive_arm == "right_arm" else 1.0 + direction = torch.where( + norm > 1.0e-6, direction / norm.clamp_min(1.0e-6), fallback + ) + final[:, :2, 3] += direction * min( + 0.12, + float(self._policy_value(policy, "relation_distance")) * 0.5, + ) + return final + + def _reference_pose(self, step: SemanticStep) -> torch.Tensor | None: + uid = step.goal.get("reference_object", step.goal.get("support_object")) + if not isinstance(uid, str) or not uid: + return None + if step.goal.get("reference_state") == "initial": + initial = getattr(self.env, "agent_initial_object_poses", {}).get(uid) + if initial is None: + raise ValueError(f"Initial pose for {uid!r} is unavailable.") + return _batched_pose(initial, self.env) + return _live_pose(self.env, uid) + + def _semantic_target( + self, + step: SemanticStep, + object_pose: torch.Tensor, + reference_pose: torch.Tensor | None, + policy: Mapping[str, Any], + *, + phase: str, + orientation_reference_pose: torch.Tensor | None = None, + ) -> torch.Tensor: + if step.operator in {"arrange_line", "place_in_line"}: + arrangement = self.arrangements.get(step.id) + if arrangement is None: + raise ValueError("arrange_line requires a live arrangement plan.") + target = arrangement.target( + step, + object_pose, + phase=phase, + policy=policy, + ) + target[:, :3, :3] = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + moved = _object(self.env, step.object_uid) + for env_id in range(int(self.env.num_envs)): + bottom = self._rotated_local_z_min( + moved, + target[env_id, :3, :3], + env_id, + ) + target[env_id, 2, 3] = ( + arrangement.table_top[env_id] + + float(policy["surface_clearance"]) + - bottom + ) + if phase == "staging": + target[:, 2, 3] += float(policy["transport_clearance"]) + return target + placement = self.placements.get(step.id) + if placement is not None: + target = placement.target( + step, + object_pose, + self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ), + surface_clearance=float(policy["surface_clearance"]), + ) + if phase == "staging": + target[:, 2, 3] += float(policy["transport_clearance"]) + return target + if step.operator == "orient_object": + initial = None + if step.goal.get("position_anchor", "initial_xy") == "initial_xy": + initial = getattr(self.env, "agent_initial_object_poses", {}).get( + step.object_uid + ) + target = ( + _batched_pose(initial, self.env).clone() + if initial is not None + else object_pose.clone() + ) + target[:, :3, :3] = self._target_rotation( + step, + target, + orientation_reference_pose=orientation_reference_pose, + ) + support_uid = str(step.goal.get("support_object", "table")) + support = _object(self.env, support_uid) + moved = _object(self.env, step.object_uid) + for env_id in range(int(self.env.num_envs)): + support_top = _world_vertices(support, self.env, env_id)[:, 2].max() + bottom = self._rotated_local_z_min( + moved, + target[env_id, :3, :3], + env_id, + ) + target[env_id, 2, 3] = ( + support_top + float(policy["surface_clearance"]) - bottom + ) + if phase == "staging": + target[:, 2, 3] += float(policy["staging_lift_height"]) + return target + target = object_pose.clone() + if reference_pose is not None: + target[:, :3, 3] = reference_pose[:, :3, 3] + # Operators without a relational goal (for example press or a + # direction-only coordinated transport) must preserve the live origin + # instead of being silently projected onto a synthetic table support. + relation = ( + normalize_placement_relation(step.goal.get("relation", "on")) + if step.operator == "place_relative" + else str(step.goal.get("relation", "none")) + ) + distance = float(self._policy_value(policy, "relation_distance")) + relation_frame = str(step.goal.get("relation_frame", "world")) + forward_distance = distance + lateral_distance = distance + if relation_frame == "robot" and reference_pose is not None: + nominal = float(policy.get("robot_relative_distance", 0.10)) + clearance = float(policy.get("relation_clearance", 0.02)) + reference_uid = str(step.goal.get("reference_object", "")) + if reference_uid: + forward_axis, lateral_axis = robot_frame_axes(self.env) + forward_distance = self._relative_object_spacing( + step.object_uid, + reference_uid, + axis=forward_axis, + nominal=nominal, + clearance=clearance, + ) + lateral_distance = self._relative_object_spacing( + step.object_uid, + reference_uid, + axis=lateral_axis, + nominal=nominal, + clearance=clearance, + ) + directional_offset = relation_offset( + self.env, + relation, + frame=relation_frame, + forward_distance=forward_distance, + lateral_distance=lateral_distance, + dtype=target.dtype, + device=target.device, + ) + offsets = { + "above": (0.0, 0.0, float(self._policy_value(policy, "hover_height"))), + "held_above_initial": ( + 0.0, + 0.0, + float(self._policy_value(policy, "hover_height")), + ), + } + if directional_offset is not None: + target[:, :3, 3] += directional_offset + elif (offset := offsets.get(relation)) is not None: + target[:, :3, 3] += torch.tensor( + offset, + dtype=target.dtype, + device=target.device, + ) + slot = str(step.goal.get("slot", "auto")) + if relation in {"on", "on_top", "on_top_of", "inside"} and slot in { + "left", + "right", + }: + slot_offset = relation_offset( + self.env, + slot, + frame=relation_frame, + forward_distance=forward_distance, + lateral_distance=lateral_distance, + dtype=target.dtype, + device=target.device, + ) + if slot_offset is not None: + target[:, :3, 3] += slot_offset + direction = str(step.goal.get("direction", "none")) + direction_offsets = { + "world_x": (distance, 0.0, 0.0), + "world_y": (0.0, distance, 0.0), + "up": (0.0, 0.0, distance), + "down": (0.0, 0.0, -distance), + } + planar_direction_offset = relation_offset( + self.env, + direction, + frame=relation_frame, + forward_distance=distance, + lateral_distance=distance, + dtype=target.dtype, + device=target.device, + ) + if planar_direction_offset is not None: + target[:, :3, 3] += planar_direction_offset + elif direction in direction_offsets: + target[:, :3, 3] += torch.tensor( + direction_offsets[direction], + dtype=target.dtype, + device=target.device, + ) + + root_stack_layer = ( + step.operator == "build_stack" + and int(step.goal.get("layer_index", 0)) == 0 + and reference_pose is None + ) + if root_stack_layer: + table = _object(self.env, "table") + for env_id in range(int(self.env.num_envs)): + vertices = _world_vertices(table, self.env, env_id) + target[env_id, :2, 3] = ( + vertices[:, :2].min(dim=0).values + + vertices[:, :2].max(dim=0).values + ) * 0.5 + + target[:, :3, :3] = self._target_rotation( + step, + object_pose, + orientation_reference_pose=orientation_reference_pose, + ) + if ( + step.operator == "coordinated_transport" + and relation not in {"on", "on_top", "on_top_of", "inside"} + and direction not in {"up", "down"} + ): + release = str(step.goal.get("terminal_behavior", "hold")) == "place" + if not release: + target[:, 2, 3] = object_pose[:, 2, 3] + float( + self._policy_value(policy, "transport_clearance") + ) + else: + table = _object(self.env, "table") + moved = _object(self.env, step.object_uid) + clearance = float(self._policy_value(policy, "surface_clearance")) + for env_id in range(int(self.env.num_envs)): + table_top = _world_vertices(table, self.env, env_id)[:, 2].max() + bottom = self._rotated_local_z_min( + moved, + target[env_id, :3, :3], + env_id, + ) + target[env_id, 2, 3] = table_top + clearance - bottom + if relation in {"on", "on_top", "on_top_of"} or root_stack_layer: + support_uid = ( + step.goal.get("reference_object") + or step.goal.get("support_object") + or "table" + ) + support = _object(self.env, str(support_uid)) + moved = _object(self.env, step.object_uid) + for env_id in range(int(self.env.num_envs)): + support_top = _world_vertices(support, self.env, env_id)[:, 2].max() + bottom = self._rotated_local_z_min( + moved, + target[env_id, :3, :3], + env_id, + ) + target[env_id, 2, 3] = ( + support_top + + float(self._policy_value(policy, "surface_clearance")) + - bottom + ) + elif relation == "inside" and reference_pose is not None: + # Grounding the final move happens after the staging lift. Preserve + # the pre-pick supported height rather than the lifted live height. + supported_pose = orientation_reference_pose + if supported_pose is None: + supported_pose = object_pose + supported_pose = _batched_pose(supported_pose, self.env) + target[:, 2, 3] = supported_pose[:, 2, 3] + if phase == "staging": + # Staging is a runtime waypoint, not a persisted coordinate. This + # keeps in-place orientation robust to the object's live height. + target[:, 2, 3] += float(self._policy_value(policy, "transport_clearance")) + elif self._is_handover_continuation(step) and relation not in { + "on", + "on_top", + "on_top_of", + "inside", + }: + # A handover can leave the live rigid-body center a few centimetres + # below the original table-supported height. Reusing that drifted + # height for the lateral placement target makes the can intersect + # the table during release and it may tip or slide. Preserve the + # predecessor's supported height for the final held-object pose. + supported_pose = orientation_reference_pose + if supported_pose is None: + supported_pose = object_pose + supported_pose = _batched_pose(supported_pose, self.env) + target[:, 2, 3] = torch.maximum( + target[:, 2, 3], + supported_pose[:, 2, 3], + ) + return target + + def _relative_object_spacing( + self, + moved_uid: str, + reference_uid: str, + *, + axis: int | torch.Tensor, + nominal: float, + clearance: float, + ) -> float: + """Return deterministic center spacing from live object extents.""" + moved = _object(self.env, moved_uid) + reference = _object(self.env, reference_uid) + required = float(nominal) + for env_id in range(int(self.env.num_envs)): + moved_vertices = _world_vertices(moved, self.env, env_id) + reference_vertices = _world_vertices(reference, self.env, env_id) + if isinstance(axis, torch.Tensor): + direction = axis[env_id].to( + dtype=moved_vertices.dtype, + device=moved_vertices.device, + ) + moved_axis = moved_vertices[:, :2] @ direction + reference_axis = reference_vertices[:, :2] @ direction + else: + moved_axis = moved_vertices[:, axis] + reference_axis = reference_vertices[:, axis] + moved_half = (moved_axis.max() - moved_axis.min()) * 0.5 + reference_half = (reference_axis.max() - reference_axis.min()) * 0.5 + required = max( + required, + float(moved_half + reference_half) + float(clearance), + ) + return required + + def _upright_local_direction(self, step: SemanticStep) -> torch.Tensor: + axis = self._upright_local_axis(step) + entity = _object(self.env, step.object_uid) + vertices = _local_vertices(entity, self.env, 0) + extents = vertices.max(dim=0).values - vertices.min(dim=0).values + if axis == "long_axis": + axis_index = int(torch.argmax(extents).item()) + else: + axis_index = {"x": 0, "y": 1, "z": 2}[axis] + direction = torch.zeros(3, dtype=torch.float32, device=self.env.device) + direction[axis_index] = 1.0 + return direction + + def _uses_upright_yaw_search( + self, + step: SemanticStep, + constraint: OrientationConstraint, + ) -> bool: + """Preserve a live upright state as a planning preference. + + Explicit full-frame matching cannot admit yaw search. With no hard + orientation terms, yaw search is enabled only when the live object's + long axis is already upright, so a preceding upright operation remains + stable without turning that state into a sticky acceptance constraint. + """ + if constraint.allows_upright_yaw_search: + return True + if ( + constraint.terms + or constraint.planning_preference != "minimize_rotation_from_current" + ): + return False + entity = _object(self.env, step.object_uid) + vertices = _local_vertices(entity, self.env, 0) + extents = vertices.max(dim=0).values - vertices.min(dim=0).values + axis_index = int(torch.argmax(extents).item()) + pose = _live_pose(self.env, step.object_uid) + cosine = pose[:, 2, axis_index].abs().clamp(0.0, 1.0) + tolerance = float(self.runtime_policy.predicate_fallbacks["upright_max_tilt"]) + return bool(torch.all(torch.arccos(cosine) <= tolerance).item()) + + @staticmethod + def _upright_local_axis(step: SemanticStep) -> str: + align_terms = tuple( + term + for term in compile_orientation_constraint(step.goal).terms + if isinstance(term, AlignAxisConstraint) + ) + if align_terms: + return align_terms[0].local_axis + axis = str(step.goal.get("upright_local_axis", "auto")) + return "long_axis" if axis == "auto" else axis + + def _target_rotation( + self, + step: SemanticStep, + object_pose: torch.Tensor, + *, + orientation_reference_pose: torch.Tensor | None = None, + ) -> torch.Tensor: + constraint = compile_orientation_constraint(step.goal) + if not constraint.terms: + return object_pose[:, :3, :3].clone() + if ( + len(constraint.terms) == 1 + and isinstance(constraint.terms[0], MatchRotationConstraint) + and constraint.terms[0].reference == "step_start" + ): + if orientation_reference_pose is not None: + reference = _batched_pose(orientation_reference_pose, self.env) + return reference[:, :3, :3].clone() + return object_pose[:, :3, :3].clone() + goal = str(step.goal.get("orientation_goal", "none")) + align_term = next( + ( + term + for term in constraint.terms + if isinstance(term, AlignAxisConstraint) + ), + None, + ) + if align_term is not None: + goal = "upright" + if goal not in {"upright", "lay_flat", "axis_align"}: + raise ValueError(f"Unsupported orientation_goal {goal!r}.") + + entity = _object(self.env, step.object_uid) + rotations = [] + for env_id in range(int(self.env.num_envs)): + vertices = _local_vertices(entity, self.env, env_id) + extents = vertices.max(dim=0).values - vertices.min(dim=0).values + longest_to_shortest = torch.argsort( + extents, + descending=True, + ).tolist() + if goal == "upright": + upright_axis = ( + align_term.local_axis + if align_term is not None + else self._upright_local_axis(step) + ) + vertical_axis = ( + int(longest_to_shortest[0]) + if upright_axis == "long_axis" + else {"x": 0, "y": 1, "z": 2}[upright_axis] + ) + horizontal_axis = next( + int(axis) + for axis in longest_to_shortest + if int(axis) != vertical_axis + ) + elif goal == "lay_flat": + vertical_axis = int(longest_to_shortest[-1]) + horizontal_axis = int(longest_to_shortest[0]) + else: + horizontal_axis = self._aligned_local_axis( + step, + longest_to_shortest, + ) + vertical_axis = next( + int(axis) + for axis in reversed(longest_to_shortest) + if int(axis) != horizontal_axis + ) + direction = self._horizontal_orientation( + step, + object_pose, + env_id, + horizontal_axis, + ) + rotations.append( + self._world_aligned_rotation( + direction, + horizontal_axis=horizontal_axis, + vertical_axis=vertical_axis, + ) + ) + return torch.stack(rotations) + + @staticmethod + def _aligned_local_axis( + step: SemanticStep, + longest_to_shortest: Sequence[int], + ) -> int: + axis = str(step.goal.get("orientation_axis", "long_axis")) + if axis == "x": + return 0 + if axis == "y": + return 1 + if axis == "long_axis": + return int(longest_to_shortest[0]) + if axis == "short_axis": + return int(longest_to_shortest[-1]) + raise ValueError(f"Unsupported axis_align orientation_axis {axis!r}.") + + def _horizontal_orientation( + self, + step: SemanticStep, + object_pose: torch.Tensor, + env_id: int, + local_axis: int, + ) -> torch.Tensor: + align_to = step.goal.get("orientation_reference_object") + if isinstance(align_to, str) and align_to: + reference = _object(self.env, align_to) + vertices = _local_vertices(reference, self.env, env_id) + extents = vertices.max(dim=0).values - vertices.min(dim=0).values + ordered = torch.argsort(extents, descending=True) + requested = str(step.goal.get("orientation_axis", "long_axis")) + reference_axis = int( + ordered[-1] if requested == "short_axis" else ordered[0] + ) + reference_pose = _live_pose(self.env, align_to) + direction = reference_pose[env_id, :3, reference_axis].clone() + elif step.operator in {"arrange_line", "place_in_line"}: + arrangement = self.arrangements.get(step.id) + axis_index = 0 if arrangement is None else arrangement.axis_index + direction = torch.zeros( + 3, + dtype=object_pose.dtype, + device=object_pose.device, + ) + direction[axis_index] = 1.0 + elif str(step.goal.get("orientation_axis", "")) in {"y", "world_y"}: + direction = object_pose.new_tensor([0.0, 1.0, 0.0]) + elif str(step.goal.get("orientation_axis", "")) in {"x", "world_x"}: + direction = object_pose.new_tensor([1.0, 0.0, 0.0]) + else: + direction = object_pose[env_id, :3, local_axis].clone() + direction[2] = 0.0 + norm = torch.linalg.vector_norm(direction) + if float(norm) < 1.0e-6: + return object_pose.new_tensor([1.0, 0.0, 0.0]) + return direction / norm + + @staticmethod + def _world_aligned_rotation( + horizontal_direction: torch.Tensor, + *, + horizontal_axis: int, + vertical_axis: int, + ) -> torch.Tensor: + world_up = horizontal_direction.new_tensor([0.0, 0.0, 1.0]) + remaining_axis = ({0, 1, 2} - {horizontal_axis, vertical_axis}).pop() + columns = [torch.zeros_like(world_up) for _ in range(3)] + columns[horizontal_axis] = horizontal_direction + columns[vertical_axis] = world_up + columns[remaining_axis] = torch.linalg.cross( + world_up, + horizontal_direction, + ) + rotation = torch.stack(columns, dim=1) + if float(torch.linalg.det(rotation)) < 0.0: + rotation[:, remaining_axis] *= -1.0 + return rotation + + def _rotated_local_z_min( + self, + entity: Any, + rotation: torch.Tensor, + env_id: int, + ) -> torch.Tensor: + vertices = _local_vertices(entity, self.env, env_id) + return (vertices @ rotation.transpose(0, 1))[:, 2].min() + + def _current_eef_pose(self, arm: str) -> torch.Tensor: + """Return the live TCP pose for one logical Action Engine arm.""" + if arm not in {"left_arm", "right_arm"}: + raise ValueError(f"Expected a physical arm, got {arm!r}.") + if hasattr(self.env, "get_current_xpos_agent"): + left, right = self.env.get_current_xpos_agent() + value = left if arm == "left_arm" else right + if value is not None: + return _batched_pose(value, self.env) + + is_left = arm == "left_arm" + if not hasattr(self.env, "get_agent_arm_control_part"): + raise ValueError("Coordinated placement requires live TCP poses.") + part = self.env.get_agent_arm_control_part(is_left) + qpos = self._arm_qpos(arm) + return _batched_pose( + self.env.robot.compute_fk(qpos=qpos, name=part, to_matrix=True), + self.env, + ) + + def _press_goal( + self, + uid: str, + object_pose: torch.Tensor, + *, + semantics: ObjectSemantics | None = None, + ) -> PressGoal: + """Ground the live top surface into the typed press-affordance contract.""" + if semantics is None: + semantics = self.semantics_factory(uid) + entity = _object(self.env, uid) + reference_pose = object_pose[0] + world_position = reference_pose[:3, 3].clone() + world_position[2] = _world_vertices(entity, self.env, 0)[:, 2].max() + rotation = reference_pose[:3, :3] + local_position = rotation.transpose(0, 1) @ ( + world_position - reference_pose[:3, 3] + ) + local_axis = rotation.transpose(0, 1) @ torch.tensor( + [0.0, 0.0, -1.0], + dtype=rotation.dtype, + device=rotation.device, + ) + press_semantics = replace( + semantics, + affordance=PressAffordance( + press_axis=local_axis, + press_position=tuple(float(value) for value in local_position), + ), + ) + return PressGoal( + semantics=press_semantics, + target_pose=object_pose.clone(), + ) + + def _retreat_pose( + self, + arm: str, + policy: Mapping[str, Any], + reference: torch.Tensor | None, + *, + clear_exchange: bool = False, + ) -> torch.Tensor: + target = self._retreat_reference_pose(arm, reference).clone() + desired = float(self._policy_value(policy, "retreat_height")) + if clear_exchange: + _, lateral = robot_frame_axes(self.env) + direction = lateral if arm == "left_arm" else -lateral + target[:, :2, 3] += direction.to( + dtype=target.dtype, + device=target.device, + ) * float(policy.get("retreat_distance", 0.10)) + desired = max( + desired, + float(self._policy_value(policy, "minimum_retreat_height")), + ) + ceiling = float(self._policy_value(policy, "maximum_eef_height")) + height = torch.clamp(ceiling - target[:, 2, 3], min=0.0, max=desired) + target[:, 2, 3] += height + return target + + def _retreat_reference_pose( + self, + arm: str, + reference: torch.Tensor | None, + ) -> torch.Tensor: + """Resolve the live or speculative TCP pose from which retreat starts.""" + pose = reference + if pose is None and hasattr(self.env, "get_current_xpos_agent"): + left, right = self.env.get_current_xpos_agent() + pose = left if arm == "left_arm" else right + if pose is None: + raise ValueError("Retreat grounding requires a live end-effector pose.") + return _batched_pose(pose, self.env) + + def _joint_target( + self, + arm: str, + control: str, + source: str, + binding: Mapping[str, Any], + ) -> torch.Tensor: + if source in {"gripper_closed", "gripper_open"}: + value = ( + getattr(self.env, "close_state") + if source == "gripper_closed" + else getattr(self.env, "open_state") + ) + return torch.as_tensor( + value, + dtype=torch.float32, + device=self.env.device, + ) + if source == "joint_delta": + current = self._arm_qpos(arm).clone() + index = int(binding["joint_index"]) + current[:, index] += torch.deg2rad( + torch.tensor( + float(binding.get("delta_degrees", 0.0)), + device=current.device, + ) + ) + return current + initial = getattr(self.env, "init_qpos", self.env.robot.get_qpos()) + joint_ids = self._joint_ids(arm, control) + return torch.as_tensor(initial, device=self.env.device)[:, joint_ids] + + def _arm_qpos(self, arm: str) -> torch.Tensor: + if hasattr(self.env, "get_current_qpos_agent"): + left, right = self.env.get_current_qpos_agent() + return torch.as_tensor( + left if arm == "left_arm" else right, + dtype=torch.float32, + device=self.env.device, + ) + return self.env.robot.get_qpos()[:, self._joint_ids(arm, "arm")] + + def _joint_ids(self, arm: str, control: str) -> list[int]: + side = "left" if arm == "left_arm" else "right" + key = f"{side}_{'eef' if control == 'hand' else 'arm'}_joints" + return list(getattr(self.env, key, ())) + + def _explicit_pose( + self, + binding: Mapping[str, Any], + object_pose: torch.Tensor, + ) -> torch.Tensor: + reference = str(binding.get("reference", "absolute")) + target = object_pose.clone() + if reference == "absolute": + values = binding.get("position_by_env", binding.get("position")) + position = torch.as_tensor( + values, + dtype=target.dtype, + device=target.device, + ) + if position.ndim == 1: + position = position.unsqueeze(0).repeat(int(self.env.num_envs), 1) + target[:, :3, 3] = position + return target + offset = torch.as_tensor( + binding.get("offset", (0.0, 0.0, 0.0)), + dtype=target.dtype, + device=target.device, + ) + target[:, :3, 3] += offset + return target + + def _coordinated_grasps( + self, + semantics: ObjectSemantics, + object_pose: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Build a deterministic opposing pair along the object's longest XY axis.""" + vertices = semantics.geometry.get("mesh_vertices") + vertices = torch.as_tensor( + vertices, + dtype=torch.float32, + device=self.env.device, + ) + lower = vertices.min(dim=0).values + upper = vertices.max(dim=0).values + axis = int(torch.argmax(upper[:2] - lower[:2]).item()) + center = (lower + upper) * 0.5 + grasp_policy = self.runtime_policy.grounding["coordinated_grasp"] + inset = max( + float(grasp_policy["minimum_inset"]), + float((upper[axis] - lower[axis]) * grasp_policy["inset_fraction"]), + ) + left = torch.eye(4, dtype=torch.float32, device=self.env.device) + right = left.clone() + left[:3, 3] = center + right[:3, 3] = center + left[axis, 3] = lower[axis] + inset + right[axis, 3] = upper[axis] - inset + # Keep TCP z horizontal and facing the object from opposite sides. + if axis == 0: + left[:3, :3] = torch.tensor( + [[0.0, 0.0, 1.0], [1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], + device=self.env.device, + ) + right[:3, :3] = torch.tensor( + [[0.0, 0.0, -1.0], [-1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], + device=self.env.device, + ) + else: + left[:3, :3] = torch.tensor( + [[1.0, 0.0, 0.0], [0.0, 0.0, 1.0], [0.0, -1.0, 0.0]], + device=self.env.device, + ) + right[:3, :3] = torch.tensor( + [[-1.0, 0.0, 0.0], [0.0, 0.0, -1.0], [0.0, -1.0, 0.0]], + device=self.env.device, + ) + batch = int(self.env.num_envs) + return left.unsqueeze(0).repeat(batch, 1, 1), right.unsqueeze(0).repeat( + batch, 1, 1 + ) diff --git a/embodichain/gen_sim/action_engine/runtime/predicates.py b/embodichain/gen_sim/action_engine/runtime/predicates.py new file mode 100644 index 000000000..aff6d1df7 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/predicates.py @@ -0,0 +1,746 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Evaluate canonical closed-loop predicates against live environment state.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any + +import torch + +from embodichain.gen_sim.action_engine.config import default_runtime_policy + +from .frames import relation_axes +from .robot_parts import arm_control_part + +__all__ = ["PREDICATE_TYPES", "evaluate_predicate"] + +PREDICATE_TYPES = frozenset( + { + "both_arms_at_initial_qpos", + "both_grippers_open", + "coordinated_placed", + "grippers_clear_of_object", + "held_by_both_grippers", + "object_axis_near", + "object_axis_offset_near", + "object_held", + "object_held_by_both_grippers", + "object_held_by_gripper", + "object_in_container", + "object_lifted", + "object_not_fallen", + "object_on_object", + "object_supported_by", + "object_position_near", + "object_relative_position", + "object_upright", + "object_xy_near", + "objects_collinear", + "objects_ordered", + "pressed", + } +) +_DEFAULT_PREDICATE_FALLBACKS = default_runtime_policy("dual_ur10").predicate_fallbacks + + +def _predicate_fallbacks(env: Any) -> Mapping[str, Any]: + policy = getattr(env, "runtime_policy", None) + value = getattr(policy, "predicate_fallbacks", None) + return value if isinstance(value, Mapping) else _DEFAULT_PREDICATE_FALLBACKS + + +def _constant(env: Any, value: bool) -> torch.Tensor: + return torch.full( + (int(env.num_envs),), + value, + dtype=torch.bool, + device=env.device, + ) + + +def _pose(env: Any, uid: str) -> torch.Tensor: + entity = env.sim.get_rigid_object(uid) + if entity is None: + raise ValueError(f"Unknown rigid object {uid!r}.") + pose = torch.as_tensor( + entity.get_local_pose(to_matrix=True), + dtype=torch.float32, + device=env.device, + ) + if pose.ndim == 2: + pose = pose.unsqueeze(0).repeat(int(env.num_envs), 1, 1) + return pose + + +def _position(env: Any, uid: str) -> torch.Tensor: + return _pose(env, uid)[:, :3, 3] + + +def _world_vertices(env: Any, uid: str, env_id: int) -> torch.Tensor: + entity = env.sim.get_rigid_object(uid) + if entity is None: + raise ValueError(f"Unknown rigid object {uid!r}.") + value = entity.get_vertices(env_ids=[env_id], scale=True) + if isinstance(value, (tuple, list)): + value = value[0] + vertices = torch.as_tensor(value, dtype=torch.float32, device=env.device) + if vertices.ndim == 3 and vertices.shape[0] == 1: + vertices = vertices[0] + if vertices.ndim != 2 or vertices.shape[-1] != 3 or vertices.numel() == 0: + raise ValueError(f"Rigid object {uid!r} has invalid mesh vertices.") + pose = _pose(env, uid)[env_id] + return vertices @ pose[:3, :3].transpose(0, 1) + pose[:3, 3] + + +def _projected_center_of_mass( + env: Any, + uid: str, + env_id: int, + world_vertices: torch.Tensor, +) -> torch.Tensor: + """Return the live COM projection, with a geometry-center fallback.""" + entity = env.sim.get_rigid_object(uid) + body_data = None if entity is None else getattr(entity, "body_data", None) + com_pose = None if body_data is None else getattr(body_data, "com_pose", None) + if callable(com_pose): + com_pose = com_pose() + if com_pose is not None: + local_com = torch.as_tensor( + com_pose, + dtype=torch.float32, + device=env.device, + ) + if local_com.ndim == 1: + local_com = local_com.unsqueeze(0).repeat(int(env.num_envs), 1) + if local_com.ndim == 2 and local_com.shape[0] == int(env.num_envs): + pose = _pose(env, uid)[env_id] + return (pose[:3, :3] @ local_com[env_id, :3] + pose[:3, 3])[:2] + return ( + world_vertices[:, :2].min(dim=0).values + + world_vertices[:, :2].max(dim=0).values + ) * 0.5 + + +def _object_supported_by( + env: Any, + spec: Mapping[str, Any], + defaults: Mapping[str, Any], +) -> torch.Tensor: + """Evaluate one-frame geometric support without advancing simulation.""" + object_uid = _object(spec) + support_uid = str( + spec.get( + "support", + spec.get("reference_object", spec.get("reference", "")), + ) + ) + if not support_uid: + raise ValueError("Support predicate requires a support object uid.") + margin = float(spec.get("com_margin", defaults["support_com_margin"])) + max_gap = float(spec.get("max_vertical_gap", defaults["support_max_vertical_gap"])) + max_penetration = float( + spec.get("max_penetration", defaults["support_max_penetration"]) + ) + min_overlap = float( + spec.get("min_overlap_ratio", defaults["support_min_overlap_ratio"]) + ) + result = _constant(env, False) + for env_id in range(int(env.num_envs)): + moved = _world_vertices(env, object_uid, env_id) + support = _world_vertices(env, support_uid, env_id) + moved_lower = moved[:, :2].min(dim=0).values + moved_upper = moved[:, :2].max(dim=0).values + support_lower = support[:, :2].min(dim=0).values + support_upper = support[:, :2].max(dim=0).values + overlap_extent = torch.clamp( + torch.minimum(moved_upper, support_upper) + - torch.maximum(moved_lower, support_lower), + min=0.0, + ) + moved_extent = torch.clamp(moved_upper - moved_lower, min=1e-6) + overlap_ratio = torch.prod(overlap_extent) / torch.prod(moved_extent) + projected_center = _projected_center_of_mass( + env, + object_uid, + env_id, + moved, + ) + center_supported = torch.all( + projected_center >= support_lower + margin + ) & torch.all(projected_center <= support_upper - margin) + local_mask = torch.all( + (support[:, :2] >= moved_lower - margin) + & (support[:, :2] <= moved_upper + margin), + dim=1, + ) + if bool(local_mask.any()): + local_support_height = support[local_mask, 2].max() + else: + # Sparse meshes may have no vertex exactly under a small payload. + # Nearest vertices are a local fallback; using the mesh-wide peak + # would confuse a remote protrusion with the candidate support pose. + distances = torch.linalg.vector_norm( + support[:, :2] - projected_center, + dim=1, + ) + count = min(8, int(support.shape[0])) + local_support_height = support[ + torch.topk(distances, count, largest=False).indices, 2 + ].max() + vertical_gap = moved[:, 2].min() - local_support_height + result[env_id] = bool( + center_supported + and overlap_ratio >= min_overlap + and vertical_gap >= -max_penetration + and vertical_gap <= max_gap + ) + return result + + +def _objects(spec: Mapping[str, Any]) -> list[str]: + values = spec.get("objects", spec.get("object_uids")) + if not isinstance(values, Sequence) or isinstance(values, (str, bytes)): + raise ValueError("Predicate requires a non-empty objects list.") + return [str(value) for value in values] + + +def _object(spec: Mapping[str, Any]) -> str: + value = spec.get("object", spec.get("object_uid")) + if not isinstance(value, str) or not value: + raise ValueError("Predicate requires a non-empty object uid.") + return value + + +def _local_axis_index(env: Any, uid: str, axis: Any) -> int: + name = str(axis).lower() + if name in {"x", "y", "z"}: + return {"x": 0, "y": 1, "z": 2}[name] + if name not in {"long", "long_axis", "longest"}: + raise ValueError(f"Unsupported upright local axis {axis!r}.") + entity = env.sim.get_rigid_object(uid) + if entity is None: + raise ValueError(f"Unknown rigid object {uid!r}.") + vertices = entity.get_vertices(env_ids=[0], scale=True) + if isinstance(vertices, (tuple, list)): + vertices = vertices[0] + vertices = torch.as_tensor(vertices, dtype=torch.float32, device=env.device) + if vertices.ndim == 3: + vertices = vertices[0] + if vertices.ndim != 2 or vertices.shape[-1] != 3 or vertices.numel() == 0: + raise ValueError(f"Rigid object {uid!r} has invalid mesh vertices.") + extents = vertices.max(dim=0).values - vertices.min(dim=0).values + return int(torch.argmax(extents).item()) + + +def _arm_values( + env: Any, kind: str +) -> tuple[torch.Tensor | None, torch.Tensor | None] | None: + getter = getattr(env, f"get_current_{kind}_agent", None) + if callable(getter): + left, right = getter() + values = [] + for value in (left, right): + if value is None: + values.append(None) + continue + item = torch.as_tensor(value, device=env.device) + if kind == "xpos" and item.ndim == 2: + item = item.unsqueeze(0).repeat(int(env.num_envs), 1, 1) + elif kind == "gripper_state" and item.ndim == 1: + item = item.unsqueeze(0) + values.append(item) + return values[0], values[1] + if kind != "gripper_state": + return None + qpos = env.robot.get_qpos() + values = [] + for side in ("left", "right"): + ids = list(getattr(env, f"{side}_eef_joints", ())) + if not ids: + return None + values.append(qpos[:, ids]) + return values[0], values[1] + + +def _gripper_has_closed( + env: Any, + gripper: torch.Tensor, + *, + tolerance: float, +) -> torch.Tensor: + """Check closure intent without requiring an impossible empty-gripper pose.""" + gripper = gripper.to(device=env.device, dtype=torch.float32) + open_state = getattr(env, "open_state", None) + close_state = getattr(env, "close_state", None) + reference = open_state if open_state is not None else close_state + if reference is None: + return _constant(env, False) + expected = torch.as_tensor( + reference, + dtype=torch.float32, + device=env.device, + ).flatten() + repeats = (gripper.shape[-1] + expected.numel() - 1) // expected.numel() + expected = expected.repeat(repeats)[: gripper.shape[-1]] + distance = torch.linalg.vector_norm(gripper - expected, dim=-1) + if open_state is not None: + return distance > tolerance + return distance <= tolerance + + +def _object_held( + env: Any, + uid: str, + *, + owners: Mapping[str, Sequence[str | None]] | None, + states: Mapping[tuple[str, str], Any] | None, + position_tolerance: float, + gripper_tolerance: float, + required_arm: str | None = None, +) -> torch.Tensor: + """Verify registry ownership against live object, TCP, and gripper state.""" + result = _constant(env, False) + if owners is None or states is None or uid not in owners: + return result + eef_values = _arm_values(env, "xpos") + gripper_values = _arm_values(env, "gripper_state") + if eef_values is None or gripper_values is None: + return result + + object_pose = _pose(env, uid) + for arm_index, arm in enumerate(("left_arm", "right_arm")): + if required_arm is not None and arm != required_arm: + continue + state = states.get((uid, arm)) + held = ( + None if state is None else state.get_held_object(arm_control_part(env, arm)) + ) + actual_eef = eef_values[arm_index] + gripper = gripper_values[arm_index] + if held is None or actual_eef is None or gripper is None: + continue + label = getattr(held.semantics, "label", None) + if not label and held.semantics.entity is not None: + label = getattr(held.semantics.entity, "uid", None) + if label != uid: + continue + actual_eef = actual_eef.to(device=env.device, dtype=object_pose.dtype) + expected_eef = torch.bmm( + object_pose, + held.object_to_eef.to(device=env.device, dtype=object_pose.dtype), + ) + position_ok = ( + torch.linalg.vector_norm( + actual_eef[:, :3, 3] - expected_eef[:, :3, 3], dim=-1 + ) + <= position_tolerance + ) + closed = _gripper_has_closed( + env, + gripper, + tolerance=gripper_tolerance, + ) + owned = torch.tensor( + [item == arm for item in owners[uid]], + dtype=torch.bool, + device=env.device, + ) + result |= owned & position_ok & closed + return result + + +def _coordinated_held( + env: Any, + uid: str, + state: Any, + *, + position_tolerance: float, + gripper_tolerance: float, +) -> torch.Tensor: + result = _constant(env, False) + if state is None: + return result + held_relations = tuple( + state.get_held_object(arm_control_part(env, arm)) + for arm in ("left_arm", "right_arm") + ) + if any(held is None for held in held_relations): + return result + for held in held_relations: + assert held is not None + label = getattr(held.semantics, "label", None) + if not label and getattr(held.semantics, "entity", None) is not None: + label = getattr(held.semantics.entity, "uid", None) + if label != uid: + return result + eef_values = _arm_values(env, "xpos") + gripper_values = _arm_values(env, "gripper_state") + if eef_values is None or gripper_values is None: + return result + + object_pose = _pose(env, uid) + result = _constant(env, True) + for arm_index, held in enumerate(held_relations): + assert held is not None + if held.env_mask is not None: + result &= held.env_mask.to(device=env.device) + actual_eef = eef_values[arm_index] + gripper = gripper_values[arm_index] + if actual_eef is None or gripper is None: + return _constant(env, False) + transform = held.object_to_eef.to( + device=env.device, + dtype=object_pose.dtype, + ) + expected_eef = torch.bmm(object_pose, transform) + actual_eef = actual_eef.to(device=env.device, dtype=object_pose.dtype) + position_ok = ( + torch.linalg.vector_norm( + actual_eef[:, :3, 3] - expected_eef[:, :3, 3], + dim=-1, + ) + <= position_tolerance + ) + closed = _gripper_has_closed( + env, + gripper, + tolerance=gripper_tolerance, + ) + result &= position_ok & closed + return result + + +def evaluate_predicate( + env: Any, + spec: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None, + *, + held_owners: Mapping[str, Sequence[str | None]] | None = None, + held_states: Mapping[tuple[str, str], Any] | None = None, + coordinated_state: Any | None = None, +) -> torch.Tensor: + """Evaluate one typed predicate or a boolean predicate tree.""" + runtime = { + "held_owners": held_owners, + "held_states": held_states, + "coordinated_state": coordinated_state, + } + defaults = _predicate_fallbacks(env) + if spec is None: + return _constant(env, True) + if isinstance(spec, Sequence) and not isinstance(spec, (str, bytes, Mapping)): + result = _constant(env, True) + for term in spec: + result &= evaluate_predicate(env, term, **runtime) + return result + if not isinstance(spec, Mapping): + raise TypeError("Predicate must be a mapping or a sequence of mappings.") + op = str(spec.get("op", "")).lower() + if not op and "terms" in spec: + op = "all" + if op in {"all", "and"}: + return evaluate_predicate(env, list(spec.get("terms", ())), **runtime) + if op in {"any", "or"}: + result = _constant(env, False) + for term in spec.get("terms", ()): + result |= evaluate_predicate(env, term, **runtime) + return result + if op == "not": + return ~evaluate_predicate(env, spec.get("term"), **runtime) + + kind = str(spec.get("type", spec.get("kind", ""))).lower() + if kind in {"semantic_goal", "line_member_placed", "stack_layer_supported"}: + raise ValueError( + f"Predicate {kind!r} is a compiler marker and requires the " + "executor's grounded target." + ) + if kind in {"object_held", "object_held_by_gripper"}: + required_arm = spec.get("arm") + if required_arm in {"left", "right"}: + required_arm = f"{required_arm}_arm" + return _object_held( + env, + _object(spec), + owners=held_owners, + states=held_states, + position_tolerance=float( + spec.get("position_tolerance", defaults["held_position_tolerance"]) + ), + gripper_tolerance=float( + spec.get("gripper_tolerance", defaults["held_gripper_tolerance"]) + ), + required_arm=str(required_arm) if required_arm else None, + ) + if kind == "handover_complete": + required_arm = spec.get("arm", "right_arm") + return _object_held( + env, + _object(spec), + owners=held_owners, + states=held_states, + position_tolerance=float( + spec.get("position_tolerance", defaults["held_position_tolerance"]) + ), + gripper_tolerance=float( + spec.get("gripper_tolerance", defaults["held_gripper_tolerance"]) + ), + required_arm=str(required_arm), + ) + if kind in {"held_by_both_grippers", "object_held_by_both_grippers"}: + return _coordinated_held( + env, + _object(spec), + coordinated_state, + position_tolerance=float( + spec.get("position_tolerance", defaults["held_position_tolerance"]) + ), + gripper_tolerance=float( + spec.get("gripper_tolerance", defaults["held_gripper_tolerance"]) + ), + ) + if kind in {"object_position_near", "position_near"}: + position = _position(env, _object(spec)) + target = torch.as_tensor( + spec.get("target_position", spec.get("target")), + dtype=position.dtype, + device=position.device, + ) + if target.ndim == 1: + target = target.unsqueeze(0) + return torch.linalg.vector_norm(position - target, dim=-1) <= float( + spec.get("tolerance", defaults["position_tolerance"]) + ) + if kind in {"object_xy_near", "xy_near"}: + position = _position(env, _object(spec))[:, :2] + target = torch.as_tensor( + spec.get("target_xy", spec.get("target")), + dtype=position.dtype, + device=position.device, + ).reshape(-1, 2) + return torch.linalg.vector_norm(position - target, dim=-1) <= float( + spec.get("tolerance", defaults["xy_tolerance"]) + ) + if kind in {"object_relative_position", "relative_position"}: + reference_uid = spec.get("reference_object", spec.get("reference")) + if not isinstance(reference_uid, str) or not reference_uid: + raise ValueError("Relative-position predicate requires a reference object.") + relation = str(spec.get("relation", "")) + axes = relation_axes( + env, + relation, + frame=str(spec.get("relation_frame", "world")), + ) + if not axes: + raise ValueError(f"Unsupported directional relation {relation!r}.") + delta = ( + _position(env, _object(spec))[:, :2] - _position(env, reference_uid)[:, :2] + ) + minimum_distance = float(spec.get("minimum_distance", 0.0)) + result = _constant(env, True) + for axis in axes: + projection = torch.sum( + delta * axis.to(dtype=delta.dtype, device=delta.device), dim=1 + ) + result &= projection >= minimum_distance + return result + if kind in {"object_in_container", "inside"}: + position = _position(env, _object(spec)) + container = _position( + env, str(spec.get("container", spec.get("reference_object"))) + ) + xy = torch.linalg.vector_norm(position[:, :2] - container[:, :2], dim=-1) + z = position[:, 2] - container[:, 2] + return ( + (xy <= float(spec.get("xy_radius", defaults["container_xy_radius"]))) + & (z >= float(spec.get("min_z_offset", defaults["container_min_z_offset"]))) + & (z <= float(spec.get("max_z_offset", defaults["container_max_z_offset"]))) + ) + if kind in {"object_supported_by", "object_on_object", "on"}: + return _object_supported_by(env, spec, defaults) + if kind == "object_not_fallen": + axis = _pose(env, _object(spec))[:, :3, 2] + cosine = axis[:, 2].clamp(-1.0, 1.0) + return torch.arccos(cosine) <= float( + spec.get("max_tilt", defaults["not_fallen_max_tilt"]) + ) + if kind == "object_upright": + uid = _object(spec) + local_axis = spec.get("local_axis", "long_axis") + axis_index = _local_axis_index( + env, + uid, + local_axis, + ) + axis = _pose(env, uid)[:, :3, axis_index] + cosine = axis[:, 2].clamp(-1.0, 1.0) + directed = spec.get( + "directed", + str(local_axis).lower() not in {"long", "long_axis", "longest"}, + ) + if not isinstance(directed, bool): + raise ValueError("object_upright directed must be a boolean.") + if not directed: + cosine = cosine.abs() + return torch.arccos(cosine) <= float( + spec.get("max_tilt", defaults["upright_max_tilt"]) + ) + if kind in {"object_axis_offset_near", "object_axis_near"}: + object_position = _position(env, _object(spec)) + axis = _axis_index(spec.get("axis", "x")) + reference_uid = spec.get( + "reference_object", + spec.get("reference", spec.get("support")), + ) + if isinstance(reference_uid, str) and reference_uid: + values = object_position[:, axis] - _position(env, reference_uid)[:, axis] + else: + values = object_position[:, axis] + target = spec.get( + "target_offset", + spec.get("offset", spec.get("target", 0.0)), + ) + target_value = torch.as_tensor( + target, + dtype=values.dtype, + device=values.device, + ) + return torch.abs(values - target_value) <= float( + spec.get("tolerance", defaults["axis_tolerance"]) + ) + if kind in {"objects_collinear", "collinear"}: + positions = torch.stack( + [_position(env, uid) for uid in _objects(spec)], + dim=1, + ) + axis = 0 if str(spec.get("axis", "x")) in {"x", "world_x"} else 1 + values = positions[:, :, 1 - axis] + return values.max(dim=1).values - values.min(dim=1).values <= float( + spec.get("tolerance", defaults["collinearity_tolerance"]) + ) + if kind in {"objects_ordered", "ordered"}: + positions = torch.stack( + [_position(env, uid) for uid in _objects(spec)], + dim=1, + ) + axis = 0 if str(spec.get("axis", "x")) in {"x", "world_x"} else 1 + differences = torch.diff(positions[:, :, axis], dim=1) + tolerance = float(spec.get("tolerance", defaults["ordering_tolerance"])) + if str(spec.get("direction", "ascending")) == "descending": + return torch.all(differences <= tolerance, dim=1) + return torch.all(differences >= -tolerance, dim=1) + if kind == "object_lifted": + position = _position(env, _object(spec))[:, 2] + initial = spec.get("initial_height") + if initial is None: + initial_pose = getattr(env, "agent_initial_object_poses", {}).get( + _object(spec) + ) + if initial_pose is None: + raise ValueError("object_lifted requires an initial object pose.") + initial = initial_pose[:, 2, 3] + initial = torch.as_tensor(initial, device=position.device) + return position >= initial + float( + spec.get("min_height", defaults["minimum_lift_height"]) + ) + if kind in {"both_arms_at_initial_qpos", "arms_home"}: + current = env.robot.get_qpos() + initial = getattr(env, "init_qpos", current) + return torch.all( + torch.abs(current - initial) + <= float(spec.get("tolerance", defaults["arm_initial_qpos_tolerance"])), + dim=-1, + ) + if kind in {"both_grippers_open", "grippers_open"}: + if not hasattr(env, "get_current_gripper_state_agent"): + return _constant(env, False) + left, right = env.get_current_gripper_state_agent() + expected = torch.as_tensor( + env.open_state, + dtype=torch.float32, + device=env.device, + ) + results = [] + for value in (left, right): + value = torch.as_tensor(value, dtype=torch.float32, device=env.device) + if value.ndim == 1: + value = value.unsqueeze(0).repeat(int(env.num_envs), 1) + results.append( + torch.linalg.vector_norm(value - expected, dim=-1) + <= float(spec.get("tolerance", defaults["gripper_state_tolerance"])) + ) + return results[0] & results[1] + if kind == "grippers_clear_of_object": + eef_values = _arm_values(env, "xpos") + if eef_values is None: + return _constant(env, False) + object_position = _position(env, _object(spec)) + clearance = float( + spec.get( + "min_distance", + spec.get("clearance", defaults["gripper_clear_min_distance"]), + ) + ) + result = _constant(env, True) + for eef in eef_values: + if eef is None: + return _constant(env, False) + result &= ( + torch.linalg.vector_norm( + eef[:, :3, 3] - object_position, + dim=-1, + ) + >= clearance + ) + return result + if kind == "pressed": + checker = getattr(env, "is_object_pressed", None) + if callable(checker): + value = checker(_object(spec), spec.get("terminal_state", "activated")) + result = torch.as_tensor(value, dtype=torch.bool, device=env.device) + return ( + result.repeat(int(env.num_envs)) + if result.ndim == 0 + else result.reshape(-1) + ) + states = getattr(env, "action_engine_semantic_states", {}) + value = states.get((_object(spec), "pressed")) + if value is None: + return _constant(env, False) + result = torch.as_tensor(value, dtype=torch.bool, device=env.device) + return ( + result.repeat(int(env.num_envs)) if result.ndim == 0 else result.reshape(-1) + ) + if kind == "coordinated_placed": + relation = str(spec.get("relation", "on")) + reference = spec.get("support_object", spec.get("reference_object")) + translated = { + "type": ( + "object_in_container" if relation == "inside" else "object_supported_by" + ), + "object": _object(spec), + ("container" if relation == "inside" else "support"): reference, + } + return evaluate_predicate(env, translated, **runtime) + raise ValueError(f"Unsupported execution predicate {kind!r}.") + + +def _axis_index(value: Any) -> int: + axis = str(value).lower().replace("world_", "") + if axis not in {"x", "y", "z"}: + raise ValueError(f"Unsupported predicate axis {value!r}.") + return {"x": 0, "y": 1, "z": 2}[axis] diff --git a/embodichain/gen_sim/action_engine/runtime/robot_parts.py b/embodichain/gen_sim/action_engine/runtime/robot_parts.py new file mode 100644 index 000000000..fe6f71c58 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/robot_parts.py @@ -0,0 +1,34 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Resolve semantic Action Engine arms to physical robot control parts.""" + +from __future__ import annotations + +from typing import Any + +__all__ = ["arm_control_part"] + + +def arm_control_part(env: Any, arm: str) -> str: + """Return the physical arm control part for a semantic arm name.""" + if arm not in {"left_arm", "right_arm"}: + raise ValueError(f"Expected a semantic arm, got {arm!r}.") + if hasattr(env, "get_agent_arm_control_part"): + part = env.get_agent_arm_control_part(arm == "left_arm") + if part: + return str(part) + return arm diff --git a/embodichain/gen_sim/action_engine/runtime/solver_compat.py b/embodichain/gen_sim/action_engine/runtime/solver_compat.py new file mode 100644 index 000000000..810bc8cf7 --- /dev/null +++ b/embodichain/gen_sim/action_engine/runtime/solver_compat.py @@ -0,0 +1,234 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Install solver compatibility corrections scoped to Action Engine.""" + +from __future__ import annotations + +from collections.abc import Mapping +import functools +import threading +from typing import Any + +import numpy as np +import torch + +from embodichain.lab.sim.solvers import PytorchSolver, URSolver, URSolverCfg + +__all__ = [ + "install_action_engine_solver_compat", + "install_pytorch_solver_tcp_compat", + "install_ur5_solver_frame_compat", + "repair_action_engine_ur5_solver_cfg", +] + +_PYTORCH_INSTALL_MARKER = "_action_engine_tcp_inverse_compat_installed" +_UR5_INSTALL_MARKER = "_action_engine_ur5_frame_compat_installed" +_UR5_ANALYTIC_TO_URDF_EE = np.eye(4, dtype=np.float32) +_UR5_ANALYTIC_TO_URDF_EE[0, 3] = -0.01 +_UR_DH_FIELDS = ("d1", "a2", "a3", "d4", "d5", "d6") + + +def repair_action_engine_ur5_solver_cfg(robot_cfg: Any) -> int: + """Repair stale UR10 DH defaults before Action Engine creates a UR5 robot. + + ``SolverCfg.from_dict`` constructs a UR10 config before assigning a + non-default ``ur_type``. Generated UR5 Action Engine configs therefore + reach the environment with UR10 DH values. Repair only that exact stale + signature so explicitly calibrated parameters remain untouched. + + Args: + robot_cfg: Robot configuration whose solver configs will be inspected. + + Returns: + Number of unique solver configs repaired by this call. + """ + configured = getattr(robot_cfg, "solver_cfg", None) + candidates = ( + configured.values() if isinstance(configured, Mapping) else (configured,) + ) + stale_defaults = URSolverCfg() + stale_dh = tuple(float(getattr(stale_defaults, name)) for name in _UR_DH_FIELDS) + + repaired = 0 + visited: set[int] = set() + for solver_cfg in candidates: + cfg_id = id(solver_cfg) + if cfg_id in visited: + continue + visited.add(cfg_id) + if not isinstance(solver_cfg, URSolverCfg): + continue + ur_type = str(getattr(solver_cfg, "ur_type", "")) + if ur_type != "ur5": + continue + current_dh = tuple(float(getattr(solver_cfg, name)) for name in _UR_DH_FIELDS) + if not np.allclose(current_dh, stale_dh, rtol=0.0, atol=1.0e-12): + continue + canonical = URSolverCfg(ur_type=ur_type) + for name in _UR_DH_FIELDS: + setattr(solver_cfg, name, getattr(canonical, name)) + repaired += 1 + return repaired + + +def install_action_engine_solver_compat(robot: Any) -> int: + """Install all solver corrections required by the Action Engine runtime.""" + return install_pytorch_solver_tcp_compat(robot) + install_ur5_solver_frame_compat( + robot + ) + + +def install_pytorch_solver_tcp_compat(robot: Any) -> int: + """Correct TCP inversion on every PytorchSolver owned by ``robot``. + + The shared solver currently transposes a rotation into an overlapping + tensor view. This wrapper transforms the requested TCP pose with a proper + matrix inverse, temporarily presents an identity TCP to the original + implementation, and otherwise preserves its sampling and ranking behavior. + + Args: + robot: Initialized robot containing its private solver registry. + + Returns: + Number of solver instances wrapped by this call. + """ + solvers = getattr(robot, "_solvers", None) + if not isinstance(solvers, Mapping): + return 0 + + installed = 0 + visited: set[int] = set() + for solver in solvers.values(): + solver_id = id(solver) + if solver_id in visited: + continue + visited.add(solver_id) + if not isinstance(solver, PytorchSolver) or bool( + getattr(solver, _PYTORCH_INSTALL_MARKER, False) + ): + continue + _wrap_solver(solver) + installed += 1 + return installed + + +def _wrap_solver(solver: PytorchSolver) -> None: + original_get_ik = solver.get_ik + call_lock = threading.RLock() + + @functools.wraps(original_get_ik) + def corrected_get_ik( + target_xpos: torch.Tensor | np.ndarray, + *args: Any, + **kwargs: Any, + ) -> Any: + target = torch.as_tensor( + target_xpos, + dtype=torch.float32, + device=solver.device, + ) + tcp = torch.as_tensor( + solver.tcp_xpos, + dtype=torch.float32, + device=solver.device, + ) + link_target = target @ torch.linalg.inv(tcp) + + # The solver instance is shared by vectorized environments. Protect the + # temporary TCP substitution in case a caller plans from another thread. + with call_lock: + active_tcp = solver.tcp_xpos + solver.tcp_xpos = np.eye(4, dtype=np.float32) + try: + return original_get_ik( + target_xpos=link_target, + *args, + **kwargs, + ) + finally: + solver.tcp_xpos = active_tcp + + solver.get_ik = corrected_get_ik + setattr(solver, _PYTORCH_INSTALL_MARKER, True) + + +def install_ur5_solver_frame_compat(robot: Any) -> int: + """Align UR5 analytic IK targets with the URDF ``ee_link`` frame. + + The UR5 asset carries a fixed ``-0.01 m`` local-x offset on ``ee_link`` + that is absent from the analytic DH model. The correction is installed + only for UR5 solvers owned by an Action Engine environment. + + Args: + robot: Initialized robot containing its private solver registry. + + Returns: + Number of solver instances wrapped by this call. + """ + solvers = getattr(robot, "_solvers", None) + if not isinstance(solvers, Mapping): + return 0 + + installed = 0 + visited: set[int] = set() + for solver in solvers.values(): + solver_id = id(solver) + if solver_id in visited: + continue + visited.add(solver_id) + if ( + not isinstance(solver, URSolver) + or str(getattr(getattr(solver, "cfg", None), "ur_type", "")) != "ur5" + or bool(getattr(solver, _UR5_INSTALL_MARKER, False)) + ): + continue + _wrap_ur5_solver(solver) + installed += 1 + return installed + + +def _wrap_ur5_solver(solver: URSolver) -> None: + original_get_ik = solver.get_ik + + @functools.wraps(original_get_ik) + def corrected_get_ik( + target_xpos: torch.Tensor | np.ndarray, + *args: Any, + **kwargs: Any, + ) -> Any: + target = torch.as_tensor( + target_xpos, + dtype=torch.float32, + device=solver.device, + ) + tcp = torch.as_tensor( + solver.tcp_xpos, + dtype=torch.float32, + device=solver.device, + ) + analytic_to_urdf = torch.as_tensor( + _UR5_ANALYTIC_TO_URDF_EE, + dtype=torch.float32, + device=solver.device, + ) + corrected_target = ( + target @ torch.linalg.inv(tcp) @ torch.linalg.inv(analytic_to_urdf) @ tcp + ) + return original_get_ik(corrected_target, *args, **kwargs) + + solver.get_ik = corrected_get_ik + setattr(solver, _UR5_INSTALL_MARKER, True) diff --git a/tests/gen_sim/action_engine/capabilities/__init__.py b/tests/gen_sim/action_engine/capabilities/__init__.py new file mode 100644 index 000000000..046cb429b --- /dev/null +++ b/tests/gen_sim/action_engine/capabilities/__init__.py @@ -0,0 +1,19 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +"""Action Engine capability tests.""" diff --git a/tests/gen_sim/action_engine/capabilities/test_atomic_v2.py b/tests/gen_sim/action_engine/capabilities/test_atomic_v2.py new file mode 100644 index 000000000..07245b274 --- /dev/null +++ b/tests/gen_sim/action_engine/capabilities/test_atomic_v2.py @@ -0,0 +1,233 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +from copy import deepcopy +from dataclasses import dataclass +from types import SimpleNamespace + +import torch + +from embodichain.gen_sim.action_engine.capabilities import ( + AtomicCapability, + build_atomic_capability_registry, +) +from embodichain.gen_sim.action_engine.runtime.actions import AtomicActionAdapter +from embodichain.gen_sim.action_engine.runtime.grounding import ActionGrounder +from embodichain.gen_sim.action_engine.runtime.loader import load_execution_program +from embodichain.gen_sim.action_engine.runtime.models import GroundedAction +from embodichain.gen_sim.action_engine.runtime.state import ExecutionState +from embodichain.gen_sim.action_engine.tasks import instantiate_seed_graph + +from ..task_fixtures import make_task_spec +from embodichain.gen_sim.action_engine.planning.linker import link_seed_graph +from embodichain.lab.sim.atomic_actions import ( + ActionBinding, + ActionOptions, + ActionPlan, + EndEffectorPoseGoal, + PlannerDiagnostics, + RuntimeCommandFrame, + TimedTrajectory, + TimedCommandSequence, +) + + +@dataclass(frozen=True, slots=True) +class _TestOptions(ActionOptions): + marker: str = "test" + + +class _TestAction: + skill_id = "test_retreat" + end_effector_roles: tuple[str, ...] = () + + +class _TestEngine: + binding_owner_id = "test-engine" + + def bind_control_parts(self, _skill_id, _endpoints): + return ActionBinding(owner_id=self.binding_owner_id) + + def plan(self, invocation, context): + assert isinstance(invocation.skill_options, _TestOptions) + positions = context.robot.qpos[:, None, :] + trajectory = TimedTrajectory.from_uniform_step( + positions, + env_ids=context.env_ids, + step_dt=context.require_control_dt(), + ) + return ActionPlan( + skill_id=invocation.skill_id, + plan_success=torch.ones(context.batch_size, dtype=torch.bool), + commands=TimedCommandSequence( + frames=( + RuntimeCommandFrame( + commands=(), + active_mask=torch.ones( + context.batch_size, + dtype=torch.bool, + ), + env_ids=context.env_ids, + hold_duration=trajectory.dt[:, 0], + ), + ), + env_ids=context.env_ids, + ), + joint_trajectory=trajectory, + recovery_policy=invocation.recovery_policy, + planned_scene_version=context.scene.version, + planned_collision_world_revision=(0,) * context.batch_size, + diagnostics=PlannerDiagnostics(backend="test"), + ) + + +class _Robot: + dof = 2 + uid = "test_robot" + control_parts = {"left_arm": [0], "right_arm": [1]} + + def get_qpos(self): + return torch.zeros((1, 2)) + + def get_joint_ids(self, *, name: str): + return self.control_parts.get(name, []) + + +class _Entity: + def get_local_pose(self, *, to_matrix: bool): + assert to_matrix + return torch.eye(4).unsqueeze(0) + + +class _Sim: + def get_rigid_object(self, _uid: str): + return _Entity() + + +def test_new_descriptor_reuses_loader_and_adapter_without_dispatch_changes() -> None: + registry = build_atomic_capability_registry() + calls = [] + + def target_hook(**kwargs): + calls.append("target") + pose = kwargs["object_pose"].clone() + return GroundedAction( + action_class="TestRetreat", + arm=kwargs["arm"], + control="arm", + target=EndEffectorPoseGoal(xpos=pose), + cfg=kwargs["policy"], + object_pose=pose, + target_object_pose=pose, + motion_policy=kwargs["policy"], + ) + + def config_hook(**_kwargs): + calls.append("config") + return _TestOptions() + + registry.register( + AtomicCapability( + "TestRetreat", + _TestAction, + _TestOptions, + frozenset({"policy_pose"}), + frozenset({"arm"}), + "single_arm", + "preserve", + "eef_pose", + motion_base="MoveEndEffector", + target_materializer_hook=target_hook, + config_materializer_hook=config_hook, + contract_resolver_hook=registry.get( + "MoveEndEffector" + ).contract_resolver_hook, + ) + ) + task, requirements = make_task_spec("E1") + bindings = { + item["role_id"]: f"uid_{item['role_id']}" for item in requirements["objects"] + } + graph = instantiate_seed_graph(task, bindings) + graph = deepcopy(graph) + graph["capability_catalog_hash"] = registry.catalog_hash() + cleanup = next( + node for node in graph["nodes"] if node["atomic_action"] == "MoveEndEffector" + ) + cleanup["atomic_action"] = "TestRetreat" + cleanup.pop("contract") + for group in graph["task_groups"]: + group.pop("contract") + graph["metadata"].pop("action_contract_linker") + graph = link_seed_graph(graph, registry=registry) + + program = load_execution_program(graph, registry=registry) + assert any( + action["atomic_action_class"] == "TestRetreat" + for edge in program.edges + for action in edge.actions + ) + + env = SimpleNamespace( + num_envs=1, + device=torch.device("cpu"), + robot=_Robot(), + sim=_Sim(), + agent_robot_profile="dual_ur10", + get_agent_arm_control_part=lambda is_left: ( + "left_arm" if is_left else "right_arm" + ), + get_agent_eef_control_part=lambda _is_left: None, + ) + adapter = AtomicActionAdapter( + env, + grasp_policy={}, + capability_registry=registry, + ) + adapter._atomic_engine = _TestEngine() + step = next( + step + for step in program.semantic_steps + if any( + action["atomic_action_class"] == "TestRetreat" + for edge_id in step.edge_ids + for action in next( + edge for edge in program.edges if edge.id == edge_id + ).actions + ) + ) + action = next( + action + for edge in program.edges + for action in edge.actions + if action["atomic_action_class"] == "TestRetreat" + ) + grounder = ActionGrounder( + program, + env, + lambda _uid: None, + capability_registry=registry, + ) + state = ExecutionState(last_qpos=torch.zeros((1, 2))) + grounded = grounder.ground(action, step, arm="left_arm", state=state) + outcome = adapter.plan( + grounded, + state, + ) + assert outcome.success.tolist() == [True] + assert calls == ["target", "config"] diff --git a/tests/gen_sim/action_engine/planning/test_linker.py b/tests/gen_sim/action_engine/planning/test_linker.py new file mode 100644 index 000000000..104438d5b --- /dev/null +++ b/tests/gen_sim/action_engine/planning/test_linker.py @@ -0,0 +1,407 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +from copy import deepcopy + +import pytest + +from embodichain.gen_sim.action_engine.domain import seed_graph_hash +from embodichain.gen_sim.action_engine.planning.linker import ( + link_seed_graph, + link_task_dependencies, +) +from embodichain.gen_sim.action_engine.protocol import ( + SEED_GRAPH_SCHEMA, + TASK_SPEC_SCHEMA, +) +from embodichain.gen_sim.action_engine.runtime.loader import load_execution_program +from embodichain.gen_sim.action_engine.tasks.recipes import instantiate_seed_graph + + +def _handover_task() -> dict: + return { + "schema_version": TASK_SPEC_SCHEMA, + "task_id": "handover_then_place", + "level": "L3", + "instruction": "Stand both cans, hand over the purple can, then place it.", + "reasoning_type": "none", + "task_instances": [ + { + "id": "task_01", + "task_type": "E2", + "params": { + "object_role": "purple", + "required_arm": "right_arm", + }, + "depends_on": [], + "role": "primary", + }, + { + "id": "task_02", + "task_type": "E2", + "params": { + "object_role": "orange", + "required_arm": "left_arm", + }, + "depends_on": [], + "role": "primary", + }, + { + "id": "task_03", + "task_type": "E4", + "params": { + "object_role": "purple", + "transfer_arm": "right_arm", + "receive_arm": "left_arm", + }, + "depends_on": ["task_02"], + "role": "primary", + }, + { + "id": "task_04", + "task_type": "E1", + "params": { + "object_role": "purple", + "target_role": "orange", + "relation": "left_of", + "required_arm": "left_arm", + }, + "depends_on": ["task_03"], + "role": "primary", + }, + ], + "success": {"type": "all_complete"}, + "oracle": {}, + "metadata": {}, + } + + +def _handover_graph() -> dict: + return instantiate_seed_graph( + _handover_task(), + {"purple": "purple_can", "orange": "orange_can"}, + ) + + +def _unlink_for_rebuild(graph: dict) -> None: + graph["metadata"].pop("action_contract_linker", None) + for group in graph["task_groups"]: + group.pop("contract", None) + + +def test_task_linker_preserves_parallel_arms_and_waits_for_both_before_handover() -> ( + None +): + linked = link_task_dependencies( + _handover_task(), + {"purple": "purple_can", "orange": "orange_can"}, + ) + by_id = {item["id"]: item for item in linked["task_instances"]} + + assert by_id["task_01"]["depends_on"] == [] + assert by_id["task_02"]["depends_on"] == [] + assert by_id["task_03"]["depends_on"] == ["task_02", "task_01"] + + +def test_resource_dependency_provenance_is_persisted_in_seed_graph() -> None: + task = _handover_task() + task["task_instances"] = task["task_instances"][:3] + handover = task["task_instances"][2] + handover["params"] = { + "object_role": "orange", + "transfer_arm": "left_arm", + "receive_arm": "right_arm", + } + + graph = instantiate_seed_graph( + task, + {"purple": "purple_can", "orange": "orange_can"}, + ) + + provenance = graph["metadata"]["action_contract_task_linker"] + assert provenance["linked_dependencies"] == [ + { + "from": "task_01", + "to": "task_03", + "reason": "resource", + "detail": "arm:right_arm", + } + ] + + +def test_same_object_e2_handover_gets_direct_causal_edge_through_a_chain() -> None: + task = _handover_task() + task["task_instances"][1]["depends_on"] = ["task_01"] + linked = link_task_dependencies( + task, + {"purple": "purple_can", "orange": "orange_can"}, + ) + handover = next( + item for item in linked["task_instances"] if item["id"] == "task_03" + ) + + assert handover["depends_on"] == ["task_02", "task_01"] + + +def test_handover_ownership_flows_through_home_terminal_barrier() -> None: + graph = _handover_graph() + groups = {group["id"]: group for group in graph["task_groups"]} + nodes = {node["id"]: node for node in graph["nodes"]} + handover_group = groups["task_03"] + terminal_id = handover_group["contract"]["terminal_node_ids"][0] + terminal = nodes[terminal_id] + receiver_entry = nodes[groups["task_04"]["contract"]["entry_node_ids"][0]] + handover = next( + node + for node in graph["nodes"] + if node["task_instance_id"] == "task_03" and node["atomic_action"] == "HandOver" + ) + + assert terminal["atomic_action"] == "MoveJoints" + assert terminal["contract"]["completion"] == "terminal_barrier" + assert terminal["contract"]["failure_policy"] == "best_effort" + retreat = next( + node + for node in graph["nodes"] + if node["task_instance_id"] == "task_03" + and node["atomic_action"] == "MoveEndEffector" + ) + assert retreat["contract"]["failure_policy"] == "safety_required" + assert terminal_id in receiver_entry["depends_on"] + assert { + (effect["op"], effect["atom"]["predicate"], effect["atom"].get("arm")) + for effect in handover["contract"]["effects"] + } >= { + ("delete", "object_held", "right_arm"), + ("add", "object_held", "left_arm"), + } + + +def test_linker_is_idempotent_and_hash_stable() -> None: + graph = _handover_graph() + relinked = link_seed_graph( + graph, + task_order=["task_01", "task_02", "task_03", "task_04"], + known_objects={"purple_can", "orange_can", "table"}, + ) + + assert relinked == graph + assert seed_graph_hash(relinked) == seed_graph_hash(graph) + + +def test_linker_rejects_missing_cleanup_wrong_holder_and_duplicate_pickup() -> None: + missing_cleanup = deepcopy(_handover_graph()) + home = next( + node + for node in missing_cleanup["nodes"] + if node["task_instance_id"] == "task_03" + and node["atomic_action"] == "MoveJoints" + ) + missing_cleanup["nodes"].remove(home) + next(group for group in missing_cleanup["task_groups"] if group["id"] == "task_03")[ + "node_ids" + ].remove(home["id"]) + for node in missing_cleanup["nodes"]: + node["depends_on"] = [ + dependency for dependency in node["depends_on"] if dependency != home["id"] + ] + _unlink_for_rebuild(missing_cleanup) + with pytest.raises(ValueError, match="terminal barrier"): + link_seed_graph(missing_cleanup) + + wrong_holder = deepcopy(_handover_graph()) + staging = next( + node + for node in wrong_holder["nodes"] + if node["task_instance_id"] == "task_03" + and node["atomic_action"] == "MoveHeldObject" + ) + staging["actor"] = {"mode": "required", "arm": "left_arm"} + staging.pop("contract") + _unlink_for_rebuild(wrong_holder) + with pytest.raises(ValueError, match="no producer|unavailable state"): + link_seed_graph(wrong_holder) + + duplicate_pickup = deepcopy(_handover_graph()) + pickup = next( + node + for node in duplicate_pickup["nodes"] + if node["task_instance_id"] == "task_01" and node["atomic_action"] == "PickUp" + ) + staging = next( + node + for node in duplicate_pickup["nodes"] + if node["task_instance_id"] == "task_01" + and node["atomic_action"] == "MoveHeldObject" + ) + repeated = deepcopy(pickup) + repeated["id"] = "task_01__duplicate_pickup" + repeated["depends_on"] = [pickup["id"]] + repeated.pop("contract") + staging["depends_on"] = [repeated["id"]] + group = next( + group for group in duplicate_pickup["task_groups"] if group["id"] == "task_01" + ) + pickup_index = group["node_ids"].index(pickup["id"]) + group["node_ids"].insert(pickup_index + 1, repeated["id"]) + duplicate_pickup["nodes"].insert( + duplicate_pickup["nodes"].index(pickup) + 1, repeated + ) + _unlink_for_rebuild(duplicate_pickup) + with pytest.raises(ValueError, match="requires unavailable state"): + link_seed_graph(duplicate_pickup) + + +def test_unavailable_arm_reports_current_holder_and_requested_object() -> None: + task = _handover_task() + placement = task["task_instances"][3] + placement["params"].update( + { + "object_role": "orange", + "target_role": "purple", + "required_arm": "left_arm", + } + ) + + with pytest.raises( + ValueError, + match=( + "left_arm.*currently holds 'purple_can'.*" "primary object is 'orange_can'" + ), + ): + instantiate_seed_graph( + task, + {"purple": "purple_can", "orange": "orange_can"}, + ) + + +def test_readers_remain_parallel_and_writer_waits_for_both() -> None: + task = { + "schema_version": TASK_SPEC_SCHEMA, + "task_id": "read_write", + "level": "L3", + "instruction": "Inspect a shared target, then manipulate it.", + "reasoning_type": "none", + "task_instances": [ + { + "id": "read_left", + "task_type": "E1", + "params": { + "object_role": "a", + "target_role": "target", + "required_arm": "left_arm", + }, + "depends_on": [], + "role": "primary", + }, + { + "id": "read_right", + "task_type": "E1", + "params": { + "object_role": "b", + "target_role": "target", + "required_arm": "right_arm", + }, + "depends_on": [], + "role": "primary", + }, + { + "id": "write_target", + "task_type": "E2", + "params": { + "object_role": "target", + "required_arm": "left_arm", + }, + "depends_on": [], + "role": "primary", + }, + ], + "success": {}, + "oracle": {}, + "metadata": {}, + } + linked = link_task_dependencies( + task, + {"a": "object_a", "b": "object_b", "target": "shared_target"}, + ) + by_id = {item["id"]: item for item in linked["task_instances"]} + + assert by_id["read_left"]["depends_on"] == [] + assert by_id["read_right"]["depends_on"] == [] + assert by_id["write_target"]["depends_on"] == ["read_left", "read_right"] + + +def test_explicit_distinct_arm_allocation_keeps_auto_groups_parallel() -> None: + task = { + "schema_version": TASK_SPEC_SCHEMA, + "task_id": "allocated_auto", + "level": "L2", + "instruction": "Stand both objects upright in parallel.", + "reasoning_type": "none", + "task_instances": [ + { + "id": "first", + "task_type": "E2", + "params": {"object_role": "first_object"}, + "depends_on": [], + "role": "primary", + }, + { + "id": "second", + "task_type": "E2", + "params": {"object_role": "second_object"}, + "depends_on": [], + "role": "primary", + }, + ], + "success": {}, + "oracle": {}, + "metadata": { + "allocation_groups": [ + { + "id": "distinct_pair", + "task_instance_ids": ["first", "second"], + "arm_constraint": "distinct_arms", + } + ] + }, + } + bindings = {"first_object": "first_uid", "second_object": "second_uid"} + linked = link_task_dependencies(task, bindings) + graph = instantiate_seed_graph(linked, bindings) + + assert all(not item["depends_on"] for item in linked["task_instances"]) + assert all(not group["depends_on"] for group in graph["task_groups"]) + + +def test_v2_and_resolver_mismatch_require_regeneration() -> None: + with pytest.raises( + ValueError, match="lacks persisted Action Contracts.*regenerate" + ): + load_execution_program({"schema_version": "action_engine_seed_graph_v2"}) + + graph = _handover_graph() + graph["nodes"][0]["contract"]["claims"][0]["access"] = "shared_read" + with pytest.raises( + ValueError, match="does not match the current capability resolver" + ): + load_execution_program(graph) + + +def test_seed_graph_schema_is_v3() -> None: + assert _handover_graph()["schema_version"] == SEED_GRAPH_SCHEMA diff --git a/tests/gen_sim/action_engine/runtime/__init__.py b/tests/gen_sim/action_engine/runtime/__init__.py new file mode 100644 index 000000000..66ef3ba1a --- /dev/null +++ b/tests/gen_sim/action_engine/runtime/__init__.py @@ -0,0 +1,19 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +"""Runtime contract tests for Action Engine.""" diff --git a/tests/gen_sim/action_engine/runtime/test_actions.py b/tests/gen_sim/action_engine/runtime/test_actions.py new file mode 100644 index 000000000..ba2d35951 --- /dev/null +++ b/tests/gen_sim/action_engine/runtime/test_actions.py @@ -0,0 +1,835 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +"""Focused contracts for the public atomic-action adapter.""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import pytest +import torch + +from embodichain.gen_sim.action_engine.runtime import actions +from embodichain.gen_sim.action_engine.runtime.actions import AtomicActionAdapter +from embodichain.gen_sim.action_engine.runtime.models import ( + ActionOutcome, + GroundedAction, +) +from embodichain.gen_sim.action_engine.runtime.state import ExecutionState +from embodichain.lab.sim.atomic_actions import ( + Affordance, + ActionBinding, + ActionPlan, + AntipodalAffordance, + CoordinatedPickGoal, + EndEffectorPoseGoal, + GraspGoal, + HeldObjectState, + JointPositionGoal, + ObjectSemantics, + PlannerDiagnostics, + RecoveryPolicy, + RuntimeCommandFrame, + SceneSnapshot, + StateDelta, + TimedCommandSequence, + TimedTrajectory, +) +from embodichain.lab.sim.planners import CuroboPlannerCfg +from embodichain.toolkits.graspkit.pg_grasp import GraspGeneratorCfg + + +class _MeshEntity: + def get_vertices(self, *, env_ids: list[int], scale: bool) -> torch.Tensor: + assert env_ids == [0] + assert scale + return torch.tensor( + [ + [0.0, 0.0, 0.0], + [1.0, 0.0, 0.0], + [0.0, 1.0, 0.0], + ], + dtype=torch.float32, + ) + + def get_triangles(self, *, env_ids: list[int]) -> torch.Tensor: + assert env_ids == [0] + return torch.tensor([[0, 1, 2]], dtype=torch.int64) + + +class _PoseEntity: + def __init__(self, pose: torch.Tensor) -> None: + self.pose = pose + + def get_local_pose(self, *, to_matrix: bool) -> torch.Tensor: + assert to_matrix + return self.pose.clone() + + +class _PlannerRobot: + uid = "test_robot" + dof = 8 + + _ids = { + "physical_left_arm": [0, 1], + "physical_left_eef": [2, 3], + "physical_right_arm": [4, 5], + "physical_right_eef": [6, 7], + } + control_parts = _ids + + def get_joint_ids(self, *, name: str) -> list[int]: + return list(self._ids[name]) + + def get_control_part_base_pose(self, *, name: str, to_matrix: bool) -> torch.Tensor: + assert to_matrix + pose = torch.eye(4).repeat(2, 1, 1) + pose[:, 1, 3] = 0.3 if name == "physical_left_arm" else -0.3 + return pose + + +def _commands_for(trajectory: TimedTrajectory) -> TimedCommandSequence: + """Build timing-only frames for retained test trajectories.""" + active = torch.ones( + trajectory.batch_size, + dtype=torch.bool, + device=trajectory.positions.device, + ) + frames = tuple( + RuntimeCommandFrame( + commands=(), + active_mask=active, + env_ids=trajectory.env_ids, + hold_duration=trajectory.dt[:, index], + ) + for index in range(trajectory.waypoint_count) + ) + return TimedCommandSequence(frames=frames, env_ids=trajectory.env_ids) + + +class _FakeEngine: + """Minimal endpoint-binding and planning surface for adapter unit tests.""" + + binding_owner_id = "action-engine-test" + + def __init__(self, plan=None) -> None: + self._plan = plan + + def bind_control_parts(self, _skill_id, _endpoints) -> ActionBinding: + return ActionBinding(owner_id=self.binding_owner_id) + + def plan(self, invocation, context) -> ActionPlan: + if self._plan is None: + raise AssertionError("This fake engine has no planning callback.") + return self._plan(invocation, context) + + +def _planner_env( + *, + table: Any | None = None, + rigid_objects: dict[str, Any] | None = None, +) -> SimpleNamespace: + entities = dict(rigid_objects or {}) + if table is not None: + entities["table"] = table + return SimpleNamespace( + num_envs=2, + device=torch.device("cpu"), + robot=_PlannerRobot(), + sim=SimpleNamespace(get_rigid_object=entities.get), + left_arm_joints=[0, 1], + left_eef_joints=[2, 3], + right_arm_joints=[4, 5], + right_eef_joints=[6, 7], + open_state=torch.zeros(2), + close_state=torch.ones(2), + get_agent_arm_control_part=lambda is_left: ( + "physical_left_arm" if is_left else "physical_right_arm" + ), + get_agent_eef_control_part=lambda is_left: ( + "physical_left_eef" if is_left else "physical_right_eef" + ), + ) + + +def test_semantics_prewarms_vhacd_cache_before_affordance( + monkeypatch: Any, +) -> None: + """The lazy shared checker must see V-HACD's pickle, never create CoACD.""" + events: list[str] = [] + observed: dict[str, Any] = {} + entity = _MeshEntity() + env = SimpleNamespace( + num_envs=1, + device=torch.device("cpu"), + sim=SimpleNamespace( + get_rigid_object=lambda uid: entity if uid == "cube" else None + ), + agent_grasp_runtime_defaults={"max_decomposition_hulls": 8}, + ) + + def fake_prepare(**kwargs: Any) -> SimpleNamespace: + events.append("cache") + observed.update(kwargs) + return SimpleNamespace(status="hit") + + def fake_affordance(**kwargs: Any) -> Affordance: + events.append("affordance") + observed["generator_cfg"] = kwargs["generator_cfg"] + observed["gripper_collision_cfg"] = kwargs["gripper_collision_cfg"] + return Affordance() + + monkeypatch.setattr( + actions, + "ensure_vhacd_grasp_collision_cache", + fake_prepare, + ) + monkeypatch.setattr(actions, "AntipodalAffordance", fake_affordance) + + adapter = AtomicActionAdapter(env) + first = adapter.semantics("cube") + second = adapter.semantics("cube") + + assert first is second + assert events == ["cache", "affordance"] + assert observed["max_decomposition_hulls"] == 8 + assert observed["mesh_vertices"].dtype == torch.float32 + assert observed["mesh_triangles"].dtype == torch.int64 + assert observed["generator_cfg"].n_deviated_approach_directions == 4 + assert observed["gripper_collision_cfg"] is not None + + +def test_planner_policy_uses_curobo_for_single_arm_and_ik_for_dual_arm() -> None: + adapter = AtomicActionAdapter(_planner_env()) + adapter._atomic_engine = _FakeEngine() + goal = JointPositionGoal(target=torch.zeros(2, 2)) + + single = adapter._invocation( + GroundedAction("MoveJoints", "left_arm", "arm", goal, {}), + adapter.capabilities.get("MoveJoints"), + ) + coordinated_goal = CoordinatedPickGoal( + semantics=ObjectSemantics( + label="tray", + geometry={}, + affordance=AntipodalAffordance(), + ), + object_target_pose=torch.eye(4), + object_initial_pose=torch.eye(4), + ) + coordinated = adapter._invocation( + GroundedAction( + "CoordinatedPickment", + "coordinated", + "arm", + coordinated_goal, + {}, + ), + adapter.capabilities.get("CoordinatedPickment"), + ) + hand = adapter._invocation( + GroundedAction("MoveJoints", "left_arm", "hand", goal, {}), + adapter.capabilities.get("MoveJoints"), + ) + + assert adapter.planner_policy["backend"] == "curobo" + assert single.motion_policy.strategy == "motion_gen" + assert coordinated.motion_policy.strategy == "ik_interp" + assert torch.allclose( + coordinated.skill_options.left_to_right_arm_direction, + torch.tensor([0.0, -1.0, 0.0]), + ) + assert hand.motion_policy.strategy == "ik_interp" + + +def test_coordinated_pickment_scopes_ground_filter_to_gensim_goal_copy() -> None: + adapter = AtomicActionAdapter(_planner_env()) + adapter._atomic_engine = _FakeEngine() + original_cfg = GraspGeneratorCfg(is_filter_ground_collision=True) + affordance = AntipodalAffordance(generator_cfg=original_cfg) + goal = CoordinatedPickGoal( + semantics=ObjectSemantics( + label="tray", + geometry={}, + affordance=affordance, + ), + object_target_pose=torch.eye(4), + object_initial_pose=torch.eye(4), + ) + grounded = GroundedAction( + "CoordinatedPickment", + "coordinated", + "arm", + goal, + { + "middle_empty_ratio": 0.7, + "is_filter_ground_collision": False, + }, + ) + + invocation = adapter._invocation( + grounded, + adapter.capabilities.get("CoordinatedPickment"), + ) + + scoped_affordance = invocation.goal.semantics.affordance + assert isinstance(scoped_affordance, AntipodalAffordance) + assert scoped_affordance is not affordance + assert affordance.generator_cfg is original_cfg + assert original_cfg.is_filter_ground_collision is True + assert scoped_affordance.generator_cfg is not original_cfg + assert scoped_affordance.generator_cfg.is_filter_ground_collision is False + assert invocation.skill_options.middle_empty_ratio == pytest.approx(0.7) + + +def test_retreat_uses_row_local_motion_planner_reachability_search( + monkeypatch: Any, +) -> None: + env = _planner_env() + adapter = AtomicActionAdapter(env) + reference = torch.eye(4).repeat(2, 1, 1) + reference[:, 2, 3] = 1.05 + requested = reference.clone() + requested[:, 2, 3] = 1.35 + height_thresholds = torch.tensor([1.24, 1.00]) + attempted_targets: list[torch.Tensor] = [] + + def plan(invocation: Any, _context: Any) -> ActionPlan: + target = invocation.goal.xpos.clone() + attempted_targets.append(target) + height_reachable = target[:, 2, 3] <= height_thresholds + baseward_reachable = target[:, 1, 3] < -0.05 + success = height_reachable | baseward_reachable + terminal = target[:, 2, 3, None].repeat(1, 8) + positions = torch.stack((torch.zeros_like(terminal), terminal), dim=1) + trajectory = TimedTrajectory.from_uniform_step( + positions, + env_ids=torch.arange(2), + step_dt=0.01, + ) + return ActionPlan( + skill_id="move_end_effector", + plan_success=success, + commands=_commands_for(trajectory), + joint_trajectory=trajectory, + recovery_policy=RecoveryPolicy(), + planned_scene_version=0, + planned_collision_world_revision=(0, 0), + diagnostics=PlannerDiagnostics(backend="fake"), + expected_effects=StateDelta(), + ) + + monkeypatch.setattr(adapter, "_engine", lambda: _FakeEngine(plan)) + grounded = GroundedAction( + "MoveEndEffector", + "right_arm", + "arm", + EndEffectorPoseGoal(xpos=requested), + { + "sample_interval": 10, + "retreat_height": 0.30, + "minimum_retreat_height": 0.05, + "retreat_distance": 0.10, + }, + motion_policy={ + "collision_safety": "required", + "retreat_reachability_search": True, + "retreat_reference_pose": reference, + "minimum_retreat_height": 0.05, + "retreat_distance": 0.10, + }, + ) + + outcome = adapter.plan( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + assert len(attempted_targets) > 1 + assert bool(outcome.success.all()) + selected_z = outcome.grounded.target.xpos[:, 2, 3] + assert selected_z.tolist() == pytest.approx([1.20, 1.35]) + assert outcome.grounded.target.xpos[:, 1, 3].tolist() == pytest.approx([0.0, -0.10]) + search = outcome.planner_trace["reachability_search"] + assert search["strategy"] == "bounded_motion_planner" + assert search["selected_target_z"].tolist() == pytest.approx([1.20, 1.35]) + assert len(search["attempts"]) == len(attempted_targets) + + +def test_curobo_generator_receives_generated_static_obstacles( + monkeypatch: Any, +) -> None: + table = object() + can = object() + captured: dict[str, Any] = {} + + def fake_motion_generator(*, cfg: Any) -> object: + captured["cfg"] = cfg + return object() + + monkeypatch.setattr(actions, "MotionGenerator", fake_motion_generator) + adapter = AtomicActionAdapter( + _planner_env(table=table, rigid_objects={"can": can}), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": ["can"], + }, + ) + + generator = adapter._generator() + + assert generator is adapter._motion_generator + planner = captured["cfg"].planner_cfg + assert isinstance(planner, CuroboPlannerCfg) + assert planner.world.rigid_objects == {"table": table, "can": can} + assert planner.world.dynamic_obstacle_names == ["can"] + assert planner.world.obstacle_representation == "cuboid" + assert planner.world.collision_cache == {"cuboid": 8, "mesh": 2} + + +def test_curobo_generator_sizes_collision_cache_for_large_scene( + monkeypatch: Any, +) -> None: + rigid_objects = {f"object_{index:02d}": object() for index in range(13)} + captured: dict[str, Any] = {} + + def fake_motion_generator(*, cfg: Any) -> object: + captured["cfg"] = cfg + return object() + + monkeypatch.setattr(actions, "MotionGenerator", fake_motion_generator) + adapter = AtomicActionAdapter( + _planner_env(rigid_objects=rigid_objects), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": list(rigid_objects), + }, + ) + + adapter._generator() + + planner = captured["cfg"].planner_cfg + assert planner.world.collision_cache == {"cuboid": 13, "mesh": 2} + + +def test_dynamic_scene_parks_contact_target_and_held_rows() -> None: + actual = torch.eye(4).repeat(2, 1, 1) + actual[:, 2, 3] = torch.tensor([0.7, 0.8]) + entities = {uid: _PoseEntity(actual.clone()) for uid in ("target", "held", "other")} + adapter = AtomicActionAdapter( + _planner_env(rigid_objects=entities), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": list(entities), + }, + ) + held_semantics = ObjectSemantics( + label="held", + entity=entities["held"], + geometry={}, + affordance=Affordance(), + ) + held = HeldObjectState( + semantics=held_semantics, + object_to_eef=torch.eye(4).repeat(2, 1, 1), + grasp_xpos=torch.eye(4).repeat(2, 1, 1), + env_mask=torch.tensor([True, False]), + ) + state = ExecutionState( + last_qpos=torch.zeros(2, 8), + held_objects={"physical_left_arm": held}, + ) + grounded = GroundedAction( + "PickUp", + "right_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + object_uid="target", + ) + + scene = adapter._scene_snapshot(grounded, state) + + assert torch.equal( + scene.entities["target"].pose[:, 2, 3], + actual[:, 2, 3] + actions._COLLISION_PARKING_Z_OFFSET, + ) + assert scene.entities["held"].pose[0, 2, 3] == ( + actual[0, 2, 3] + actions._COLLISION_PARKING_Z_OFFSET + ) + assert scene.entities["held"].pose[1, 2, 3] == actual[1, 2, 3] + assert torch.equal(scene.entities["other"].pose, actual) + + +def test_released_object_returns_to_live_dynamic_collision_pose() -> None: + actual = torch.eye(4).repeat(2, 1, 1) + actual[:, 0, 3] = torch.tensor([0.2, 0.4]) + entity = _PoseEntity(actual) + adapter = AtomicActionAdapter( + _planner_env(rigid_objects={"released": entity}), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": ["released"], + }, + ) + grounded = GroundedAction( + "MoveJoints", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + object_uid="released", + ) + + scene = adapter._scene_snapshot( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + assert torch.equal(scene.entities["released"].pose, actual) + + +def test_default_scene_provider_advances_only_after_material_change() -> None: + actual = torch.eye(4).repeat(2, 1, 1) + entity = _PoseEntity(actual.clone()) + adapter = AtomicActionAdapter( + _planner_env(rigid_objects={"can": entity}), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": ["can"], + }, + ) + grounded = GroundedAction( + "MoveJoints", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + object_uid="can", + ) + state = ExecutionState(last_qpos=torch.zeros(2, 8)) + + first = adapter._scene_snapshot(grounded, state) + unchanged = adapter._scene_snapshot(grounded, state) + entity.pose[:, 0, 3] += 0.1 + changed = adapter._scene_snapshot(grounded, state) + + assert first.version == unchanged.version == 0 + assert changed.version == 1 + assert changed.collision_world_revisions(2) == (1, 1) + + +def test_external_scene_provider_is_used_by_planning_snapshot() -> None: + pose = torch.eye(4).repeat(2, 1, 1) + + class _Provider: + def snapshot(self, *, timestamp: float, env_ids: torch.Tensor) -> SceneSnapshot: + assert timestamp == 0.0 + assert torch.equal(env_ids, torch.tensor([0, 1])) + return SceneSnapshot( + timestamp=timestamp, + version=7, + entities={"can": actions.EntityState(pose)}, + ) + + adapter = AtomicActionAdapter( + _planner_env(), + scene_provider=_Provider(), + ) + grounded = GroundedAction( + "MoveJoints", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + object_uid="can", + ) + + scene = adapter._scene_snapshot( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + assert scene.version == 7 + assert torch.equal(scene.entities["can"].pose, pose) + + +def test_start_session_delegates_to_shared_atomic_engine(monkeypatch: Any) -> None: + adapter = AtomicActionAdapter(_planner_env()) + grounded = GroundedAction( + "MoveJoints", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + ) + state = ExecutionState(last_qpos=torch.zeros(2, 8)) + marker = object() + captured: dict[str, Any] = {} + + monkeypatch.setattr(adapter, "_planning_context", lambda *_args: "context") + monkeypatch.setattr(adapter, "_invocation", lambda *_args: "invocation") + + class _Engine: + def start(self, invocations: tuple[Any, ...], context: Any) -> object: + captured["invocations"] = invocations + captured["context"] = context + return marker + + monkeypatch.setattr(adapter, "_engine", lambda: _Engine()) + + result = adapter.start_session(grounded, state) + + assert result is marker + assert captured == {"invocations": ("invocation",), "context": "context"} + + +def test_retreat_parks_intentional_contact_objects() -> None: + actual = torch.eye(4).repeat(2, 1, 1) + entities = { + uid: _PoseEntity(actual.clone()) for uid in ("released", "container", "other") + } + adapter = AtomicActionAdapter( + _planner_env(rigid_objects=entities), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": list(entities), + }, + ) + grounded = GroundedAction( + "MoveEndEffector", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + motion_policy={ + "collision_exclusion_uids": ["released", "container"], + }, + object_uid="released", + ) + + scene = adapter._scene_snapshot( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + parked_z = actual[:, 2, 3] + actions._COLLISION_PARKING_Z_OFFSET + assert torch.equal(scene.entities["released"].pose[:, 2, 3], parked_z) + assert torch.equal(scene.entities["container"].pose[:, 2, 3], parked_z) + assert torch.equal(scene.entities["other"].pose, actual) + + +def test_action_outcome_commits_state_delta_only_for_verified_rows() -> None: + semantics = ObjectSemantics( + label="cube", + entity=object(), + geometry={}, + affordance=Affordance(), + ) + held = HeldObjectState( + semantics=semantics, + object_to_eef=torch.eye(4).repeat(2, 1, 1), + grasp_xpos=torch.eye(4).repeat(2, 1, 1), + ) + prior = ExecutionState(last_qpos=torch.zeros(2, 3)) + trajectory = torch.stack( + (torch.zeros(2, 3), torch.ones(2, 3)), + dim=1, + ) + delta = StateDelta(held_object_updates={"physical_left_arm": held}) + projected = ExecutionState.from_task_state( + delta.apply(prior.to_task_state(), torch.ones(2, dtype=torch.bool)), + last_qpos=trajectory[:, -1], + ) + grounded = GroundedAction( + "PickUp", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + ) + outcome = ActionOutcome( + trajectory=trajectory, + success=torch.ones(2, dtype=torch.bool), + next_state=projected, + grounded=grounded, + prior_state=prior, + expected_effects=delta, + ) + + committed = outcome.state_after(torch.tensor([True, False])) + + assert torch.equal(committed.last_qpos[0], torch.ones(3)) + assert torch.equal(committed.last_qpos[1], torch.zeros(3)) + committed_held = committed.get_held_object("physical_left_arm") + assert committed_held is not None + assert torch.equal(committed_held.env_mask, torch.tensor([True, False])) + + +def test_fallback_rows_keep_the_fallback_plan_effects(monkeypatch: Any) -> None: + env = _planner_env() + adapter = AtomicActionAdapter(env) + semantics = ObjectSemantics( + label="cube", + entity=object(), + geometry={}, + affordance=Affordance(), + ) + + def held_at(x: float) -> HeldObjectState: + relation = torch.eye(4).repeat(2, 1, 1) + relation[:, 0, 3] = x + return HeldObjectState( + semantics=semantics, + object_to_eef=relation, + grasp_xpos=torch.eye(4).repeat(2, 1, 1), + ) + + def action_plan( + success: torch.Tensor, + terminal: float, + held: HeldObjectState, + ) -> ActionPlan: + positions = torch.full((2, 2, 8), terminal) + trajectory = TimedTrajectory.from_uniform_step( + positions, + env_ids=torch.arange(2), + step_dt=0.01, + ) + return ActionPlan( + skill_id="pick_up", + plan_success=success, + commands=_commands_for(trajectory), + joint_trajectory=trajectory, + recovery_policy=RecoveryPolicy(), + planned_scene_version=0, + planned_collision_world_revision=(0, 0), + diagnostics=PlannerDiagnostics(backend="fake"), + expected_effects=StateDelta( + held_object_updates={"physical_left_arm": held} + ), + ) + + plans = iter( + ( + action_plan(torch.tensor([True, False]), 1.0, held_at(1.0)), + action_plan(torch.tensor([True, True]), 2.0, held_at(2.0)), + ) + ) + strategies: list[str] = [] + + def plan(invocation: Any, _context: Any) -> ActionPlan: + strategies.append(invocation.motion_policy.strategy) + return next(plans) + + monkeypatch.setattr( + adapter, + "_engine", + lambda: _FakeEngine(plan), + ) + grounded = GroundedAction( + "PickUp", + "left_arm", + "arm", + GraspGoal(semantics=semantics), + {}, + ) + + outcome = adapter.plan( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + assert strategies == ["motion_gen", "ik_interp"] + assert torch.equal(outcome.success, torch.tensor([True, True])) + assert torch.equal(outcome.next_state.last_qpos[0], torch.ones(8)) + assert torch.equal(outcome.next_state.last_qpos[1], torch.full((8,), 2.0)) + held = outcome.next_state.get_held_object("physical_left_arm") + assert held is not None + assert held.object_to_eef[0, 0, 3] == 1.0 + assert held.object_to_eef[1, 0, 3] == 2.0 + assert torch.equal( + outcome.planner_trace["primary_success"], torch.tensor([True, False]) + ) + assert torch.equal( + outcome.planner_trace["fallback_attempted"], torch.tensor([False, True]) + ) + assert torch.equal( + outcome.planner_trace["fallback_used"], torch.tensor([False, True]) + ) + + +def test_collision_required_cleanup_does_not_use_unsafe_fallback( + monkeypatch: Any, +) -> None: + pose = torch.eye(4).repeat(2, 1, 1) + adapter = AtomicActionAdapter( + _planner_env(rigid_objects={"released": _PoseEntity(pose)}), + planner_policy={ + "dynamic_collision": True, + "dynamic_obstacle_uids": ["released"], + }, + ) + failed_trajectory = TimedTrajectory.from_uniform_step( + torch.zeros(2, 2, 8), + env_ids=torch.arange(2), + step_dt=0.01, + ) + failed_plan = ActionPlan( + skill_id="move_joints", + plan_success=torch.tensor([False, False]), + commands=_commands_for(failed_trajectory), + joint_trajectory=failed_trajectory, + recovery_policy=RecoveryPolicy(), + planned_scene_version=1, + planned_collision_world_revision=(1, 1), + diagnostics=PlannerDiagnostics(backend="fake"), + expected_effects=StateDelta(), + ) + strategies: list[str] = [] + + def plan(invocation: Any, _context: Any) -> ActionPlan: + strategies.append(invocation.motion_policy.strategy) + assert invocation.motion_policy.dynamic_collision_mode.value == "required" + return failed_plan + + monkeypatch.setattr(adapter, "_engine", lambda: _FakeEngine(plan)) + grounded = GroundedAction( + "MoveJoints", + "left_arm", + "arm", + JointPositionGoal(target=torch.zeros(2, 2)), + {}, + motion_policy={"collision_safety": "required"}, + object_uid="released", + ) + + outcome = adapter.plan( + grounded, + ExecutionState(last_qpos=torch.zeros(2, 8)), + ) + + assert strategies == ["motion_gen"] + assert not bool(outcome.success.any()) + assert outcome.planner_trace["fallback_allowed"] is False + assert not bool(outcome.planner_trace["fallback_attempted"].any()) + assert not bool(outcome.planner_trace["fallback_used"].any()) + assert outcome.planner_trace["collision_obstacle_positions"]["released"].shape == ( + 2, + 3, + ) diff --git a/tests/gen_sim/action_engine/runtime/test_atomic_compat.py b/tests/gen_sim/action_engine/runtime/test_atomic_compat.py new file mode 100644 index 000000000..e6c2d7573 --- /dev/null +++ b/tests/gen_sim/action_engine/runtime/test_atomic_compat.py @@ -0,0 +1,87 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from embodichain.gen_sim.action_engine.runtime.actions import AtomicActionAdapter +from embodichain.gen_sim.action_engine.runtime.atomic_compat import ( + ExactTargetMoveHeldObject, + ExactTargetMoveHeldObjectOptions, +) +from embodichain.lab.sim.atomic_actions import MoveHeldObject, MoveHeldObjectOptions + + +def test_exact_target_transport_only_disables_rotation_when_requested( + monkeypatch: pytest.MonkeyPatch, +) -> None: + applied = [] + result = object() + + def fake_apply(self, move_eef_xpos, end_arm_xpos) -> None: + del self, move_eef_xpos, end_arm_xpos + applied.append(True) + + def fake_plan(self, request, context): + del request, context + self._apply_automatic_transport_rotation(torch.eye(4), torch.eye(4)) + return result + + monkeypatch.setattr( + MoveHeldObject, + "_apply_automatic_transport_rotation", + fake_apply, + ) + monkeypatch.setattr(MoveHeldObject, "_plan", fake_plan) + action = ExactTargetMoveHeldObject() + disabled_request = SimpleNamespace( + skill_options=ExactTargetMoveHeldObjectOptions( + allow_automatic_transport_rotation=False, + ) + ) + enabled_request = SimpleNamespace( + skill_options=ExactTargetMoveHeldObjectOptions(), + ) + + assert action._plan(disabled_request, object()) is result + assert not applied + assert action._plan(enabled_request, object()) is result + assert applied == [True] + + +@pytest.mark.parametrize( + ("yaw_samples", "expected"), + [(1, True), (8, False)], +) +def test_semantic_transport_config_scopes_rotation_override( + yaw_samples: int, + expected: bool, +) -> None: + adapter = AtomicActionAdapter.__new__(AtomicActionAdapter) + action = SimpleNamespace(cfg={"upright_yaw_samples": yaw_samples}) + capability = SimpleNamespace( + config_type=MoveHeldObjectOptions, + target_materializer="semantic_held_object", + ) + + options = adapter._build_single_arm_config(action, capability) + + assert isinstance(options, ExactTargetMoveHeldObjectOptions) + assert options.allow_automatic_transport_rotation is expected diff --git a/tests/gen_sim/action_engine/runtime/test_grasp_collision_cache.py b/tests/gen_sim/action_engine/runtime/test_grasp_collision_cache.py new file mode 100644 index 000000000..a08771054 --- /dev/null +++ b/tests/gen_sim/action_engine/runtime/test_grasp_collision_cache.py @@ -0,0 +1,354 @@ +# ---------------------------------------------------------------------------- +# Copyright (c) 2021-2026 DexForce Technology Co., Ltd. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ---------------------------------------------------------------------------- + +from __future__ import annotations + +import hashlib +import json +import os +from pathlib import Path +import pickle +from typing import Callable + +import numpy as np +import pytest +import torch + +from embodichain.gen_sim.action_engine.runtime import grasp_collision_cache +from embodichain.gen_sim.action_engine.runtime.grasp_collision_cache import ( + GraspCollisionCacheError, + ensure_vhacd_grasp_collision_cache, + grasp_collision_cache_path, +) + + +def _tetrahedron() -> tuple[torch.Tensor, torch.Tensor]: + vertices = torch.tensor( + [ + [0.0, 0.0, 0.0], + [1.0, 0.0, 0.0], + [0.0, 1.0, 0.0], + [0.0, 0.0, 1.0], + ], + dtype=torch.float32, + ) + triangles = torch.tensor( + [ + [0, 2, 1], + [0, 1, 3], + [0, 3, 2], + [1, 2, 3], + ], + dtype=torch.int64, + ) + return vertices, triangles + + +def _plane_equations() -> list[tuple[np.ndarray, np.ndarray]]: + return [ + ( + np.asarray( + [ + [1.0, 0.0, 0.0], + [0.0, 1.0, 0.0], + [0.0, 0.0, 1.0], + ], + dtype=np.float32, + ), + np.asarray([-1.0, -1.0, -1.0], dtype=np.float32), + ), + ( + np.asarray([[1.0, 1.0, 1.0]], dtype=np.float32), + np.asarray([-1.0], dtype=np.float32), + ), + ] + + +def _install_fake_decomposer( + monkeypatch: pytest.MonkeyPatch, +) -> list[tuple[tuple[int, ...], tuple[int, ...], int]]: + calls: list[tuple[tuple[int, ...], tuple[int, ...], int]] = [] + + def fake_decompose( + vertices: np.ndarray, + triangles: np.ndarray, + max_decomposition_hulls: int, + ) -> list[tuple[np.ndarray, np.ndarray]]: + calls.append( + ( + tuple(vertices.shape), + tuple(triangles.shape), + max_decomposition_hulls, + ) + ) + return _plane_equations() + + monkeypatch.setattr( + grasp_collision_cache, + "_compute_vhacd_plane_equations", + fake_decompose, + ) + return calls + + +def test_cache_key_and_payload_match_main_checker_contract( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + _install_fake_decomposer(monkeypatch) + + result = ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=16, + cache_dir=tmp_path, + ) + + expected_hash = hashlib.md5( + vertices.numpy().tobytes() + triangles.numpy().tobytes() + ).hexdigest() + assert result.cache_path == tmp_path / f"{expected_hash}_16.pkl" + with result.cache_path.open("rb") as cache_file: + payload = pickle.load(cache_file) + assert set(payload) == {"plane_equations", "plane_equation_counts"} + assert payload["plane_equations"].shape == (2, 3, 4) + assert payload["plane_equations"].dtype == torch.float32 + assert payload["plane_equation_counts"].tolist() == [3, 1] + assert payload["plane_equation_counts"].dtype == torch.int32 + + +def test_main_checker_loads_prepared_cache_without_running_coacd( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + import embodichain.lab.sim + from embodichain.toolkits.graspkit.pg_grasp import collision_checker + + vertices, triangles = _tetrahedron() + _install_fake_decomposer(monkeypatch) + result = ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=16, + cache_dir=tmp_path, + ) + + def fail_coacd(*args: object, **kwargs: object) -> None: + raise AssertionError("The prepared V-HACD cache must bypass CoACD.") + + monkeypatch.setattr(embodichain.lab.sim, "CONVEX_DECOMP_DIR", tmp_path) + monkeypatch.setattr(collision_checker, "convex_decomposition_coacd", fail_coacd) + checker = collision_checker.ConvexCollisionChecker( + vertices, + triangles, + max_decomposition_hulls=16, + ) + + assert checker.cache_path == result.cache_path.as_posix() + assert checker.plane_equations["plane_equation_counts"].tolist() == [3, 1] + + +def test_matching_vhacd_metadata_returns_cache_hit( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + calls = _install_fake_decomposer(monkeypatch) + kwargs = { + "mesh_vertices": vertices, + "mesh_triangles": triangles, + "max_decomposition_hulls": 16, + "cache_dir": tmp_path, + } + + first = ensure_vhacd_grasp_collision_cache(**kwargs) + second = ensure_vhacd_grasp_collision_cache(**kwargs) + + assert first.status == "generated" + assert second.status == "hit" + assert len(calls) == 1 + + +def test_non_vhacd_metadata_forces_cache_replacement( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + calls = _install_fake_decomposer(monkeypatch) + kwargs = { + "mesh_vertices": vertices, + "mesh_triangles": triangles, + "max_decomposition_hulls": 16, + "cache_dir": tmp_path, + } + first = ensure_vhacd_grasp_collision_cache(**kwargs) + metadata = json.loads(first.metadata_path.read_text(encoding="utf-8")) + metadata["backend"] = "coacd" + first.metadata_path.write_text(json.dumps(metadata), encoding="utf-8") + + replaced = ensure_vhacd_grasp_collision_cache(**kwargs) + + assert replaced.status == "replaced" + assert len(calls) == 2 + repaired = json.loads(replaced.metadata_path.read_text(encoding="utf-8")) + assert repaired["backend"] == "vhacd" + + +def test_modified_cache_fails_checksum_and_is_rebuilt( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + calls = _install_fake_decomposer(monkeypatch) + kwargs = { + "mesh_vertices": vertices, + "mesh_triangles": triangles, + "max_decomposition_hulls": 16, + "cache_dir": tmp_path, + } + first = ensure_vhacd_grasp_collision_cache(**kwargs) + first.cache_path.write_bytes(b"not a valid collision cache") + + replaced = ensure_vhacd_grasp_collision_cache(**kwargs) + + assert replaced.status == "replaced" + assert len(calls) == 2 + with replaced.cache_path.open("rb") as cache_file: + assert "plane_equations" in pickle.load(cache_file) + + +def test_cache_and_metadata_are_published_by_atomic_replace( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + _install_fake_decomposer(monkeypatch) + replacements: list[tuple[Path, Path]] = [] + real_replace: Callable[[os.PathLike[str], os.PathLike[str]], None] = os.replace + + def recording_replace( + source: os.PathLike[str], + destination: os.PathLike[str], + ) -> None: + replacements.append((Path(source), Path(destination))) + real_replace(source, destination) + + monkeypatch.setattr(grasp_collision_cache.os, "replace", recording_replace) + + result = ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=16, + cache_dir=tmp_path, + ) + + assert [destination for _, destination in replacements] == [ + result.cache_path, + result.metadata_path, + ] + assert all( + source.parent == destination.parent for source, destination in replacements + ) + assert all(not source.exists() for source, _ in replacements) + + +@pytest.mark.parametrize( + ("vertices", "triangles", "message"), + [ + ( + torch.empty((0, 3), dtype=torch.float32), + torch.tensor([[0, 1, 2]], dtype=torch.int64), + "mesh_vertices", + ), + ( + torch.zeros((3, 3), dtype=torch.float32), + torch.tensor([[0, 1]], dtype=torch.int64), + "mesh_triangles", + ), + ( + torch.tensor([[0.0, 0.0, 0.0], [1.0, float("nan"), 0.0], [0.0, 1.0, 0.0]]), + torch.tensor([[0, 1, 2]], dtype=torch.int64), + "finite", + ), + ( + torch.zeros((3, 3), dtype=torch.float32), + torch.tensor([[0, 1, 3]], dtype=torch.int64), + "indices", + ), + ], +) +def test_invalid_mesh_is_rejected_before_decomposition( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, + vertices: torch.Tensor, + triangles: torch.Tensor, + message: str, +) -> None: + calls = _install_fake_decomposer(monkeypatch) + + with pytest.raises(ValueError, match=message): + ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=16, + cache_dir=tmp_path, + ) + + assert calls == [] + + +@pytest.mark.parametrize("max_decomposition_hulls", [True, 0, -1, 1.5]) +def test_invalid_hull_limit_is_rejected( + tmp_path: Path, + max_decomposition_hulls: object, +) -> None: + vertices, triangles = _tetrahedron() + + with pytest.raises((TypeError, ValueError), match="max_decomposition_hulls"): + ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=max_decomposition_hulls, # type: ignore[arg-type] + cache_dir=tmp_path, + ) + + +def test_symlinked_cache_path_is_refused( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + vertices, triangles = _tetrahedron() + _install_fake_decomposer(monkeypatch) + cache_path = grasp_collision_cache_path( + vertices, + triangles, + 16, + cache_dir=tmp_path, + ) + victim = tmp_path / "victim.pkl" + victim.write_bytes(b"do not overwrite") + cache_path.symlink_to(victim) + + with pytest.raises(GraspCollisionCacheError, match="symlink"): + ensure_vhacd_grasp_collision_cache( + mesh_vertices=vertices, + mesh_triangles=triangles, + max_decomposition_hulls=16, + cache_dir=tmp_path, + ) + + assert victim.read_bytes() == b"do not overwrite"