diff --git a/embodichain/lab/sim/atomic_actions/__init__.py b/embodichain/lab/sim/atomic_actions/__init__.py index fd1e0429c..8611766fb 100644 --- a/embodichain/lab/sim/atomic_actions/__init__.py +++ b/embodichain/lab/sim/atomic_actions/__init__.py @@ -80,7 +80,6 @@ ActionPlan, CompiledTrajectory, EffectVerificationRequirement, - ExecutionFeedbackMode, PlannerDiagnostics, TimedTrajectory, TrajectorySegment, @@ -108,6 +107,46 @@ TimedCommandSequence, ) from .transports import EndpointCommandRouter, EndpointCommandTransport +from .tracking import ( + BASE_POSE_CHANNEL, + JOINT_POSITION_CHANNEL, + WHOLE_BODY_POSE_CHANNEL, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, + FeedbackTerminalAcceptance, + InFlightTrackingPolicy, + JointPositionTrackingEvaluator, + JointPositionTrackingMetric, + JointPositionTrackingProjector, + JointPositionTrackingState, + PlanningContextTrackingFeedbackProvider, + PoseTrackingEvaluator, + PoseTrackingMetric, + PoseTrackingState, + TerminalAcceptance, + TimedTerminalAcceptance, + TimedTrackingSequence, + TrackingCommandProjector, + TrackingEvaluation, + TrackingEvaluatorRegistry, + TrackingFeedbackAddress, + TrackingFeedbackBatch, + TrackingFeedbackProvider, + TrackingFeedbackProviderRegistry, + TrackingFeedbackSourceRef, + TrackingFrame, + TrackingMetricCfg, + TrackingMetricEvaluator, + TrackingPolicy, + TrackingProjectorRef, + TrackingProjectorRegistry, + TrackingRuntime, + TrackingSetpoint, + TrackingState, + WholeBodyPoseTrackingEvaluator, + WholeBodyPoseTrackingMetric, + WholeBodyPoseTrackingState, +) from .primitives import ( AssembleGoal, BUILTIN_ACTION_TYPES, @@ -225,7 +264,6 @@ "EffectVerificationResult", "EffectVerifier", "ExecutionClock", - "ExecutionFeedbackMode", "ExecutionEvent", "ExecutionEventKind", "ExecutionPlanAttempt", @@ -234,6 +272,9 @@ "ExecutionSession", "ExecutionStatus", "ExecutionTick", + "EndpointTrackingChannelBinding", + "EndpointTrackingFeedbackAddress", + "FeedbackTerminalAcceptance", "GRASP_COMMAND", "GRASP_CAPABILITY", "GraspGoal", @@ -243,12 +284,16 @@ "HeldObjectState", "FORWARD_KINEMATICS_CAPABILITY", "INVERSE_KINEMATICS_CAPABILITY", + "InFlightTrackingPolicy", "InteractionPoints", "JointPositionGoal", "JointPositionCommand", "JointPositionPayload", "JointPositionTarget", "JOINT_POSITION_CAPABILITY", + "JOINT_POSITION_CHANNEL", + "JointPositionTrackingMetric", + "JointPositionTrackingState", "MotionPolicy", "MonotonicExecutionClock", "MoveEndEffector", @@ -273,6 +318,8 @@ "PlannerDiagnostics", "PlanningContext", "PoseGoalValue", + "PoseTrackingMetric", + "PoseTrackingState", "Press", "PressGoal", "PressOptions", @@ -300,8 +347,37 @@ "SimulationExecutionAdapter", "TaskState", "TimedCommandSequence", + "TimedTerminalAcceptance", + "TimedTrackingSequence", "TimedTrajectory", + "TerminalAcceptance", "TrajectorySegment", + "TrackingCommandProjector", + "TrackingEvaluation", + "TrackingEvaluatorRegistry", + "TrackingFeedbackAddress", + "TrackingFeedbackBatch", + "TrackingFeedbackProvider", + "TrackingFeedbackProviderRegistry", + "TrackingFeedbackSourceRef", + "TrackingFrame", + "TrackingMetricCfg", + "TrackingMetricEvaluator", + "TrackingPolicy", + "TrackingProjectorRef", + "TrackingProjectorRegistry", + "TrackingRuntime", + "TrackingSetpoint", + "TrackingState", + "BASE_POSE_CHANNEL", + "JointPositionTrackingEvaluator", + "JointPositionTrackingProjector", + "PlanningContextTrackingFeedbackProvider", + "PoseTrackingEvaluator", + "WHOLE_BODY_POSE_CHANNEL", + "WholeBodyPoseTrackingEvaluator", + "WholeBodyPoseTrackingMetric", + "WholeBodyPoseTrackingState", "get_registered_actions", "register_action", "unregister_action", diff --git a/embodichain/lab/sim/atomic_actions/bindings.py b/embodichain/lab/sim/atomic_actions/bindings.py index 3043c5b5c..1180244a4 100644 --- a/embodichain/lab/sim/atomic_actions/bindings.py +++ b/embodichain/lab/sim/atomic_actions/bindings.py @@ -28,6 +28,10 @@ import torch from .control import ControlCommand +from .tracking import ( + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, +) def _validate_identifier(value: str, *, field_name: str) -> str: @@ -79,6 +83,49 @@ def _snapshot_commands( return MappingProxyType(commands) +def _snapshot_tracking_channels( + values: Mapping[str, EndpointTrackingChannelBinding], + *, + target: RuntimeEndpointTarget, +) -> Mapping[str, EndpointTrackingChannelBinding]: + """Validate and own endpoint-local tracking-channel bindings.""" + if not isinstance(values, Mapping): + raise TypeError("EndpointBinding.tracking_channels must be a mapping.") + channels: dict[str, EndpointTrackingChannelBinding] = {} + for channel_id, binding in values.items(): + _validate_identifier( + channel_id, + field_name="EndpointBinding tracking channel IDs", + ) + if not isinstance(binding, EndpointTrackingChannelBinding): + raise TypeError( + "EndpointBinding.tracking_channels values must be " + "EndpointTrackingChannelBinding instances." + ) + if binding.channel_id != channel_id: + raise ValueError( + f"Tracking channel key {channel_id!r} disagrees with its binding " + f"channel {binding.channel_id!r}." + ) + snapshot = binding.snapshot() + if snapshot is binding: + raise TypeError( + "EndpointTrackingChannelBinding.snapshot() must return an " + "independently owned value." + ) + address = snapshot.source.address + if ( + isinstance(address, EndpointTrackingFeedbackAddress) + and address.target.address_fingerprint != target.address_fingerprint + ): + raise ValueError( + f"Tracking channel {channel_id!r} addresses a different runtime " + "endpoint target." + ) + channels[channel_id] = snapshot + return MappingProxyType(channels) + + def _validate_target_fingerprint( target: RuntimeEndpointTarget, *, @@ -192,6 +239,9 @@ class EndpointBinding: task_state_key: str | None = None """Symbolic task-state key; direct-core defaults to ``target.target_id``.""" + tracking_channels: Mapping[str, EndpointTrackingChannelBinding] = field( + default_factory=dict + ) capabilities: frozenset[str] = frozenset() commands: Mapping[str, ControlCommand] = field(default_factory=dict) claim_tokens: frozenset[str] = frozenset() @@ -246,6 +296,11 @@ def __post_init__(self) -> None: field_name="EndpointBinding.task_state_key", ) object.__setattr__(self, "task_state_key", task_state_key) + object.__setattr__( + self, + "tracking_channels", + _snapshot_tracking_channels(self.tracking_channels, target=target), + ) object.__setattr__( self, "capabilities", @@ -317,6 +372,18 @@ def command(self, name: str) -> ControlCommand: ) from exc return command.snapshot() + def tracking_channel(self, channel_id: str) -> EndpointTrackingChannelBinding: + """Return one independently owned typed tracking-channel binding.""" + try: + binding = self.tracking_channels[channel_id] + except KeyError as exc: + raise KeyError( + f"Endpoint {self.slot_id}.{self.endpoint_id} has no tracking " + f"channel {channel_id!r}; available channels are " + f"{sorted(self.tracking_channels)}." + ) from exc + return binding.snapshot() + def joint_positions( self, name: str, @@ -356,6 +423,7 @@ def with_commands( adapter_id=self.adapter_id, target=self.target, task_state_key=self.task_state_key, + tracking_channels=self.tracking_channels, capabilities=self.capabilities, commands=merged, claim_tokens=self.claim_tokens, @@ -371,6 +439,7 @@ def snapshot(self) -> EndpointBinding: adapter_id=self.adapter_id, target=self.target, task_state_key=self.task_state_key, + tracking_channels=self.tracking_channels, capabilities=self.capabilities, commands=self.commands, claim_tokens=self.claim_tokens, diff --git a/embodichain/lab/sim/atomic_actions/core.py b/embodichain/lab/sim/atomic_actions/core.py index a3d37d55e..892bec50a 100644 --- a/embodichain/lab/sim/atomic_actions/core.py +++ b/embodichain/lab/sim/atomic_actions/core.py @@ -42,7 +42,6 @@ from .plans import ( ActionPlan, EffectVerificationRequirement, - ExecutionFeedbackMode, PlannerDiagnostics, TimedTrajectory, TrajectorySegment, @@ -56,6 +55,12 @@ RuntimeCommandFrame, TimedCommandSequence, ) +from .tracking import ( + FeedbackTerminalAcceptance, + TimedTrackingSequence, + TrackingFrame, + TrackingSetpoint, +) if TYPE_CHECKING: from embodichain.lab.sim.objects import Robot @@ -370,6 +375,7 @@ def resolve_request( invocation.control_overrides, ), motion_policy=invocation.motion_policy, + tracking_policy=invocation.tracking_policy, recovery_policy=invocation.recovery_policy, skill_options=options, invocation_id=invocation.invocation_id, @@ -574,7 +580,6 @@ def build_plan( diagnostics=diagnostics, segment_lengths=segment_lengths, scene_dependency_monitor_until=scene_dependency_monitor_until, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, joint_trajectory=timed, ) @@ -591,14 +596,13 @@ def build_command_plan( diagnostics: PlannerDiagnostics | None = None, segment_lengths: Mapping[str, int] | None = None, scene_dependency_monitor_until: Mapping[str, int] | None = None, - feedback_mode: ExecutionFeedbackMode = ExecutionFeedbackMode.TIMED, joint_trajectory: TimedTrajectory | None = None, ) -> ActionPlan: """Build a plan from transport-neutral runtime command frames. - Non-joint command sequences use timed completion unless a future - endpoint-specific feedback evaluator is installed. Semantic effects - remain externally verified through the execution session. + Tracking targets are projected from the command payloads through the + typed channels declared by each bound endpoint. Semantic effects remain + externally verified through the execution session. Args: request: Resolved invocation snapshot being planned. @@ -617,9 +621,8 @@ def build_command_plan( its bound. ``0`` disables monitoring immediately; omitted dependencies remain monitored for the full action. Once the bound is reached, all pose changes for that entity are ignored. - feedback_mode: Feedback contract used to determine target completion. - joint_trajectory: Optional joint trajectory retained for joint-position - feedback and inspection. + joint_trajectory: Optional joint trajectory retained for offline + compilation and inspection. Returns: Side-effect-free action plan. @@ -647,6 +650,7 @@ def build_command_plan( ), env_ids=commands.env_ids, ) + tracking = self._tracking_sequence(request, masked_commands) segments = self._build_segments( segment_lengths, frame_count=masked_commands.frame_count, @@ -660,12 +664,13 @@ def build_command_plan( plan_success=success_mask, commands=masked_commands, recovery_policy=request.recovery_policy, + tracking_policy=request.tracking_policy, planned_scene_version=context.scene.version, planned_collision_world_revision=( context.scene.collision_world_revisions(context.batch_size) ), diagnostics=diagnostics, - feedback_mode=feedback_mode, + tracking=tracking, joint_trajectory=joint_trajectory, segments=segments, scene_dependencies=self._scene_dependencies(request), @@ -682,6 +687,67 @@ def build_command_plan( invocation_revision=request.revision, ) + def _tracking_sequence( + self, + request: ResolvedActionRequest[GoalT, OptionsT], + commands: TimedCommandSequence, + ) -> TimedTrackingSequence | None: + """Project command payloads through binding-owned tracking channels.""" + policy = request.tracking_policy + metrics = list(() if policy.in_flight is None else policy.in_flight.metrics) + if isinstance(policy.terminal, FeedbackTerminalAcceptance): + metrics.extend(policy.terminal.metrics) + if not metrics: + return None + + runtime = self.planning_services.tracking_runtime + for metric in metrics: + runtime.evaluators.resolve(metric) + metrics_by_channel = {metric.channel_id: metric for metric in metrics} + + endpoints_by_destination: dict[ + tuple[str, str], + tuple[EndpointBinding, ...], + ] = {} + for endpoint in request.binding.endpoints: + endpoints_by_destination.setdefault(endpoint.destination_key, ()) + endpoints_by_destination[endpoint.destination_key] += (endpoint,) + + tracking_frames: list[TrackingFrame] = [] + for frame_index, frame in enumerate(commands.frames): + setpoints: list[TrackingSetpoint] = [] + for command in frame.commands: + endpoints = endpoints_by_destination[command.destination_key] + for endpoint in endpoints: + for channel_id in metrics_by_channel: + channel = endpoint.tracking_channels.get(channel_id) + if channel is None: + continue + runtime.providers.resolve(channel.source) + runtime.projectors.resolve(channel.projector) + setpoints.append( + TrackingSetpoint( + endpoint_key=endpoint.key, + binding=channel, + desired=runtime.project(command, channel), + ) + ) + covered_channels = {setpoint.binding.channel_id for setpoint in setpoints} + missing_channels = sorted( + set(metrics_by_channel).difference(covered_channels) + ) + if missing_channels: + raise ValueError( + f"Command frame {frame_index} cannot project configured " + f"tracking channels {missing_channels}; bound endpoints must " + "declare a typed feedback source and projector." + ) + tracking_frames.append(TrackingFrame(tuple(setpoints))) + return TimedTrackingSequence( + env_ids=commands.env_ids, + frames=tuple(tracking_frames), + ) + @staticmethod def _authorize_command_targets( request: ResolvedActionRequest[GoalT, OptionsT], diff --git a/embodichain/lab/sim/atomic_actions/engine.py b/embodichain/lab/sim/atomic_actions/engine.py index ff1d5eb30..c9f6dbce2 100644 --- a/embodichain/lab/sim/atomic_actions/engine.py +++ b/embodichain/lab/sim/atomic_actions/engine.py @@ -30,6 +30,7 @@ from .plans import ActionPlan, CompiledTrajectory, TimedTrajectory from .runtime import ActionPlanningServices from .state import PlanningContext, RobotObservation, SceneSnapshot, TaskState +from .tracking import TrackingRuntime if TYPE_CHECKING: from embodichain.lab.sim.objects import Robot @@ -100,6 +101,7 @@ def __init__( endpoint_adapters: ( Mapping[type[ResourceEndpoint], ResourceEndpointAdapter] | None ) = None, + tracking_runtime: TrackingRuntime | None = None, ) -> None: """Initialize one engine and bind its built-in action implementations. @@ -114,6 +116,9 @@ def __init__( ``skill_profile`` are mutually exclusive. endpoint_adapters: Optional exact-type endpoint adapters used when binding ``skill_profile``. Invalid without a profile. + tracking_runtime: Optional exact-version feedback, projector, and + metric registries. Built-in joint tracking is installed when + omitted. """ if endpoint_adapters is not None and skill_profile is None: raise ValueError("endpoint_adapters requires skill_profile.") @@ -131,6 +136,7 @@ def __init__( self._planning_services = ActionPlanningServices( motion_generator, control_profiles=control_profiles, + tracking_runtime=tracking_runtime, ) self._actions: dict[str, AtomicAction] = {} self._skill_catalog_revision = 0 @@ -163,6 +169,11 @@ def planning_services(self) -> ActionPlanningServices: """Engine-owned resources shared by every bound atomic action.""" return self._planning_services + @property + def tracking_runtime(self) -> TrackingRuntime: + """Typed endpoint-feedback runtime used by plans and sessions.""" + return self._planning_services.tracking_runtime + @property def binding_owner_id(self) -> str: """Return the opaque owner identity required by action bindings.""" @@ -617,6 +628,11 @@ def _validate_plan( raise ValueError( "ActionPlan.invocation_revision must preserve the request revision." ) + if plan.tracking_policy != request.tracking_policy: + raise ValueError( + "ActionPlan.tracking_policy must preserve the resolved request " + "tracking policy." + ) commands = plan.commands if commands.batch_size != context.batch_size: raise ValueError("Action plan batch size does not match the context.") diff --git a/embodichain/lab/sim/atomic_actions/execution.py b/embodichain/lab/sim/atomic_actions/execution.py index c30691c28..4876a0c25 100644 --- a/embodichain/lab/sim/atomic_actions/execution.py +++ b/embodichain/lab/sim/atomic_actions/execution.py @@ -27,20 +27,22 @@ from .effects import StateDelta from .invocation import ActionInvocation, ResolvedActionRequest -from .bindings import JointPositionTarget, RuntimeEndpointTarget +from .bindings import RuntimeEndpointTarget from .plans import ( ActionPlan, EffectVerificationRequirement, - ExecutionFeedbackMode, TrajectorySegment, ) from .policies import RecoveryPolicy -from .runtime_commands import ( - JointPositionPayload, - RuntimeCommandFrame, - TimedCommandSequence, -) +from .runtime_commands import RuntimeCommandFrame, TimedCommandSequence from .state import EntityState, PlanningContext, SceneSnapshot, TaskState +from .tracking import ( + FeedbackTerminalAcceptance, + TimedTerminalAcceptance, + TrackingEvaluation, + TrackingFrame, + TrackingMetricCfg, +) if TYPE_CHECKING: from .engine import AtomicActionEngine @@ -60,7 +62,10 @@ class ExecutionEventKind(str, Enum): ACTION_PLANNED = "action_planned" INVOCATION_REVISED = "invocation_revised" REPLANNED = "replanned" - TRACKING_ERROR = "tracking_error" + TRACKING_DIVERGED = "tracking_diverged" + TRACKING_FEEDBACK_FAILED = "tracking_feedback_failed" + TERMINAL_ACCEPTANCE_PENDING = "terminal_acceptance_pending" + TERMINAL_ACCEPTANCE_FAILED = "terminal_acceptance_failed" DYNAMIC_GOAL_CHANGED = "dynamic_goal_changed" COLLISION_WORLD_CHANGED = "collision_world_changed" ACTION_PLANNING_FAILED = "action_planning_failed" @@ -441,14 +446,27 @@ def __init__( tuple[str, str], RuntimeEndpointTarget, ] = {} + self._active_tracking_routes: dict[ + tuple[str, str, str], + tuple[object, str, str], + ] = {} self._planned_scene = context.scene self._action_started_at = context.robot.timestamp self._attempt_generation = -1 - self._last_joint_command: torch.Tensor | None = None - self._last_joint_ids: tuple[int, ...] = () + self._last_tracking_frame: TrackingFrame | None = None self._last_command_mask = torch.zeros( context.batch_size, dtype=torch.bool, device=context.robot.qpos.device ) + self._tracking_violation_counts = torch.zeros( + context.batch_size, + dtype=torch.long, + device=context.robot.qpos.device, + ) + self._terminal_acceptance_counts = torch.zeros_like( + self._tracking_violation_counts + ) + self._terminal_started_at: float | None = None + self._terminal_pending_reported = False self._eligible = ( torch.ones_like(self._last_command_mask) if eligible_mask is None @@ -655,6 +673,10 @@ def _install_prepared_revision( replacement_plan, ExecutionEventKind.INVOCATION_REVISED, ) + self._validate_tracking_continuity( + replacement_plan, + ExecutionEventKind.INVOCATION_REVISED, + ) requests = list(self._requests) requests[self._invocation_index] = replacement @@ -888,14 +910,7 @@ def tick( events.extend(recovery_events) if self._status is not ExecutionStatus.RUNNING: return self._tick_result(command=None, events=events) - if recovery_events and any( - event.kind - in { - ExecutionEventKind.REPLANNED, - ExecutionEventKind.RECOVERY_EXHAUSTED, - } - for event in recovery_events - ): + if recovery_events: assert self._plan is not None plan = self._plan execution_mask = self._pending & self._plan.plan_success @@ -929,59 +944,97 @@ def tick( self._waypoint_index += 1 return self._tick_result(command=command, events=events) - terminal_error = self._terminal_error(plan) - not_reached = execution_mask & ( - terminal_error > plan.recovery_policy.tracking_error_threshold - ) - if not_reached.any(): - max_terminal_error = float(terminal_error[not_reached].amax().item()) - events.extend( - self._attempt_replan( - not_reached, - ExecutionEventKind.TRACKING_ERROR, - "Terminal command has not been reached " - f"(max_error={max_terminal_error:.6f}, " - "threshold=" - f"{plan.recovery_policy.tracking_error_threshold:.6f}).", + terminal = plan.tracking_policy.terminal + if self._terminal_started_at is None: + self._terminal_started_at = self._context.robot.timestamp + elapsed_terminal = self._context.robot.timestamp - self._terminal_started_at + terminal_pending = torch.zeros_like(execution_mask) + if isinstance(terminal, TimedTerminalAcceptance): + if elapsed_terminal < terminal.settle_duration: + terminal_pending = execution_mask.clone() + elif isinstance(terminal, FeedbackTerminalAcceptance): + if plan.tracking is None or not plan.tracking.frames: + raise RuntimeError( + "Feedback terminal acceptance requires a terminal tracking " + "frame." ) - ) - if self._status is not ExecutionStatus.RUNNING: - return self._tick_result(command=None, events=events) - assert self._plan is not None - plan = self._plan - execution_mask = self._pending & self._plan.plan_success - if not self._pending.any(): - command, hold_targets, completion_events = self._finish_action( - self._pending, - None, + try: + accepted, valid, normalized_error = self._evaluate_tracking_frame( + plan.tracking.frames[-1], + terminal.metrics, ) - events.extend(completion_events) - return self._tick_result( - command=command, - hold_targets=hold_targets, - events=events, + except Exception as exc: # noqa: BLE001 - fail required feedback closed + events.extend( + self._fail_tracking_feedback( + execution_mask, + "Terminal tracking feedback evaluation failed: " + f"{type(exc).__name__}: {exc}", + ) ) - if plan.commands.frame_count > 0: - command = self._command_at(plan, 0, execution_mask) - self._waypoint_index = 1 - return self._tick_result(command=command, events=events) - events.append( - self._event( - ExecutionEventKind.TRAJECTORY_COMPLETED, - execution_mask, - "Replanned action has no executable command frame.", + return self._tick_result(command=None, events=events) + invalid = execution_mask & ~valid + if invalid.any(): + events.extend( + self._fail_tracking_feedback( + invalid, + "Required terminal tracking feedback was invalid.", + ) ) + if self._status is not ExecutionStatus.RUNNING: + return self._tick_result(command=None, events=events) + execution_mask = self._pending & plan.plan_success + accepted_now = execution_mask & valid & accepted + self._terminal_acceptance_counts[accepted_now] += 1 + self._terminal_acceptance_counts[execution_mask & ~accepted_now] = 0 + terminal_pending = execution_mask & ( + self._terminal_acceptance_counts < terminal.consecutive_acceptances ) - command, hold_targets, completion_events = self._finish_action( - execution_mask, - effect_result, + if terminal_pending.any() and elapsed_terminal >= terminal.settle_timeout: + max_error = float(normalized_error[terminal_pending].amax().item()) + events.extend( + self._attempt_action_retry( + terminal_pending, + ExecutionEventKind.TERMINAL_ACCEPTANCE_FAILED, + "Terminal feedback did not satisfy the acceptance " + "contract before its settle timeout " + f"(max_normalized_error={max_error:.6f}).", + ) + ) + if self._status is not ExecutionStatus.RUNNING: + return self._tick_result(command=None, events=events) + assert self._plan is not None + plan = self._plan + execution_mask = self._pending & plan.plan_success + if plan.commands.frame_count > 0 and execution_mask.any(): + command = self._command_at(plan, 0, execution_mask) + self._waypoint_index = 1 + return self._tick_result(command=command, events=events) + terminal_pending.zero_() + else: # pragma: no cover - TrackingPolicy validates exact alternatives + raise AssertionError( + f"Unsupported terminal acceptance {type(terminal).__name__}." ) - events.extend(completion_events) - return self._tick_result( - command=command, - hold_targets=hold_targets, - events=events, + + if terminal_pending.any(): + if not self._terminal_pending_reported: + events.append( + self._event( + ExecutionEventKind.TERMINAL_ACCEPTANCE_PENDING, + terminal_pending, + "Maintaining the terminal command while acceptance is " + "pending.", + ) + ) + self._terminal_pending_reported = True + if plan.commands.frame_count == 0: + raise RuntimeError( + "Terminal settling requires an executable terminal command " + "frame." + ) + terminal_command = plan.commands.frames[-1].with_active_mask( + plan.commands.frames[-1].active_mask & terminal_pending ) + return self._tick_result(command=terminal_command, events=events) events.append( self._event( @@ -1054,7 +1107,9 @@ def _install_plan( for target in plan.commands.targets } replacement_destinations = frozenset(replacement_targets) + replacement_tracking_routes = self._tracking_routes(plan) self._validate_destination_continuity(plan, event_kind) + self._validate_tracking_continuity(plan, event_kind) if ( event_kind not in ( @@ -1064,14 +1119,26 @@ def _install_plan( or replacement_destinations ): self._active_targets = replacement_targets + if ( + event_kind + not in ( + ExecutionEventKind.REPLANNED, + ExecutionEventKind.INVOCATION_REVISED, + ) + or replacement_tracking_routes + ): + self._active_tracking_routes = replacement_tracking_routes self._plan = plan self._attempt_generation += 1 self._waypoint_index = 0 self._planned_scene = context.scene self._action_started_at = context.robot.timestamp - self._last_joint_command = None - self._last_joint_ids = () + self._last_tracking_frame = None self._last_command_mask.zero_() + self._tracking_violation_counts.zero_() + self._terminal_acceptance_counts.zero_() + self._terminal_started_at = None + self._terminal_pending_reported = False self._pending_effect = None self._effect_failures.zero_() self._effect_requested_at = None @@ -1158,6 +1225,56 @@ def _validate_destination_continuity( f"replacement={sorted(replacement_destinations)}.{guidance}" ) + def _validate_tracking_continuity( + self, + plan: ActionPlan, + event_kind: ExecutionEventKind, + ) -> None: + """Reject in-place replacement of feedback ownership or projection.""" + if event_kind not in ( + ExecutionEventKind.REPLANNED, + ExecutionEventKind.INVOCATION_REVISED, + ): + return + if self._plan is None: + return + previous_routes = self._active_tracking_routes + replacement_routes = self._tracking_routes(plan) + if ( + event_kind is ExecutionEventKind.REPLANNED + and not plan.commands.targets + and not replacement_routes + ): + return + if previous_routes == replacement_routes: + return + prefix = ( + "Recovery replans" + if event_kind is ExecutionEventKind.REPLANNED + else "Invocation revisions" + ) + raise ValueError( + f"{prefix} must preserve endpoint tracking source fingerprints and " + "projector routes; start a new invocation to change feedback " + "ownership." + ) + + @staticmethod + def _tracking_routes( + plan: ActionPlan, + ) -> dict[tuple[str, str, str], tuple[object, str, str]]: + """Return the complete feedback/projector route owned by one plan.""" + if plan.tracking is None or not plan.tracking.frames: + return {} + return { + setpoint.key: ( + setpoint.binding.source.source_fingerprint, + setpoint.binding.projector.projector_id, + setpoint.binding.projector.revision, + ) + for setpoint in plan.tracking.frames[0].setpoints + } + def _recover_if_needed( self, plan: ActionPlan, @@ -1180,34 +1297,49 @@ def _recover_if_needed( ExecutionEventKind.COLLISION_WORLD_CHANGED, "The collision world changed after this trajectory was planned.", ) + in_flight = plan.tracking_policy.in_flight if ( - plan.feedback_mode is ExecutionFeedbackMode.JOINT_POSITION - and self._last_joint_command is not None - and self._last_joint_ids + in_flight is not None + and self._last_tracking_frame is not None + and self._waypoint_index < plan.commands.frame_count ): - joint_ids = list(self._last_joint_ids) - tracking_error = torch.amax( - torch.abs( - self._context.robot.qpos[:, joint_ids] - - self._last_joint_command[:, joint_ids] - ), - dim=1, - ) - tracking_mask = ( - execution_mask - & self._last_command_mask - & (tracking_error > plan.recovery_policy.tracking_error_threshold) - ) - if tracking_mask.any(): - max_tracking_error = float(tracking_error[tracking_mask].amax().item()) - return self._attempt_replan( - tracking_mask, - ExecutionEventKind.TRACKING_ERROR, - "Observed joint tracking error exceeded the policy threshold " - f"(max_error={max_tracking_error:.6f}, " - "threshold=" - f"{plan.recovery_policy.tracking_error_threshold:.6f}).", + tracking_mask = execution_mask & self._last_command_mask + if ( + tracking_mask.any() + and self._context.robot.timestamp - self._action_started_at + >= in_flight.grace_period + ): + try: + accepted, valid, normalized_error = self._evaluate_tracking_frame( + self._last_tracking_frame, + in_flight.metrics, + ) + except Exception as exc: # noqa: BLE001 - fail required feedback closed + return self._fail_tracking_feedback( + tracking_mask, + "In-flight tracking feedback evaluation failed: " + f"{type(exc).__name__}: {exc}", + ) + invalid = tracking_mask & ~valid + if invalid.any(): + return self._fail_tracking_feedback( + invalid, + "Required in-flight tracking feedback was invalid.", + ) + violated = tracking_mask & valid & ~accepted + self._tracking_violation_counts[violated] += 1 + self._tracking_violation_counts[tracking_mask & ~violated] = 0 + diverged = tracking_mask & ( + self._tracking_violation_counts >= in_flight.consecutive_violations ) + if diverged.any(): + max_error = float(normalized_error[diverged].amax().item()) + return self._attempt_replan( + diverged, + ExecutionEventKind.TRACKING_DIVERGED, + "Observed in-flight tracking diverged from the commanded " + f"setpoint (max_normalized_error={max_error:.6f}).", + ) scene_mask, scene_message = self._dynamic_scene_change( plan, execution_mask, @@ -1513,76 +1645,72 @@ def _command_at( waypoint_index: int, active_mask: torch.Tensor, ) -> RuntimeCommandFrame: - """Return one frame and retain joint targets when feedback requires it.""" + """Return one frame and retain its generic typed tracking targets.""" frame = plan.commands.frames[waypoint_index] frame = frame.with_active_mask(frame.active_mask & active_mask) - if plan.feedback_mode is ExecutionFeedbackMode.JOINT_POSITION: - positions = self._context.robot.qpos.clone() - commanded_joint_ids: list[int] = [] - for command in frame.commands: - if not isinstance( - command.target, JointPositionTarget - ) or not isinstance( - command.payload, - JointPositionPayload, - ): - raise TypeError( - "joint_position feedback requires only joint-position " - "targets and payloads." - ) - joint_ids = list(command.target.joint_ids) - commanded_joint_ids.extend(joint_ids) - positions[:, joint_ids] = torch.where( - frame.active_mask[:, None], - command.payload.positions, - positions[:, joint_ids], - ) - self._last_joint_command = positions - self._last_joint_ids = tuple(commanded_joint_ids) - self._last_command_mask = frame.active_mask.clone() - else: - self._last_joint_command = None - self._last_joint_ids = () - self._last_command_mask.zero_() + self._last_tracking_frame = ( + None + if plan.tracking is None + else plan.tracking.frames[waypoint_index].snapshot() + ) + self._last_command_mask = frame.active_mask.clone() + if waypoint_index == plan.commands.frame_count - 1: + self._terminal_started_at = self._context.robot.timestamp + self._terminal_acceptance_counts.zero_() + self._terminal_pending_reported = False return frame - def _terminal_error(self, plan: ActionPlan) -> torch.Tensor: - """Return terminal error for the plan's explicit feedback contract.""" - if plan.feedback_mode is ExecutionFeedbackMode.TIMED: - return torch.zeros( - self._context.batch_size, - dtype=self._context.robot.qpos.dtype, - device=self._context.robot.qpos.device, - ) - if plan.commands.frame_count == 0: - return torch.full_like( - self._eligible, - float("inf"), - dtype=self._context.robot.qpos.dtype, - ) - errors: list[torch.Tensor] = [] - for command in plan.commands.frames[-1].commands: - if not isinstance(command.target, JointPositionTarget) or not isinstance( - command.payload, - JointPositionPayload, - ): + def _evaluate_tracking_frame( + self, + frame: TrackingFrame, + metrics: tuple[TrackingMetricCfg, ...], + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Aggregate typed endpoint predicates without mixing physical units.""" + evaluations = self._engine.tracking_runtime.evaluate_frame( + frame, + metrics, + self._context, + ) + accepted = torch.ones_like(self._eligible) + valid = torch.ones_like(self._eligible) + normalized_error = torch.zeros( + self._context.batch_size, + dtype=self._context.robot.qpos.dtype, + device=self._context.robot.qpos.device, + ) + for evaluation in evaluations.values(): + if not isinstance(evaluation, TrackingEvaluation): raise TypeError( - "joint_position feedback requires only joint-position targets " - "and payloads." - ) - joint_ids = list(command.target.joint_ids) - errors.append( - torch.abs( - self._context.robot.qpos[:, joint_ids] - command.payload.positions + "TrackingRuntime.evaluate_frame() must return " + "TrackingEvaluation values." ) + accepted &= evaluation.accepted_mask + valid &= evaluation.valid_mask + normalized_error = torch.maximum( + normalized_error, + evaluation.normalized_error.to(normalized_error.dtype), ) - if not errors: - return torch.full_like( - self._eligible, - float("inf"), - dtype=self._context.robot.qpos.dtype, + return accepted, valid, normalized_error + + def _fail_tracking_feedback( + self, + failed_mask: torch.Tensor, + message: str, + ) -> list[ExecutionEvent]: + """Fail affected rows closed when required feedback is unavailable.""" + self._eligible &= ~failed_mask + self._pending &= ~failed_mask + events = [ + self._event( + ExecutionEventKind.TRACKING_FEEDBACK_FAILED, + failed_mask, + message, ) - return torch.amax(torch.cat(errors, dim=1), dim=1) + ] + terminal_event = self._update_terminal_status() + if terminal_event is not None: + events.append(terminal_event) + return events def _dynamic_scene_change( self, diff --git a/embodichain/lab/sim/atomic_actions/invocation.py b/embodichain/lab/sim/atomic_actions/invocation.py index a600ff30b..cba5fac0f 100644 --- a/embodichain/lab/sim/atomic_actions/invocation.py +++ b/embodichain/lab/sim/atomic_actions/invocation.py @@ -29,6 +29,7 @@ from .control import ActionControlOverrides from .goals import ActionGoal from .policies import MotionPolicy, RecoveryPolicy +from .tracking import TrackingPolicy GoalT = TypeVar("GoalT", bound=ActionGoal) @@ -101,6 +102,11 @@ class ActionInvocation(Generic[GoalT, OptionsT]): motion_policy: MotionPolicy = field(default_factory=MotionPolicy) """Reusable motion-generation settings.""" + tracking_policy: TrackingPolicy = field( + default_factory=TrackingPolicy.joint_position + ) + """Typed in-flight tracking and terminal-acceptance settings.""" + recovery_policy: RecoveryPolicy = field(default_factory=RecoveryPolicy) """Bounded local execution recovery settings.""" @@ -131,6 +137,8 @@ def __post_init__(self) -> None: raise TypeError("binding must be an ActionBinding.") if not isinstance(self.motion_policy, MotionPolicy): raise TypeError("motion_policy must be a MotionPolicy.") + if not isinstance(self.tracking_policy, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") if not isinstance(self.recovery_policy, RecoveryPolicy): raise TypeError("recovery_policy must be a RecoveryPolicy.") if self.skill_options is not None and not isinstance( @@ -161,6 +169,7 @@ class ResolvedActionRequest(Generic[GoalT, OptionsT]): goal: GoalT binding: ActionBinding motion_policy: MotionPolicy + tracking_policy: TrackingPolicy recovery_policy: RecoveryPolicy skill_options: OptionsT invocation_id: str | None = None @@ -173,6 +182,8 @@ def __post_init__(self) -> None: raise TypeError("binding must be an ActionBinding.") if not isinstance(self.motion_policy, MotionPolicy): raise TypeError("motion_policy must be a MotionPolicy.") + if not isinstance(self.tracking_policy, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") if not isinstance(self.recovery_policy, RecoveryPolicy): raise TypeError("recovery_policy must be a RecoveryPolicy.") if not isinstance(self.skill_options, ActionOptions): @@ -197,6 +208,7 @@ def __post_init__(self) -> None: ), ) object.__setattr__(self, "motion_policy", deepcopy(self.motion_policy)) + object.__setattr__(self, "tracking_policy", deepcopy(self.tracking_policy)) object.__setattr__(self, "recovery_policy", deepcopy(self.recovery_policy)) object.__setattr__(self, "skill_options", deepcopy(self.skill_options)) @@ -207,6 +219,7 @@ def snapshot(self) -> ResolvedActionRequest[GoalT, OptionsT]: goal=self.goal, binding=self.binding, motion_policy=self.motion_policy, + tracking_policy=self.tracking_policy, recovery_policy=self.recovery_policy, skill_options=self.skill_options, invocation_id=self.invocation_id, diff --git a/embodichain/lab/sim/atomic_actions/plans.py b/embodichain/lab/sim/atomic_actions/plans.py index 7e86f16cf..b3c901b90 100644 --- a/embodichain/lab/sim/atomic_actions/plans.py +++ b/embodichain/lab/sim/atomic_actions/plans.py @@ -20,7 +20,6 @@ from copy import deepcopy from dataclasses import dataclass, field -from enum import Enum from types import MappingProxyType from typing import Any, Mapping, Sequence @@ -28,11 +27,15 @@ from embodichain.lab.sim.planners.utils import normalize_success_mask -from .bindings import JointPositionTarget from .effects import StateDelta from .policies import RecoveryPolicy -from .runtime_commands import JointPositionPayload, TimedCommandSequence +from .runtime_commands import TimedCommandSequence from .state import PlanningContext +from .tracking import ( + FeedbackTerminalAcceptance, + TimedTrackingSequence, + TrackingPolicy, +) def _validate_optional_trajectory_field( @@ -384,13 +387,6 @@ def __post_init__(self) -> None: ) -class ExecutionFeedbackMode(str, Enum): - """Feedback contract used to decide whether an action reached its target.""" - - JOINT_POSITION = "joint_position" - TIMED = "timed" - - @dataclass(frozen=True, slots=True) class EffectVerificationRequirement: """Explicit physical-effect verification independent of symbolic state. @@ -478,10 +474,11 @@ class ActionPlan: plan_success: torch.Tensor commands: TimedCommandSequence recovery_policy: RecoveryPolicy + tracking_policy: TrackingPolicy planned_scene_version: int planned_collision_world_revision: tuple[int, ...] diagnostics: PlannerDiagnostics - feedback_mode: ExecutionFeedbackMode = ExecutionFeedbackMode.TIMED + tracking: TimedTrackingSequence | None = None joint_trajectory: TimedTrajectory | None = None segments: tuple[TrajectorySegment, ...] = () scene_dependencies: tuple[str, ...] = () @@ -511,8 +508,8 @@ def __post_init__(self) -> None: raise ValueError("plan_success batch must match the command sequence.") if self.commands.device != self.plan_success.device: raise ValueError("plan_success and commands must share a device.") - if not isinstance(self.feedback_mode, ExecutionFeedbackMode): - raise TypeError("feedback_mode must be an ExecutionFeedbackMode.") + if not isinstance(self.tracking_policy, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") expected_target_types: dict[tuple[str, str], type[object]] | None = None expected_target_fingerprints: dict[tuple[str, str], object] | None = None for frame_index, frame in enumerate(self.commands.frames): @@ -573,101 +570,88 @@ def __post_init__(self) -> None: ) if self.joint_trajectory.positions.device != self.commands.device: raise ValueError("joint_trajectory and commands must share a device.") - if ( - self.feedback_mode is ExecutionFeedbackMode.JOINT_POSITION - and self.joint_trajectory is None + required_channels = { + metric.channel_id + for metric in ( + () + if self.tracking_policy.in_flight is None + else self.tracking_policy.in_flight.metrics + ) + } + if isinstance( + self.tracking_policy.terminal, + FeedbackTerminalAcceptance, ): - raise ValueError( - "joint_position feedback requires an owned joint_trajectory." + required_channels.update( + metric.channel_id for metric in self.tracking_policy.terminal.metrics ) - if self.feedback_mode is ExecutionFeedbackMode.JOINT_POSITION: - if bool(self.plan_success.any().item()) and self.commands.frame_count == 0: + if self.tracking is None: + if required_channels: + raise ValueError( + "Feedback tracking policies require an owned tracking sequence." + ) + else: + if not isinstance(self.tracking, TimedTrackingSequence): + raise TypeError("tracking must be a TimedTrackingSequence or None.") + if self.tracking.batch_size != self.commands.batch_size: + raise ValueError("tracking batch must match the command sequence.") + if self.tracking.frame_count != self.commands.frame_count: + raise ValueError("tracking frames must match command sequence frames.") + if not torch.equal(self.tracking.env_ids, self.commands.env_ids): + raise ValueError("tracking env_ids must match the command sequence.") + if self.tracking.device != self.commands.device: + raise ValueError("tracking and commands must share a device.") + if not required_channels: + raise ValueError( + "A tracking sequence requires an in-flight or terminal " + "feedback metric." + ) + if bool(self.plan_success.any().item()) and not self.tracking.frames: raise ValueError( - "joint_position feedback requires command frames when any " + "Feedback tracking requires command frames when any " "environment planned successfully." ) - assert self.joint_trajectory is not None - expected_destinations: dict[tuple[str, str], tuple[int, ...]] | None = None - for frame_index, frame in enumerate(self.commands.frames): - if not frame.commands: + expected_setpoint_keys: set[tuple[str, str, str]] | None = None + expected_setpoint_routes: ( + dict[ + tuple[str, str, str], + tuple[object, str, str], + ] + | None + ) = None + for frame_index, frame in enumerate(self.tracking.frames): + frame_keys = {setpoint.key for setpoint in frame.setpoints} + frame_routes = { + setpoint.key: ( + setpoint.binding.source.source_fingerprint, + setpoint.binding.projector.projector_id, + setpoint.binding.projector.revision, + ) + for setpoint in frame.setpoints + } + frame_channels = { + setpoint.binding.channel_id for setpoint in frame.setpoints + } + if frame_channels != required_channels: raise ValueError( - "joint_position feedback requires at least one endpoint " - f"command in frame {frame_index}." + "Every tracking frame must cover exactly the configured " + f"feedback channels; frame {frame_index} has " + f"{sorted(frame_channels)}, expected " + f"{sorted(required_channels)}." ) - if any( - not isinstance(command.target, JointPositionTarget) - or not isinstance(command.payload, JointPositionPayload) - for command in frame.commands - ): + if expected_setpoint_keys is None: + expected_setpoint_keys = frame_keys + expected_setpoint_routes = frame_routes + elif frame_keys != expected_setpoint_keys: raise ValueError( - "joint_position feedback accepts only JointPositionTarget " - "and JointPositionPayload commands." + "Tracking frames must preserve the same endpoint/channel " + f"set; frame {frame_index} differs from frame 0." ) - for command in frame.commands: - target = command.target - payload = command.payload - assert isinstance(target, JointPositionTarget) - assert isinstance(payload, JointPositionPayload) - if any( - joint_id >= self.joint_trajectory.robot_dof - for joint_id in target.joint_ids - ): - raise ValueError( - f"Joint target {command.destination_key} contains joint " - "IDs outside joint_trajectory robot_dof " - f"{self.joint_trajectory.robot_dof}." - ) - joint_ids = list(target.joint_ids) - expected_positions = self.joint_trajectory.positions[ - :, frame_index, joint_ids - ] - if ( - payload.positions.dtype != expected_positions.dtype - or not torch.equal(payload.positions, expected_positions) - ): - raise ValueError( - f"Joint payload positions for {command.destination_key} " - "must exactly match the corresponding joint_trajectory " - f"slice at frame {frame_index}." - ) - trajectory_velocities = self.joint_trajectory.velocities - if (payload.velocities is None) != (trajectory_velocities is None): - raise ValueError( - f"Joint payload velocities for {command.destination_key} " - "must have the same presence as joint_trajectory " - "velocities." - ) - if ( - payload.velocities is not None - and trajectory_velocities is not None - ): - expected_velocities = trajectory_velocities[ - :, frame_index, joint_ids - ] - if ( - payload.velocities.dtype != expected_velocities.dtype - or not torch.equal( - payload.velocities, - expected_velocities, - ) - ): - raise ValueError( - "Joint payload velocities for " - f"{command.destination_key} must exactly match the " - "corresponding joint_trajectory slice at frame " - f"{frame_index}." - ) - destinations = { - command.destination_key: command.target.joint_ids - for command in frame.commands - if isinstance(command.target, JointPositionTarget) - } - if expected_destinations is None: - expected_destinations = destinations - elif destinations != expected_destinations: + elif frame_routes != expected_setpoint_routes: raise ValueError( - "joint_position feedback requires a stable joint endpoint " - "set across every command frame." + "Tracking frames must preserve each endpoint/channel " + "source fingerprint and projector route; " + f"frame {frame_index} differs from frame 0." ) if not isinstance(self.recovery_policy, RecoveryPolicy): raise TypeError("recovery_policy must be a RecoveryPolicy.") @@ -753,6 +737,16 @@ def __post_init__(self) -> None: ) object.__setattr__(self, "plan_success", self.plan_success.clone()) object.__setattr__(self, "commands", self.commands.snapshot()) + object.__setattr__( + self, + "tracking_policy", + self.tracking_policy.snapshot(), + ) + object.__setattr__( + self, + "tracking", + None if self.tracking is None else self.tracking.snapshot(), + ) object.__setattr__( self, "joint_trajectory", @@ -811,10 +805,11 @@ def snapshot(self) -> ActionPlan: plan_success=self.plan_success, commands=self.commands, recovery_policy=self.recovery_policy, + tracking_policy=self.tracking_policy, planned_scene_version=self.planned_scene_version, planned_collision_world_revision=self.planned_collision_world_revision, diagnostics=self.diagnostics, - feedback_mode=self.feedback_mode, + tracking=self.tracking, joint_trajectory=self.joint_trajectory, segments=self.segments, scene_dependencies=self.scene_dependencies, @@ -915,7 +910,6 @@ def segment(self, action_index: int, name: str) -> TrajectorySegment: "ActionPlan", "CompiledTrajectory", "EffectVerificationRequirement", - "ExecutionFeedbackMode", "PlannerDiagnostics", "TimedTrajectory", "TrajectorySegment", diff --git a/embodichain/lab/sim/atomic_actions/policies.py b/embodichain/lab/sim/atomic_actions/policies.py index 5ef36d4c6..c9d536b4d 100644 --- a/embodichain/lab/sim/atomic_actions/policies.py +++ b/embodichain/lab/sim/atomic_actions/policies.py @@ -152,9 +152,6 @@ class RecoveryPolicy: max_action_retries: int = 2 """Maximum whole-action retries after planning, execution, or effect failure.""" - tracking_error_threshold: float = 0.05 - """Joint tracking-error threshold in radians.""" - goal_translation_threshold: float = 0.02 """Dynamic-goal translation threshold in metres.""" @@ -170,7 +167,6 @@ def __post_init__(self) -> None: if self.max_action_retries < 0: raise ValueError("max_action_retries must be non-negative.") threshold_fields = ( - "tracking_error_threshold", "goal_translation_threshold", "goal_rotation_threshold", "action_timeout", diff --git a/embodichain/lab/sim/atomic_actions/runtime.py b/embodichain/lab/sim/atomic_actions/runtime.py index 13228ff37..868d31bc6 100644 --- a/embodichain/lab/sim/atomic_actions/runtime.py +++ b/embodichain/lab/sim/atomic_actions/runtime.py @@ -33,6 +33,14 @@ DisjointSlotEndpoints, SkillBindingContract, ) +from .tracking import ( + JOINT_POSITION_CHANNEL, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, + TrackingFeedbackSourceRef, + TrackingProjectorRef, + TrackingRuntime, +) if TYPE_CHECKING: from embodichain.lab.sim.objects import Robot @@ -46,11 +54,18 @@ def __init__( self, motion_generator: MotionGenerator, control_profiles: Mapping[str, ControlPartCommandProfile] | None = None, + tracking_runtime: TrackingRuntime | None = None, ) -> None: self._motion_generator = motion_generator self._robot: Robot = motion_generator.robot self._device = resolve_runtime_device(motion_generator.device) self._binding_owner_id = uuid4().hex + if tracking_runtime is not None and not isinstance( + tracking_runtime, + TrackingRuntime, + ): + raise TypeError("tracking_runtime must be a TrackingRuntime or None.") + self._tracking_runtime = tracking_runtime or TrackingRuntime.with_builtins() self._control_profiles = self._snapshot_control_profiles( {} if control_profiles is None else control_profiles ) @@ -75,6 +90,11 @@ def binding_owner_id(self) -> str: """Return the opaque identity required by this engine's bindings.""" return self._binding_owner_id + @property + def tracking_runtime(self) -> TrackingRuntime: + """Return the engine-owned typed tracking runtime.""" + return self._tracking_runtime + @property def control_profiles(self) -> Mapping[str, ControlPartCommandProfile]: """Return owned direct-core command profiles by control-part name.""" @@ -232,14 +252,32 @@ def bind_control_parts( f"Endpoint {slot_id}.{endpoint_id} requires command {name!r} " f"of type {command_type.__name__}." ) + target = JointPositionTarget(control_part, joint_ids) resolved.append( EndpointBinding( slot_id=slot_id, endpoint_id=endpoint_id, resource_id=f"direct.{slot_id}", adapter_id="control_part", - target=JointPositionTarget(control_part, joint_ids), + target=target, task_state_key=resolved_task_state_keys[slot_id], + tracking_channels={ + JOINT_POSITION_CHANNEL: EndpointTrackingChannelBinding( + channel_id=JOINT_POSITION_CHANNEL, + source=TrackingFeedbackSourceRef( + provider_id="planning_context.robot", + revision="1", + address=EndpointTrackingFeedbackAddress( + target=target, + channel_id=JOINT_POSITION_CHANNEL, + ), + ), + projector=TrackingProjectorRef( + projector_id="joint_position_payload", + revision="1", + ), + ) + }, capabilities=requirement.capabilities, commands=commands, claim_tokens=frozenset({f"robot.control_part:{control_part}"}), diff --git a/embodichain/lab/sim/atomic_actions/state.py b/embodichain/lab/sim/atomic_actions/state.py index 578a43d86..d24812628 100644 --- a/embodichain/lab/sim/atomic_actions/state.py +++ b/embodichain/lab/sim/atomic_actions/state.py @@ -511,6 +511,38 @@ def __post_init__(self) -> None: raise ValueError("RobotObservation.qeffort must match qpos shape.") if self.qeffort.device != self.qpos.device: raise ValueError("RobotObservation.qeffort must share the qpos device.") + if self.root_pose is not None: + if not isinstance(self.root_pose, torch.Tensor): + raise TypeError("RobotObservation.root_pose must be a tensor or None.") + if self.root_pose.shape != (self.qpos.shape[0], 4, 4): + raise ValueError( + "RobotObservation.root_pose must have shape " + f"({self.qpos.shape[0]}, 4, 4)." + ) + if not self.root_pose.is_floating_point(): + raise TypeError("RobotObservation.root_pose must be floating point.") + if self.root_pose.device != self.qpos.device: + raise ValueError( + "RobotObservation.root_pose must share the qpos device." + ) + if not torch.isfinite(self.root_pose).all(): + raise ValueError("RobotObservation.root_pose must be finite.") + if self.root_twist is not None: + if not isinstance(self.root_twist, torch.Tensor): + raise TypeError("RobotObservation.root_twist must be a tensor or None.") + if self.root_twist.shape != (self.qpos.shape[0], 6): + raise ValueError( + "RobotObservation.root_twist must have shape " + f"({self.qpos.shape[0]}, 6)." + ) + if not self.root_twist.is_floating_point(): + raise TypeError("RobotObservation.root_twist must be floating point.") + if self.root_twist.device != self.qpos.device: + raise ValueError( + "RobotObservation.root_twist must share the qpos device." + ) + if not torch.isfinite(self.root_twist).all(): + raise ValueError("RobotObservation.root_twist must be finite.") object.__setattr__(self, "qpos", self.qpos.clone()) object.__setattr__(self, "qvel", self.qvel.clone()) if self.qeffort is not None: diff --git a/embodichain/lab/sim/atomic_actions/tracking.py b/embodichain/lab/sim/atomic_actions/tracking.py new file mode 100644 index 000000000..2a9e12acc --- /dev/null +++ b/embodichain/lab/sim/atomic_actions/tracking.py @@ -0,0 +1,1210 @@ +# ---------------------------------------------------------------------------- +# 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. +# ---------------------------------------------------------------------------- + +"""Typed, transport-neutral tracking contracts for atomic-action execution.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from copy import deepcopy +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import TYPE_CHECKING, ClassVar, Hashable, Iterable, Mapping, Protocol + +import torch + +if TYPE_CHECKING: + from .bindings import RuntimeEndpointTarget + from .runtime_commands import EndpointCommand + from .state import PlanningContext + + +TrackingChannelId = str +"""Open string identifier for one typed endpoint-feedback channel.""" + +JOINT_POSITION_CHANNEL: TrackingChannelId = "joint.position" +BASE_POSE_CHANNEL: TrackingChannelId = "base.pose" +WHOLE_BODY_POSE_CHANNEL: TrackingChannelId = "whole_body.pose" + + +def _identifier(value: str, *, field_name: str) -> str: + if not isinstance(value, str) or not value or value != value.strip(): + raise ValueError(f"{field_name} must be a non-empty trimmed string.") + return value + + +def _positive_float(value: float, *, field_name: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError(f"{field_name} must be a number.") + normalized = float(value) + if not torch.isfinite(torch.tensor(normalized)).item() or normalized <= 0.0: + raise ValueError(f"{field_name} must be finite and positive.") + return normalized + + +def _non_negative_float(value: float, *, field_name: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError(f"{field_name} must be a number.") + normalized = float(value) + if not torch.isfinite(torch.tensor(normalized)).item() or normalized < 0.0: + raise ValueError(f"{field_name} must be finite and non-negative.") + return normalized + + +def _tensor(value: torch.Tensor, *, field_name: str, dimensions: int) -> torch.Tensor: + if not isinstance(value, torch.Tensor): + raise TypeError(f"{field_name} must be a torch.Tensor.") + if value.dim() != dimensions or any(size < 1 for size in value.shape): + raise ValueError(f"{field_name} must be a non-empty {dimensions}-D tensor.") + if not torch.is_floating_point(value) or not torch.isfinite(value).all().item(): + raise ValueError(f"{field_name} must contain finite floating-point values.") + return value.clone() + + +class TrackingFeedbackAddress(ABC): + """Immutable address understood by one tracking-feedback provider.""" + + @property + @abstractmethod + def address_fingerprint(self) -> Hashable: + """Return a stable, hashable address identity.""" + + def snapshot(self) -> TrackingFeedbackAddress: + """Return an independently owned address snapshot.""" + return deepcopy(self) + + +@dataclass(frozen=True, slots=True) +class EndpointTrackingFeedbackAddress(TrackingFeedbackAddress): + """Feedback address for one runtime endpoint and open tracking channel.""" + + target: RuntimeEndpointTarget + channel_id: TrackingChannelId + + def __post_init__(self) -> None: + from .bindings import RuntimeEndpointTarget + + if not isinstance(self.target, RuntimeEndpointTarget): + raise TypeError("target must be a RuntimeEndpointTarget.") + snapshot = self.target.snapshot() + if type(snapshot) is not type(self.target) or snapshot is self.target: + raise TypeError("RuntimeEndpointTarget.snapshot() must own a new value.") + if snapshot.address_fingerprint != self.target.address_fingerprint: + raise ValueError("Target snapshot must preserve its address fingerprint.") + _identifier(self.channel_id, field_name="channel_id") + object.__setattr__(self, "target", snapshot) + + @property + def address_fingerprint(self) -> Hashable: + """Return the endpoint- and channel-scoped address identity.""" + return self.target.address_fingerprint, self.channel_id + + +@dataclass(frozen=True, slots=True) +class TrackingFeedbackSourceRef: + """Versioned provider route plus one immutable feedback address.""" + + provider_id: str + revision: str + address: TrackingFeedbackAddress + + def __post_init__(self) -> None: + _identifier(self.provider_id, field_name="provider_id") + _identifier(self.revision, field_name="revision") + if not isinstance(self.address, TrackingFeedbackAddress): + raise TypeError("address must be a TrackingFeedbackAddress.") + snapshot = self.address.snapshot() + if type(snapshot) is not type(self.address) or snapshot is self.address: + raise TypeError("TrackingFeedbackAddress.snapshot() must own a new value.") + if snapshot.address_fingerprint != self.address.address_fingerprint: + raise ValueError("Address snapshot must preserve its fingerprint.") + hash(snapshot.address_fingerprint) + object.__setattr__(self, "address", snapshot) + + @property + def source_fingerprint(self) -> Hashable: + """Return the exact versioned source identity.""" + return self.provider_id, self.revision, self.address.address_fingerprint + + def snapshot(self) -> TrackingFeedbackSourceRef: + """Return an independently owned source reference.""" + return TrackingFeedbackSourceRef(self.provider_id, self.revision, self.address) + + +@dataclass(frozen=True, slots=True) +class TrackingProjectorRef: + """Exact version of a command-to-tracking-state projector.""" + + projector_id: str + revision: str + + def __post_init__(self) -> None: + _identifier(self.projector_id, field_name="projector_id") + _identifier(self.revision, field_name="revision") + + def snapshot(self) -> TrackingProjectorRef: + """Return an independently owned projector route.""" + return TrackingProjectorRef(self.projector_id, self.revision) + + +@dataclass(frozen=True, slots=True) +class EndpointTrackingChannelBinding: + """Resolved source and projector for one endpoint tracking channel.""" + + channel_id: TrackingChannelId + source: TrackingFeedbackSourceRef + projector: TrackingProjectorRef + + def __post_init__(self) -> None: + _identifier(self.channel_id, field_name="channel_id") + if not isinstance(self.source, TrackingFeedbackSourceRef): + raise TypeError("source must be a TrackingFeedbackSourceRef.") + if not isinstance(self.projector, TrackingProjectorRef): + raise TypeError("projector must be a TrackingProjectorRef.") + address = self.source.address + if isinstance(address, EndpointTrackingFeedbackAddress): + if address.channel_id != self.channel_id: + raise ValueError("Binding and feedback-address channels must match.") + object.__setattr__(self, "source", self.source.snapshot()) + object.__setattr__(self, "projector", self.projector.snapshot()) + + def snapshot(self) -> EndpointTrackingChannelBinding: + """Return an independently owned channel binding.""" + return EndpointTrackingChannelBinding( + self.channel_id, self.source, self.projector + ) + + @property + def route_fingerprint(self) -> tuple[str, Hashable, str, str]: + """Return the exact channel, source, and projector route identity.""" + return ( + self.channel_id, + self.source.source_fingerprint, + self.projector.projector_id, + self.projector.revision, + ) + + +class TrackingState(ABC): + """Immutable-by-ownership typed desired or observed tracking state.""" + + channel_id: ClassVar[TrackingChannelId] + + @property + @abstractmethod + def batch_size(self) -> int: + """Return the represented environment count.""" + + @property + @abstractmethod + def device(self) -> torch.device: + """Return the tensor device.""" + + @abstractmethod + def snapshot(self) -> TrackingState: + """Return an independently owned state snapshot.""" + + +@dataclass(frozen=True, slots=True, eq=False) +class JointPositionTrackingState(TrackingState): + """Batched joint positions with shape ``(B, D)``.""" + + channel_id: ClassVar[str] = JOINT_POSITION_CHANNEL + positions: torch.Tensor + + def __post_init__(self) -> None: + object.__setattr__( + self, + "positions", + _tensor(self.positions, field_name="positions", dimensions=2), + ) + + @property + def batch_size(self) -> int: + return int(self.positions.shape[0]) + + @property + def device(self) -> torch.device: + return self.positions.device + + def snapshot(self) -> JointPositionTrackingState: + return JointPositionTrackingState(self.positions) + + +@dataclass(frozen=True, slots=True, eq=False) +class PoseTrackingState(TrackingState): + """Batched homogeneous poses with shape ``(B, 4, 4)``.""" + + channel_id: ClassVar[str] = BASE_POSE_CHANNEL + poses: torch.Tensor + + def __post_init__(self) -> None: + poses = _tensor(self.poses, field_name="poses", dimensions=3) + if poses.shape[1:] != (4, 4): + raise ValueError("poses must have shape (batch_size, 4, 4).") + object.__setattr__(self, "poses", poses) + + @property + def batch_size(self) -> int: + return int(self.poses.shape[0]) + + @property + def device(self) -> torch.device: + return self.poses.device + + def snapshot(self) -> PoseTrackingState: + return PoseTrackingState(self.poses) + + +@dataclass(frozen=True, slots=True, eq=False) +class WholeBodyPoseTrackingState(TrackingState): + """Batched base poses and joint positions for whole-body tracking.""" + + channel_id: ClassVar[str] = WHOLE_BODY_POSE_CHANNEL + root_poses: torch.Tensor + joint_positions: torch.Tensor + + def __post_init__(self) -> None: + root_poses = _tensor(self.root_poses, field_name="root_poses", dimensions=3) + joints = _tensor( + self.joint_positions, + field_name="joint_positions", + dimensions=2, + ) + if root_poses.shape[1:] != (4, 4): + raise ValueError("root_poses must have shape (batch_size, 4, 4).") + if root_poses.shape[0] != joints.shape[0]: + raise ValueError("root_poses and joint_positions batches must match.") + if root_poses.device != joints.device: + raise ValueError("root_poses and joint_positions must share a device.") + object.__setattr__(self, "root_poses", root_poses) + object.__setattr__(self, "joint_positions", joints) + + @property + def batch_size(self) -> int: + return int(self.root_poses.shape[0]) + + @property + def device(self) -> torch.device: + return self.root_poses.device + + def snapshot(self) -> WholeBodyPoseTrackingState: + return WholeBodyPoseTrackingState(self.root_poses, self.joint_positions) + + +class TrackingMetricCfg(ABC): + """Immutable tolerance configuration dispatched by exact metric ID/revision.""" + + metric_id: ClassVar[str] + revision: ClassVar[str] = "1" + channel_id: ClassVar[TrackingChannelId] + + def snapshot(self) -> TrackingMetricCfg: + """Return an independently owned metric configuration.""" + return deepcopy(self) + + +@dataclass(frozen=True, slots=True) +class JointPositionTrackingMetric(TrackingMetricCfg): + """Maximum absolute joint-error tolerance.""" + + metric_id: ClassVar[str] = "joint.max_abs" + channel_id: ClassVar[str] = JOINT_POSITION_CHANNEL + tolerance: float = 0.05 + + def __post_init__(self) -> None: + object.__setattr__( + self, "tolerance", _positive_float(self.tolerance, field_name="tolerance") + ) + + +@dataclass(frozen=True, slots=True) +class PoseTrackingMetric(TrackingMetricCfg): + """Independent translation and rotation tolerances for base pose.""" + + metric_id: ClassVar[str] = "pose.se3" + channel_id: ClassVar[str] = BASE_POSE_CHANNEL + translation_tolerance: float = 0.02 + rotation_tolerance: float = 0.05 + + def __post_init__(self) -> None: + object.__setattr__( + self, + "translation_tolerance", + _positive_float( + self.translation_tolerance, field_name="translation_tolerance" + ), + ) + object.__setattr__( + self, + "rotation_tolerance", + _positive_float(self.rotation_tolerance, field_name="rotation_tolerance"), + ) + + +@dataclass(frozen=True, slots=True) +class WholeBodyPoseTrackingMetric(TrackingMetricCfg): + """Independent base-pose and joint-position tolerances.""" + + metric_id: ClassVar[str] = "whole_body.pose" + channel_id: ClassVar[str] = WHOLE_BODY_POSE_CHANNEL + translation_tolerance: float = 0.02 + rotation_tolerance: float = 0.05 + joint_position_tolerance: float = 0.05 + + def __post_init__(self) -> None: + for field_name in ( + "translation_tolerance", + "rotation_tolerance", + "joint_position_tolerance", + ): + object.__setattr__( + self, + field_name, + _positive_float(getattr(self, field_name), field_name=field_name), + ) + + +def _own_metrics( + metrics: Iterable[TrackingMetricCfg], *, field_name: str +) -> tuple[TrackingMetricCfg, ...]: + snapshots: list[TrackingMetricCfg] = [] + channels: set[str] = set() + for metric in metrics: + if not isinstance(metric, TrackingMetricCfg): + raise TypeError(f"{field_name} must contain TrackingMetricCfg values.") + _identifier(metric.metric_id, field_name=f"{field_name}.metric_id") + _identifier(metric.revision, field_name=f"{field_name}.revision") + _identifier(metric.channel_id, field_name=f"{field_name}.channel_id") + if metric.channel_id in channels: + raise ValueError( + f"{field_name} contains duplicate channel {metric.channel_id!r}." + ) + snapshot = metric.snapshot() + if type(snapshot) is not type(metric) or snapshot is metric: + raise TypeError("TrackingMetricCfg.snapshot() must own a same-type value.") + channels.add(metric.channel_id) + snapshots.append(snapshot) + if not snapshots: + raise ValueError(f"{field_name} must contain at least one metric.") + return tuple(snapshots) + + +@dataclass(frozen=True, slots=True) +class InFlightTrackingPolicy: + """Feedback checks used while a command sequence is still in flight.""" + + metrics: tuple[TrackingMetricCfg, ...] + consecutive_violations: int = 1 + grace_period: float = 0.0 + + def __post_init__(self) -> None: + object.__setattr__( + self, "metrics", _own_metrics(self.metrics, field_name="metrics") + ) + if ( + not isinstance(self.consecutive_violations, int) + or isinstance(self.consecutive_violations, bool) + or self.consecutive_violations < 1 + ): + raise ValueError("consecutive_violations must be a positive integer.") + object.__setattr__( + self, + "grace_period", + _non_negative_float(self.grace_period, field_name="grace_period"), + ) + + def snapshot(self) -> InFlightTrackingPolicy: + return InFlightTrackingPolicy( + self.metrics, self.consecutive_violations, self.grace_period + ) + + +@dataclass(frozen=True, slots=True) +class FeedbackTerminalAcceptance: + """Terminal acceptance proven by typed endpoint feedback.""" + + metrics: tuple[TrackingMetricCfg, ...] + settle_timeout: float = 0.0 + consecutive_acceptances: int = 1 + + def __post_init__(self) -> None: + object.__setattr__( + self, "metrics", _own_metrics(self.metrics, field_name="metrics") + ) + object.__setattr__( + self, + "settle_timeout", + _non_negative_float(self.settle_timeout, field_name="settle_timeout"), + ) + if ( + not isinstance(self.consecutive_acceptances, int) + or isinstance(self.consecutive_acceptances, bool) + or self.consecutive_acceptances < 1 + ): + raise ValueError("consecutive_acceptances must be a positive integer.") + + def snapshot(self) -> FeedbackTerminalAcceptance: + return FeedbackTerminalAcceptance( + self.metrics, self.settle_timeout, self.consecutive_acceptances + ) + + +@dataclass(frozen=True, slots=True) +class TimedTerminalAcceptance: + """Explicit terminal acceptance without endpoint feedback.""" + + settle_duration: float = 0.0 + + def __post_init__(self) -> None: + object.__setattr__( + self, + "settle_duration", + _non_negative_float(self.settle_duration, field_name="settle_duration"), + ) + + def snapshot(self) -> TimedTerminalAcceptance: + return TimedTerminalAcceptance(self.settle_duration) + + +TerminalAcceptance = FeedbackTerminalAcceptance | TimedTerminalAcceptance + + +@dataclass(frozen=True, slots=True) +class TrackingPolicy: + """Independent in-flight recovery signal and terminal acceptance contract.""" + + in_flight: InFlightTrackingPolicy | None + terminal: TerminalAcceptance + + def __post_init__(self) -> None: + if self.in_flight is not None and not isinstance( + self.in_flight, InFlightTrackingPolicy + ): + raise TypeError("in_flight must be InFlightTrackingPolicy or None.") + if not isinstance( + self.terminal, (FeedbackTerminalAcceptance, TimedTerminalAcceptance) + ): + raise TypeError("terminal must be a terminal-acceptance contract.") + if self.in_flight is not None: + object.__setattr__(self, "in_flight", self.in_flight.snapshot()) + object.__setattr__(self, "terminal", self.terminal.snapshot()) + in_flight = self.in_flight + terminal = self.terminal + if in_flight is not None and isinstance(terminal, FeedbackTerminalAcceptance): + in_flight_by_channel = { + metric.channel_id: metric for metric in in_flight.metrics + } + for terminal_metric in terminal.metrics: + in_flight_metric = in_flight_by_channel.get(terminal_metric.channel_id) + if in_flight_metric is None: + continue + if ( + in_flight_metric.metric_id != terminal_metric.metric_id + or in_flight_metric.revision != terminal_metric.revision + or type(in_flight_metric) is not type(terminal_metric) + ): + raise ValueError( + "In-flight and terminal metrics sharing a channel must " + "use the same exact metric ID, revision, and type." + ) + + def snapshot(self) -> TrackingPolicy: + return TrackingPolicy(self.in_flight, self.terminal) + + @classmethod + def timed(cls, *, settle_duration: float = 0.0) -> TrackingPolicy: + """Create an explicit time-only terminal contract with no tracking.""" + return cls( + in_flight=None, + terminal=TimedTerminalAcceptance(settle_duration=settle_duration), + ) + + @classmethod + def joint_position( + cls, + *, + in_flight_max_abs_error: float = 0.05, + terminal_max_abs_error: float = 0.05, + terminal_settle_timeout: float = 0.5, + consecutive_violations: int = 1, + consecutive_acceptances: int = 1, + grace_period: float = 0.0, + ) -> TrackingPolicy: + """Create the built-in joint-position tracking and acceptance contract.""" + return cls( + in_flight=InFlightTrackingPolicy( + metrics=(JointPositionTrackingMetric(in_flight_max_abs_error),), + consecutive_violations=consecutive_violations, + grace_period=grace_period, + ), + terminal=FeedbackTerminalAcceptance( + metrics=(JointPositionTrackingMetric(terminal_max_abs_error),), + settle_timeout=terminal_settle_timeout, + consecutive_acceptances=consecutive_acceptances, + ), + ) + + +@dataclass(frozen=True, slots=True, eq=False) +class TrackingSetpoint: + """One endpoint-local desired state and its typed feedback route.""" + + endpoint_key: tuple[str, str] + binding: EndpointTrackingChannelBinding + desired: TrackingState + + def __post_init__(self) -> None: + if not isinstance(self.endpoint_key, tuple) or len(self.endpoint_key) != 2: + raise TypeError("endpoint_key must be a (slot_id, endpoint_id) tuple.") + _identifier(self.endpoint_key[0], field_name="endpoint_key.slot_id") + _identifier(self.endpoint_key[1], field_name="endpoint_key.endpoint_id") + if not isinstance(self.binding, EndpointTrackingChannelBinding): + raise TypeError("binding must be an EndpointTrackingChannelBinding.") + if not isinstance(self.desired, TrackingState): + raise TypeError("desired must be a TrackingState.") + if self.binding.channel_id != self.desired.channel_id: + raise ValueError("Binding and desired-state channels must match.") + desired = self.desired.snapshot() + if type(desired) is not type(self.desired) or desired is self.desired: + raise TypeError("TrackingState.snapshot() must own a same-type value.") + object.__setattr__(self, "binding", self.binding.snapshot()) + object.__setattr__(self, "desired", desired) + + @property + def key(self) -> tuple[str, str, str]: + return self.endpoint_key[0], self.endpoint_key[1], self.binding.channel_id + + def snapshot(self) -> TrackingSetpoint: + return TrackingSetpoint(self.endpoint_key, self.binding, self.desired) + + +@dataclass(frozen=True, slots=True) +class TrackingFrame: + """Desired endpoint states associated with one command frame.""" + + setpoints: tuple[TrackingSetpoint, ...] = () + + def __post_init__(self) -> None: + snapshots: list[TrackingSetpoint] = [] + keys: set[tuple[str, str, str]] = set() + for setpoint in self.setpoints: + if not isinstance(setpoint, TrackingSetpoint): + raise TypeError("setpoints must contain TrackingSetpoint values.") + if setpoint.key in keys: + raise ValueError(f"Duplicate tracking setpoint {setpoint.key!r}.") + keys.add(setpoint.key) + snapshots.append(setpoint.snapshot()) + object.__setattr__(self, "setpoints", tuple(snapshots)) + + def snapshot(self) -> TrackingFrame: + return TrackingFrame(self.setpoints) + + +@dataclass(frozen=True, slots=True) +class TimedTrackingSequence: + """Tracking frames aligned by index with an authoritative command sequence.""" + + env_ids: torch.Tensor + frames: tuple[TrackingFrame, ...] + + def __post_init__(self) -> None: + if not isinstance(self.env_ids, torch.Tensor): + raise TypeError("env_ids must be a torch.Tensor.") + if ( + self.env_ids.dtype != torch.long + or self.env_ids.dim() != 1 + or self.env_ids.numel() < 1 + ): + raise ValueError("env_ids must be a non-empty one-dimensional long tensor.") + if torch.unique(self.env_ids).numel() != self.env_ids.numel(): + raise ValueError("env_ids must be unique.") + frames: list[TrackingFrame] = [] + for frame in self.frames: + if not isinstance(frame, TrackingFrame): + raise TypeError("frames must contain TrackingFrame values.") + snapshot = frame.snapshot() + for setpoint in snapshot.setpoints: + if setpoint.desired.batch_size != self.env_ids.numel(): + raise ValueError("Every setpoint batch must match env_ids.") + if setpoint.desired.device != self.env_ids.device: + raise ValueError("Every setpoint and env_ids must share a device.") + frames.append(snapshot) + object.__setattr__(self, "env_ids", self.env_ids.clone()) + object.__setattr__(self, "frames", tuple(frames)) + + @property + def batch_size(self) -> int: + """Return the represented environment count.""" + return int(self.env_ids.numel()) + + @property + def device(self) -> torch.device: + """Return the sequence tensor device.""" + return self.env_ids.device + + @property + def frame_count(self) -> int: + """Return the number of command-aligned tracking frames.""" + return len(self.frames) + + def snapshot(self) -> TimedTrackingSequence: + return TimedTrackingSequence(self.env_ids, self.frames) + + +@dataclass(frozen=True, slots=True, eq=False) +class TrackingFeedbackBatch: + """One synchronized typed observation from an exact feedback source.""" + + source: TrackingFeedbackSourceRef + state: TrackingState + valid_mask: torch.Tensor + timestamp: float + + def __post_init__(self) -> None: + if not isinstance(self.source, TrackingFeedbackSourceRef): + raise TypeError("source must be a TrackingFeedbackSourceRef.") + if not isinstance(self.state, TrackingState): + raise TypeError("state must be a TrackingState.") + if self.valid_mask.dtype != torch.bool or self.valid_mask.shape != ( + self.state.batch_size, + ): + raise ValueError("valid_mask must have shape (batch_size,) and bool dtype.") + if self.valid_mask.device != self.state.device: + raise ValueError("valid_mask and state must share a device.") + object.__setattr__(self, "source", self.source.snapshot()) + object.__setattr__(self, "state", self.state.snapshot()) + object.__setattr__(self, "valid_mask", self.valid_mask.clone()) + object.__setattr__( + self, + "timestamp", + _non_negative_float(self.timestamp, field_name="timestamp"), + ) + + def snapshot(self) -> TrackingFeedbackBatch: + return TrackingFeedbackBatch( + self.source, self.state, self.valid_mask, self.timestamp + ) + + +@dataclass(frozen=True, slots=True, eq=False) +class TrackingEvaluation: + """Per-row metric result with unit-preserving component errors.""" + + channel_id: TrackingChannelId + accepted_mask: torch.Tensor + valid_mask: torch.Tensor + normalized_error: torch.Tensor + component_errors: Mapping[str, torch.Tensor] = field(default_factory=dict) + + def __post_init__(self) -> None: + _identifier(self.channel_id, field_name="channel_id") + expected = self.accepted_mask.shape + if self.accepted_mask.dtype != torch.bool or self.accepted_mask.dim() != 1: + raise ValueError("accepted_mask must be a one-dimensional bool tensor.") + if self.valid_mask.dtype != torch.bool or self.valid_mask.shape != expected: + raise ValueError("valid_mask must match accepted_mask with bool dtype.") + if self.normalized_error.shape != expected or not torch.is_floating_point( + self.normalized_error + ): + raise ValueError("normalized_error must be a floating tensor per row.") + if not ( + self.accepted_mask.device + == self.valid_mask.device + == self.normalized_error.device + ): + raise ValueError("Evaluation tensors must share a device.") + components: dict[str, torch.Tensor] = {} + for name, value in self.component_errors.items(): + _identifier(name, field_name="component_errors key") + if value.shape != expected or value.device != self.normalized_error.device: + raise ValueError("Every component error must be a per-row tensor.") + components[name] = value.clone() + object.__setattr__(self, "accepted_mask", self.accepted_mask.clone()) + object.__setattr__(self, "valid_mask", self.valid_mask.clone()) + object.__setattr__(self, "normalized_error", self.normalized_error.clone()) + object.__setattr__(self, "component_errors", MappingProxyType(components)) + + def snapshot(self) -> TrackingEvaluation: + return TrackingEvaluation( + self.channel_id, + self.accepted_mask, + self.valid_mask, + self.normalized_error, + self.component_errors, + ) + + +class TrackingFeedbackProvider(Protocol): + """Versioned live port that reads one exact tracking source.""" + + provider_id: str + revision: str + + def observe( + self, source: TrackingFeedbackSourceRef, context: PlanningContext + ) -> TrackingFeedbackBatch: + """Read one synchronized typed feedback batch.""" + + +class TrackingCommandProjector(Protocol): + """Versioned pure projector from an endpoint command to desired state.""" + + projector_id: str + revision: str + + def project( + self, command: EndpointCommand, binding: EndpointTrackingChannelBinding + ) -> TrackingState: + """Project one command into the binding's desired tracking channel.""" + + +class TrackingMetricEvaluator(Protocol): + """Versioned evaluator for one exact metric configuration type.""" + + metric_id: str + revision: str + metric_type: type[TrackingMetricCfg] + + def evaluate( + self, + desired: TrackingState, + observed: TrackingState, + valid_mask: torch.Tensor, + metric: TrackingMetricCfg, + ) -> TrackingEvaluation: + """Evaluate a desired and observed batch row by row.""" + + +class _ExactRegistry: + __slots__ = ("_values", "_kind") + + def __init__(self, values: Iterable[object], *, kind: str) -> None: + normalized: dict[tuple[str, str], object] = {} + for value in values: + identifier = _identifier( + getattr(value, f"{kind}_id"), field_name=f"{kind}_id" + ) + revision = _identifier(getattr(value, "revision"), field_name="revision") + key = identifier, revision + if key in normalized: + raise ValueError(f"Duplicate {kind} registration {key!r}.") + normalized[key] = value + self._values = MappingProxyType(normalized) + self._kind = kind + + @property + def values(self) -> Mapping[tuple[str, str], object]: + return self._values + + def _resolve(self, identifier: str, revision: str) -> object: + key = identifier, revision + try: + return self._values[key] + except KeyError as exc: + raise KeyError(f"Unknown {self._kind} registration {key!r}.") from exc + + +class TrackingFeedbackProviderRegistry(_ExactRegistry): + """Immutable exact-version feedback-provider registry.""" + + def __init__(self, providers: Iterable[TrackingFeedbackProvider] = ()) -> None: + super().__init__(providers, kind="provider") + + def resolve(self, source: TrackingFeedbackSourceRef) -> TrackingFeedbackProvider: + return self._resolve(source.provider_id, source.revision) # type: ignore[return-value] + + +class TrackingProjectorRegistry(_ExactRegistry): + """Immutable exact-version command-projector registry.""" + + def __init__(self, projectors: Iterable[TrackingCommandProjector] = ()) -> None: + super().__init__(projectors, kind="projector") + + def resolve(self, route: TrackingProjectorRef) -> TrackingCommandProjector: + return self._resolve(route.projector_id, route.revision) # type: ignore[return-value] + + +class TrackingEvaluatorRegistry(_ExactRegistry): + """Immutable exact-version metric-evaluator registry.""" + + def __init__(self, evaluators: Iterable[TrackingMetricEvaluator] = ()) -> None: + super().__init__(evaluators, kind="metric") + + def resolve(self, metric: TrackingMetricCfg) -> TrackingMetricEvaluator: + evaluator = self._resolve(metric.metric_id, metric.revision) + if type(metric) is not evaluator.metric_type: # type: ignore[attr-defined] + raise TypeError( + f"Metric {metric.metric_id!r} requires " + f"{evaluator.metric_type.__name__}." # type: ignore[attr-defined] + ) + return evaluator # type: ignore[return-value] + + +class PlanningContextTrackingFeedbackProvider: + """Built-in provider backed by :class:`PlanningContext.robot`.""" + + provider_id = "planning_context.robot" + revision = "1" + + def observe( + self, source: TrackingFeedbackSourceRef, context: PlanningContext + ) -> TrackingFeedbackBatch: + from .bindings import JointPositionTarget + from .state import PlanningContext + + if not isinstance(context, PlanningContext): + raise TypeError("context must be a PlanningContext.") + address = source.address + if not isinstance(address, EndpointTrackingFeedbackAddress): + raise TypeError( + "Built-in provider requires EndpointTrackingFeedbackAddress." + ) + target = address.target + if address.channel_id == JOINT_POSITION_CHANNEL: + if not isinstance(target, JointPositionTarget): + raise TypeError("joint.position requires a JointPositionTarget.") + state: TrackingState = JointPositionTrackingState( + context.robot.qpos[:, target.joint_ids] + ) + elif address.channel_id == BASE_POSE_CHANNEL: + if context.robot.root_pose is None: + raise RuntimeError("RobotObservation.root_pose is unavailable.") + state = PoseTrackingState(context.robot.root_pose) + elif address.channel_id == WHOLE_BODY_POSE_CHANNEL: + if context.robot.root_pose is None: + raise RuntimeError("RobotObservation.root_pose is unavailable.") + joints = ( + context.robot.qpos[:, target.joint_ids] + if isinstance(target, JointPositionTarget) + else context.robot.qpos + ) + state = WholeBodyPoseTrackingState(context.robot.root_pose, joints) + else: + raise KeyError( + f"Unsupported built-in tracking channel {address.channel_id!r}." + ) + return TrackingFeedbackBatch( + source=source, + state=state, + valid_mask=torch.ones( + context.batch_size, dtype=torch.bool, device=state.device + ), + timestamp=context.robot.timestamp, + ) + + +class JointPositionTrackingProjector: + """Built-in projector for joint-position endpoint commands.""" + + projector_id = "joint_position_payload" + revision = "1" + + def project( + self, command: EndpointCommand, binding: EndpointTrackingChannelBinding + ) -> JointPositionTrackingState: + from .runtime_commands import EndpointCommand, JointPositionPayload + + if not isinstance(command, EndpointCommand): + raise TypeError("command must be an EndpointCommand.") + if binding.channel_id != JOINT_POSITION_CHANNEL: + raise ValueError("Joint projector requires the joint.position channel.") + if not isinstance(command.payload, JointPositionPayload): + raise TypeError("Joint projector requires JointPositionPayload.") + address = binding.source.address + if isinstance(address, EndpointTrackingFeedbackAddress): + if address.target.address_fingerprint != command.target.address_fingerprint: + raise ValueError( + "Command and feedback binding target different endpoints." + ) + return JointPositionTrackingState(command.payload.positions) + + +def _compatible( + desired: TrackingState, + observed: TrackingState, + valid_mask: torch.Tensor, + expected_type: type[TrackingState], +) -> None: + if type(desired) is not expected_type or type(observed) is not expected_type: + raise TypeError(f"Metric requires {expected_type.__name__} values.") + if desired.batch_size != observed.batch_size or desired.device != observed.device: + raise ValueError("Desired and observed batches must match.") + if valid_mask.dtype != torch.bool or valid_mask.shape != (desired.batch_size,): + raise ValueError("valid_mask must be a bool tensor with one value per row.") + if valid_mask.device != desired.device: + raise ValueError("valid_mask and states must share a device.") + + +def _pose_errors( + desired: torch.Tensor, observed: torch.Tensor +) -> tuple[torch.Tensor, torch.Tensor]: + translation = torch.linalg.vector_norm( + desired[:, :3, 3] - observed[:, :3, 3], dim=1 + ) + relative = desired[:, :3, :3].transpose(1, 2) @ observed[:, :3, :3] + cosine = ((relative.diagonal(dim1=1, dim2=2).sum(dim=1) - 1.0) * 0.5).clamp( + -1.0, 1.0 + ) + return translation, torch.acos(cosine) + + +class JointPositionTrackingEvaluator: + """Evaluator for :class:`JointPositionTrackingMetric`.""" + + metric_id = JointPositionTrackingMetric.metric_id + revision = JointPositionTrackingMetric.revision + metric_type = JointPositionTrackingMetric + + def evaluate(self, desired, observed, valid_mask, metric) -> TrackingEvaluation: + _compatible(desired, observed, valid_mask, JointPositionTrackingState) + if type(metric) is not JointPositionTrackingMetric: + raise TypeError("metric must be JointPositionTrackingMetric.") + if desired.positions.shape != observed.positions.shape: + raise ValueError("Joint-position state shapes must match.") + error = (desired.positions - observed.positions).abs().amax(dim=1) + normalized = error / metric.tolerance + return TrackingEvaluation( + JOINT_POSITION_CHANNEL, + valid_mask & (error <= metric.tolerance), + valid_mask, + normalized, + {"joint_max_abs": error}, + ) + + +class PoseTrackingEvaluator: + """Evaluator for :class:`PoseTrackingMetric`.""" + + metric_id = PoseTrackingMetric.metric_id + revision = PoseTrackingMetric.revision + metric_type = PoseTrackingMetric + + def evaluate(self, desired, observed, valid_mask, metric) -> TrackingEvaluation: + _compatible(desired, observed, valid_mask, PoseTrackingState) + if type(metric) is not PoseTrackingMetric: + raise TypeError("metric must be PoseTrackingMetric.") + translation, rotation = _pose_errors(desired.poses, observed.poses) + normalized = torch.maximum( + translation / metric.translation_tolerance, + rotation / metric.rotation_tolerance, + ) + return TrackingEvaluation( + BASE_POSE_CHANNEL, + valid_mask & (normalized <= 1.0), + valid_mask, + normalized, + {"translation": translation, "rotation": rotation}, + ) + + +class WholeBodyPoseTrackingEvaluator: + """Evaluator for :class:`WholeBodyPoseTrackingMetric`.""" + + metric_id = WholeBodyPoseTrackingMetric.metric_id + revision = WholeBodyPoseTrackingMetric.revision + metric_type = WholeBodyPoseTrackingMetric + + def evaluate(self, desired, observed, valid_mask, metric) -> TrackingEvaluation: + _compatible(desired, observed, valid_mask, WholeBodyPoseTrackingState) + if type(metric) is not WholeBodyPoseTrackingMetric: + raise TypeError("metric must be WholeBodyPoseTrackingMetric.") + if desired.joint_positions.shape != observed.joint_positions.shape: + raise ValueError("Whole-body joint-position shapes must match.") + translation, rotation = _pose_errors(desired.root_poses, observed.root_poses) + joint = (desired.joint_positions - observed.joint_positions).abs().amax(dim=1) + normalized = torch.maximum( + torch.maximum( + translation / metric.translation_tolerance, + rotation / metric.rotation_tolerance, + ), + joint / metric.joint_position_tolerance, + ) + return TrackingEvaluation( + WHOLE_BODY_POSE_CHANNEL, + valid_mask & (normalized <= 1.0), + valid_mask, + normalized, + {"translation": translation, "rotation": rotation, "joint_max_abs": joint}, + ) + + +class TrackingRuntime: + """Runtime facade for projecting commands and evaluating typed feedback.""" + + __slots__ = ("_providers", "_projectors", "_evaluators") + + def __init__( + self, + providers: TrackingFeedbackProviderRegistry, + projectors: TrackingProjectorRegistry, + evaluators: TrackingEvaluatorRegistry, + ) -> None: + if type(providers) is not TrackingFeedbackProviderRegistry: + raise TypeError( + "providers must be exactly TrackingFeedbackProviderRegistry." + ) + if type(projectors) is not TrackingProjectorRegistry: + raise TypeError("projectors must be exactly TrackingProjectorRegistry.") + if type(evaluators) is not TrackingEvaluatorRegistry: + raise TypeError("evaluators must be exactly TrackingEvaluatorRegistry.") + self._providers = providers + self._projectors = projectors + self._evaluators = evaluators + + @property + def providers(self) -> TrackingFeedbackProviderRegistry: + """Return the immutable exact-version provider registry.""" + return self._providers + + @property + def projectors(self) -> TrackingProjectorRegistry: + """Return the immutable exact-version projector registry.""" + return self._projectors + + @property + def evaluators(self) -> TrackingEvaluatorRegistry: + """Return the immutable exact-version evaluator registry.""" + return self._evaluators + + @classmethod + def with_builtins(cls) -> TrackingRuntime: + """Create a runtime with context feedback and built-in typed metrics.""" + return cls( + TrackingFeedbackProviderRegistry( + [PlanningContextTrackingFeedbackProvider()] + ), + TrackingProjectorRegistry([JointPositionTrackingProjector()]), + TrackingEvaluatorRegistry( + [ + JointPositionTrackingEvaluator(), + PoseTrackingEvaluator(), + WholeBodyPoseTrackingEvaluator(), + ] + ), + ) + + def project( + self, command: EndpointCommand, binding: EndpointTrackingChannelBinding + ) -> TrackingState: + """Project one command through the exact binding-owned projector.""" + return self.projectors.resolve(binding.projector).project(command, binding) + + def observe( + self, setpoint: TrackingSetpoint, context: PlanningContext + ) -> TrackingFeedbackBatch: + """Read the exact feedback source for one setpoint.""" + feedback = self.providers.resolve(setpoint.binding.source).observe( + setpoint.binding.source, context + ) + if ( + feedback.source.source_fingerprint + != setpoint.binding.source.source_fingerprint + ): + raise ValueError("Feedback provider returned a different source.") + if feedback.state.channel_id != setpoint.binding.channel_id: + raise TypeError("Feedback state does not match the bound channel.") + if feedback.timestamp != context.robot.timestamp: + raise ValueError( + "Tracking feedback must use the current planning-context timestamp." + ) + if feedback.state.batch_size != context.batch_size: + raise ValueError("Tracking feedback batch must match the context batch.") + if feedback.state.device != context.robot.qpos.device: + raise ValueError("Tracking feedback and context must share a device.") + return feedback + + def evaluate( + self, + setpoint: TrackingSetpoint, + feedback: TrackingFeedbackBatch, + metric: TrackingMetricCfg, + ) -> TrackingEvaluation: + """Evaluate one observed setpoint with an exact metric implementation.""" + if metric.channel_id != setpoint.binding.channel_id: + raise ValueError("Metric and setpoint channels must match.") + if ( + feedback.source.source_fingerprint + != setpoint.binding.source.source_fingerprint + ): + raise ValueError("Feedback source does not match the setpoint binding.") + return self.evaluators.resolve(metric).evaluate( + setpoint.desired, feedback.state, feedback.valid_mask, metric + ) + + def evaluate_frame( + self, + frame: TrackingFrame, + metrics: Iterable[TrackingMetricCfg], + context: PlanningContext, + ) -> Mapping[tuple[str, str, str], TrackingEvaluation]: + """Observe and evaluate every setpoint required by one frame.""" + by_channel = {metric.channel_id: metric for metric in metrics} + results: dict[tuple[str, str, str], TrackingEvaluation] = {} + for setpoint in frame.setpoints: + try: + metric = by_channel[setpoint.binding.channel_id] + except KeyError as exc: + raise KeyError( + f"No metric configured for channel {setpoint.binding.channel_id!r}." + ) from exc + results[setpoint.key] = self.evaluate( + setpoint, self.observe(setpoint, context), metric + ) + return MappingProxyType(results) + + +__all__ = [ + "BASE_POSE_CHANNEL", + "FeedbackTerminalAcceptance", + "InFlightTrackingPolicy", + "JOINT_POSITION_CHANNEL", + "JointPositionTrackingEvaluator", + "JointPositionTrackingMetric", + "JointPositionTrackingProjector", + "JointPositionTrackingState", + "EndpointTrackingChannelBinding", + "EndpointTrackingFeedbackAddress", + "PlanningContextTrackingFeedbackProvider", + "PoseTrackingEvaluator", + "PoseTrackingMetric", + "PoseTrackingState", + "TerminalAcceptance", + "TimedTerminalAcceptance", + "TimedTrackingSequence", + "TrackingChannelId", + "TrackingCommandProjector", + "TrackingEvaluation", + "TrackingEvaluatorRegistry", + "TrackingFeedbackAddress", + "TrackingFeedbackBatch", + "TrackingFeedbackProvider", + "TrackingFeedbackProviderRegistry", + "TrackingFeedbackSourceRef", + "TrackingFrame", + "TrackingMetricCfg", + "TrackingMetricEvaluator", + "TrackingPolicy", + "TrackingProjectorRef", + "TrackingProjectorRegistry", + "TrackingRuntime", + "TrackingSetpoint", + "TrackingState", + "WHOLE_BODY_POSE_CHANNEL", + "WholeBodyPoseTrackingEvaluator", + "WholeBodyPoseTrackingMetric", + "WholeBodyPoseTrackingState", +] diff --git a/embodichain/lab/sim/skills/__init__.py b/embodichain/lab/sim/skills/__init__.py index 576d2f24a..9bd87c54a 100644 --- a/embodichain/lab/sim/skills/__init__.py +++ b/embodichain/lab/sim/skills/__init__.py @@ -200,6 +200,7 @@ ResolvedCorePolicyTrace, SkillCallTrace, SkillEndpointBindingTrace, + SkillEndpointTrackingChannelTrace, SkillEffectTrace, SkillFailure, SkillPlanAttemptTrace, @@ -367,6 +368,7 @@ "SkillPolicyPreset", "SkillCallTrace", "SkillEndpointBindingTrace", + "SkillEndpointTrackingChannelTrace", "SkillEffectTrace", "SkillFailure", "SkillPlanAttemptTrace", diff --git a/embodichain/lab/sim/skills/compiler.py b/embodichain/lab/sim/skills/compiler.py index 69cc0832b..def756fc5 100644 --- a/embodichain/lab/sim/skills/compiler.py +++ b/embodichain/lab/sim/skills/compiler.py @@ -1083,6 +1083,7 @@ def ground( goal=lowering.goal, binding=bound.binding.action_binding, motion_policy=bound.preset.motion_policy, + tracking_policy=bound.preset.tracking_policy, recovery_policy=bound.preset.recovery_policy, skill_options=lowering.skill_options, control_overrides=lowering.control_overrides, diff --git a/embodichain/lab/sim/skills/integration.py b/embodichain/lab/sim/skills/integration.py index f98ab1fe4..0892f2284 100644 --- a/embodichain/lab/sim/skills/integration.py +++ b/embodichain/lab/sim/skills/integration.py @@ -1386,6 +1386,7 @@ def link_call( preset.motion_policy, dynamic_collision_mode=DynamicCollisionMode.REQUIRED, ), + tracking_policy=preset.tracking_policy, recovery_policy=preset.recovery_policy, runner_cfg=preset.runner_cfg, effect_monitors=preset.effect_monitors, diff --git a/embodichain/lab/sim/skills/profiles.py b/embodichain/lab/sim/skills/profiles.py index 9623e8a6d..2fc875d5d 100644 --- a/embodichain/lab/sim/skills/profiles.py +++ b/embodichain/lab/sim/skills/profiles.py @@ -38,6 +38,14 @@ ) from embodichain.lab.sim.atomic_actions.core import SkillDescriptor from embodichain.lab.sim.atomic_actions.policies import MotionPolicy, RecoveryPolicy +from embodichain.lab.sim.atomic_actions.tracking import ( + JOINT_POSITION_CHANNEL, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, + TrackingFeedbackSourceRef, + TrackingPolicy, + TrackingProjectorRef, +) from embodichain.lab.sim.atomic_actions.requirements import ( BATCH_INVERSE_KINEMATICS_CAPABILITY, DisjointResourceSlots, @@ -171,6 +179,37 @@ def _snapshot_effect_sources( return MappingProxyType(snapshots) +def _snapshot_tracking_channels( + values: Mapping[str, EndpointTrackingChannelBinding], + *, + field_name: str, +) -> Mapping[str, EndpointTrackingChannelBinding]: + """Validate, own, and freeze endpoint tracking bindings by channel.""" + if not isinstance(values, Mapping): + raise TypeError(f"{field_name} must be a mapping.") + snapshots: dict[str, EndpointTrackingChannelBinding] = {} + for channel_id, binding in values.items(): + _validate_identifier(channel_id, field_name=f"{field_name} channel IDs") + if not isinstance(binding, EndpointTrackingChannelBinding): + raise TypeError( + f"{field_name} values must be EndpointTrackingChannelBinding " + "instances." + ) + if binding.channel_id != channel_id: + raise ValueError( + f"{field_name}[{channel_id!r}] disagrees with binding channel " + f"{binding.channel_id!r}." + ) + snapshot = binding.snapshot() + if snapshot is binding: + raise TypeError( + f"{field_name}[{channel_id!r}].snapshot() must return an " + "independent channel binding." + ) + snapshots[channel_id] = snapshot + return MappingProxyType(snapshots) + + @dataclass(frozen=True, slots=True, kw_only=True) class ResourceEndpoint(ABC): """Extensible execution endpoint in a robot resource graph. @@ -242,6 +281,11 @@ class EndpointResolution: effect_sources: Mapping[str, EffectEvidenceSourceRef] = field(default_factory=dict) """Provider-routed raw observation sources keyed by open channel ID.""" + tracking_channels: Mapping[str, EndpointTrackingChannelBinding] = field( + default_factory=dict + ) + """Typed feedback source and desired-state projector by channel ID.""" + command_profile_key: str | None = None """Profile key that owns semantic commands for this endpoint, when any.""" @@ -293,6 +337,14 @@ def __post_init__(self) -> None: field_name="EndpointResolution.effect_sources", ), ) + object.__setattr__( + self, + "tracking_channels", + _snapshot_tracking_channels( + self.tracking_channels, + field_name="EndpointResolution.tracking_channels", + ), + ) if self.command_profile_key is not None: _validate_identifier( self.command_profile_key, @@ -435,11 +487,24 @@ def resolve( FORCE_EFFECT_CHANNEL, } ) - return EndpointResolution( - runtime_target=JointPositionTarget( - control_part=endpoint.control_part, - joint_ids=joint_ids, + runtime_target = JointPositionTarget( + control_part=endpoint.control_part, + joint_ids=joint_ids, + ) + tracking_channel = EndpointTrackingChannelBinding( + JOINT_POSITION_CHANNEL, + TrackingFeedbackSourceRef( + "planning_context.robot", + "1", + EndpointTrackingFeedbackAddress( + runtime_target, + JOINT_POSITION_CHANNEL, + ), ), + TrackingProjectorRef("joint_position_payload", "1"), + ) + return EndpointResolution( + runtime_target=runtime_target, command_profile_key=( endpoint.control_part if endpoint.command_profile is None @@ -454,6 +519,7 @@ def resolve( ) for channel in sorted(effect_channels) }, + tracking_channels={JOINT_POSITION_CHANNEL: tracking_channel}, claim_tokens=frozenset({f"robot.control_part:{endpoint.control_part}"}), joint_ids=joint_ids, ) @@ -468,6 +534,9 @@ class ResolvedResourceEndpoint: runtime_target: RuntimeEndpointTarget task_state_key: str | None = None effect_sources: Mapping[str, EffectEvidenceSourceRef] = field(default_factory=dict) + tracking_channels: Mapping[str, EndpointTrackingChannelBinding] = field( + default_factory=dict + ) command_profile_key: str | None = None requires_command_profile: bool = False commands: Mapping[str, ControlCommand] = field(default_factory=dict) @@ -496,6 +565,7 @@ def __post_init__(self) -> None: runtime_target=self.runtime_target, task_state_key=self.task_state_key, effect_sources=self.effect_sources, + tracking_channels=self.tracking_channels, command_profile_key=self.command_profile_key, requires_command_profile=self.requires_command_profile, claim_tokens=self.claim_tokens, @@ -510,6 +580,7 @@ def __post_init__(self) -> None: ) object.__setattr__(self, "task_state_key", resolved_state_key) object.__setattr__(self, "effect_sources", resolution.effect_sources) + object.__setattr__(self, "tracking_channels", resolution.tracking_channels) object.__setattr__( self, "command_profile_key", @@ -646,11 +717,12 @@ def __post_init__(self) -> None: @dataclass(frozen=True, slots=True, init=False) class SkillPolicyPreset: - """Versioned planning, recovery, runner, and effect-monitor bundle.""" + """Versioned planning, tracking, recovery, runner, and monitor bundle.""" preset_id: str schema_version: int _motion_policy: MotionPolicy + _tracking_policy: TrackingPolicy _recovery_policy: RecoveryPolicy _runner_cfg: ExecutionRunnerCfg _effect_monitors: Mapping[str, EffectMonitorRef] @@ -660,6 +732,7 @@ def __init__( preset_id: str, schema_version: int = 1, motion_policy: MotionPolicy | None = None, + tracking_policy: TrackingPolicy | None = None, recovery_policy: RecoveryPolicy | None = None, runner_cfg: ExecutionRunnerCfg | None = None, effect_monitors: Mapping[str, EffectMonitorRef] | None = None, @@ -674,12 +747,19 @@ def __init__( f"{schema_version}; supported versions are [1]." ) selected_motion = MotionPolicy() if motion_policy is None else motion_policy + selected_tracking = ( + TrackingPolicy.joint_position() + if tracking_policy is None + else tracking_policy + ) selected_recovery = ( RecoveryPolicy() if recovery_policy is None else recovery_policy ) selected_runner = ExecutionRunnerCfg() if runner_cfg is None else runner_cfg if not isinstance(selected_motion, MotionPolicy): raise TypeError("motion_policy must be a MotionPolicy.") + if not isinstance(selected_tracking, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") if not isinstance(selected_recovery, RecoveryPolicy): raise TypeError("recovery_policy must be a RecoveryPolicy.") if not isinstance(selected_runner, ExecutionRunnerCfg): @@ -716,6 +796,7 @@ def __init__( object.__setattr__(self, "preset_id", preset_id) object.__setattr__(self, "schema_version", schema_version) object.__setattr__(self, "_motion_policy", deepcopy(selected_motion)) + object.__setattr__(self, "_tracking_policy", deepcopy(selected_tracking)) object.__setattr__(self, "_recovery_policy", deepcopy(selected_recovery)) object.__setattr__(self, "_runner_cfg", deepcopy(selected_runner)) object.__setattr__( @@ -734,6 +815,11 @@ def recovery_policy(self) -> RecoveryPolicy: """Return an independently owned recovery policy.""" return deepcopy(self._recovery_policy) + @property + def tracking_policy(self) -> TrackingPolicy: + """Return independently owned endpoint-tracking settings.""" + return deepcopy(self._tracking_policy) + @property def runner_cfg(self) -> ExecutionRunnerCfg: """Return an independently owned runner configuration.""" @@ -755,6 +841,7 @@ def snapshot(self) -> SkillPolicyPreset: preset_id=self.preset_id, schema_version=self.schema_version, motion_policy=self.motion_policy, + tracking_policy=self.tracking_policy, recovery_policy=self.recovery_policy, runner_cfg=self.runner_cfg, effect_monitors=self.effect_monitors, @@ -1605,6 +1692,7 @@ def _resolve_resources(self) -> Mapping[str, ResolvedRobotResource]: else resolution.task_state_key ), effect_sources=resolution.effect_sources, + tracking_channels=resolution.tracking_channels, command_profile_key=resolution.command_profile_key, requires_command_profile=resolution.requires_command_profile, commands=( @@ -2005,6 +2093,7 @@ def _lower_binding( adapter_id=endpoint.adapter_id, target=endpoint.runtime_target, task_state_key=endpoint.task_state_key, + tracking_channels=endpoint.tracking_channels, capabilities=endpoint.capabilities, commands=endpoint.commands, claim_tokens=endpoint.claim_tokens, diff --git a/embodichain/lab/sim/skills/runtime.py b/embodichain/lab/sim/skills/runtime.py index de98f5ed3..46fe7da6e 100644 --- a/embodichain/lab/sim/skills/runtime.py +++ b/embodichain/lab/sim/skills/runtime.py @@ -19,7 +19,7 @@ from __future__ import annotations from collections.abc import Iterable, Mapping -from dataclasses import dataclass, replace +from dataclasses import dataclass, fields, is_dataclass, replace from enum import Enum import math from types import MappingProxyType @@ -35,7 +35,7 @@ ExecutionEvent, ExecutionPlanAttempt, ) -from ..atomic_actions.plans import ExecutionFeedbackMode, TrajectorySegment +from ..atomic_actions.plans import TrajectorySegment from ..atomic_actions.policies import MotionPolicy, RecoveryPolicy from ..atomic_actions.runner import ( CommandSink, @@ -48,6 +48,12 @@ RunnerStep, ) from ..atomic_actions.state import PlanningContext, TaskState +from ..atomic_actions.tracking import ( + FeedbackTerminalAcceptance, + TimedTrackingSequence, + TrackingMetricCfg, + TrackingPolicy, +) from .calls import SemanticCallSpec from .compiler import SemanticSkillCompiler from .effects import ( @@ -109,6 +115,8 @@ def _metadata_value(value: object, *, depth: int = 0) -> object: return _metadata_value(value.detach().cpu().tolist(), depth=depth + 1) if isinstance(value, torch.device): return str(value) + if isinstance(value, type): + return {"__type__": f"{value.__module__}.{value.__qualname__}"} if isinstance(value, Mapping): items = sorted(value.items(), key=lambda item: str(item[0])) if all(type(key) is str and key and key == key.strip() for key, _ in items): @@ -134,6 +142,17 @@ def _metadata_value(value: object, *, depth: int = 0) -> object: return {"type": f"{type(value).__module__}.{type(value).__qualname__}"} +def _freeze_metadata_value(value: object) -> object: + """Recursively freeze already JSON-safe metadata for immutable traces.""" + if isinstance(value, dict): + return MappingProxyType( + {key: _freeze_metadata_value(nested) for key, nested in value.items()} + ) + if isinstance(value, list): + return tuple(_freeze_metadata_value(nested) for nested in value) + return value + + def _snapshot_metadata_mapping(value: Mapping[str, object]) -> Mapping[str, object]: """Own one JSON-safe string-keyed metadata mapping.""" if not isinstance(value, Mapping): @@ -217,6 +236,59 @@ class SkillStatus(str, Enum): CANCELLED = "cancelled" +@dataclass(frozen=True, slots=True) +class SkillEndpointTrackingChannelTrace: + """Stable provider and projector route for one endpoint feedback channel.""" + + channel_id: str + provider_id: str + provider_revision: str + projector_id: str + projector_revision: str + feedback_address_type: str + address_fingerprint: object + route_fingerprint: object + + def __post_init__(self) -> None: + for name in ( + "channel_id", + "provider_id", + "provider_revision", + "projector_id", + "projector_revision", + "feedback_address_type", + ): + if type(getattr(self, name)) is not str or not getattr(self, name): + raise ValueError(f"{name} must be a non-empty string.") + object.__setattr__( + self, + "address_fingerprint", + _freeze_metadata_value(_metadata_value(self.address_fingerprint)), + ) + object.__setattr__( + self, + "route_fingerprint", + _freeze_metadata_value(_metadata_value(self.route_fingerprint)), + ) + + def to_metadata(self) -> dict[str, object]: + """Return the exact immutable tracking route without live objects.""" + return { + "channel_id": self.channel_id, + "feedback_source": { + "provider_id": self.provider_id, + "revision": self.provider_revision, + "address_type": self.feedback_address_type, + "address_fingerprint": _metadata_value(self.address_fingerprint), + }, + "projector": { + "projector_id": self.projector_id, + "revision": self.projector_revision, + }, + "route_fingerprint": _metadata_value(self.route_fingerprint), + } + + @dataclass(frozen=True, slots=True) class SkillEndpointBindingTrace: """JSON-safe typed projection of one resolved execution endpoint.""" @@ -231,6 +303,7 @@ class SkillEndpointBindingTrace: task_state_key: str capabilities: tuple[str, ...] command_ids: tuple[str, ...] + tracking_channels: tuple[SkillEndpointTrackingChannelTrace, ...] claim_tokens: tuple[str, ...] joint_ids: tuple[int, ...] @@ -255,6 +328,21 @@ def __post_init__(self) -> None: ): raise ValueError(f"{name} must contain sorted unique identifiers.") object.__setattr__(self, name, values) + tracking_channels = tuple(self.tracking_channels) + if not all( + type(value) is SkillEndpointTrackingChannelTrace + for value in tracking_channels + ): + raise TypeError( + "tracking_channels must contain exact " + "SkillEndpointTrackingChannelTrace values." + ) + channel_ids = tuple(value.channel_id for value in tracking_channels) + if tuple(sorted(set(channel_ids))) != channel_ids: + raise ValueError( + "tracking_channels must use sorted unique channel identifiers." + ) + object.__setattr__(self, "tracking_channels", tracking_channels) joint_ids = tuple(self.joint_ids) if len(set(joint_ids)) != len(joint_ids) or not all( type(value) is int and value >= 0 for value in joint_ids @@ -279,6 +367,22 @@ def from_binding(cls, binding: EndpointBinding) -> SkillEndpointBindingTrace: task_state_key=binding.task_state_key, capabilities=tuple(sorted(binding.capabilities)), command_ids=tuple(sorted(binding.commands)), + tracking_channels=tuple( + SkillEndpointTrackingChannelTrace( + channel_id=channel_id, + provider_id=channel.source.provider_id, + provider_revision=channel.source.revision, + projector_id=channel.projector.projector_id, + projector_revision=channel.projector.revision, + feedback_address_type=( + f"{type(channel.source.address).__module__}." + f"{type(channel.source.address).__qualname__}" + ), + address_fingerprint=(channel.source.address.address_fingerprint), + route_fingerprint=channel.route_fingerprint, + ) + for channel_id, channel in sorted(binding.tracking_channels.items()) + ), claim_tokens=tuple(sorted(binding.claim_tokens)), joint_ids=binding.joint_ids, ) @@ -296,6 +400,9 @@ def to_metadata(self) -> dict[str, object]: "task_state_key": self.task_state_key, "capabilities": list(self.capabilities), "command_ids": list(self.command_ids), + "tracking_channels": [ + channel.to_metadata() for channel in self.tracking_channels + ], "claim_tokens": list(self.claim_tokens), "joint_ids": list(self.joint_ids), } @@ -332,7 +439,6 @@ def _recovery_policy_to_metadata(policy: RecoveryPolicy) -> dict[str, object]: return { "max_replans": policy.max_replans, "max_action_retries": policy.max_action_retries, - "tracking_error_threshold": _metadata_value(policy.tracking_error_threshold), "goal_translation_threshold": _metadata_value( policy.goal_translation_threshold ), @@ -341,6 +447,97 @@ def _recovery_policy_to_metadata(policy: RecoveryPolicy) -> dict[str, object]: } +def _tracking_metric_to_metadata(metric: TrackingMetricCfg) -> dict[str, object]: + """Serialize one exact typed metric and its unit-preserving tolerances.""" + parameters = ( + { + value.name: _metadata_value(getattr(metric, value.name)) + for value in fields(metric) + } + if is_dataclass(metric) + else {} + ) + return { + "metric_id": metric.metric_id, + "revision": metric.revision, + "channel_id": metric.channel_id, + "type": f"{type(metric).__module__}.{type(metric).__qualname__}", + "parameters": parameters, + } + + +def _tracking_policy_to_metadata(policy: TrackingPolicy) -> dict[str, object]: + """Serialize independent in-flight and terminal tracking contracts.""" + in_flight = policy.in_flight + terminal = policy.terminal + return { + "in_flight": ( + None + if in_flight is None + else { + "metrics": [ + _tracking_metric_to_metadata(metric) for metric in in_flight.metrics + ], + "consecutive_violations": in_flight.consecutive_violations, + "grace_period": _metadata_value(in_flight.grace_period), + } + ), + "terminal": ( + { + "mode": "feedback", + "metrics": [ + _tracking_metric_to_metadata(metric) for metric in terminal.metrics + ], + "settle_timeout": _metadata_value(terminal.settle_timeout), + "consecutive_acceptances": terminal.consecutive_acceptances, + } + if isinstance(terminal, FeedbackTerminalAcceptance) + else { + "mode": "timed", + "settle_duration": _metadata_value(terminal.settle_duration), + } + ), + } + + +def _tracking_sequence_to_metadata( + sequence: TimedTrackingSequence | None, +) -> dict[str, object] | None: + """Serialize the provider/projector shape of one plan-owned contract.""" + if sequence is None: + return None + first_frame = None if not sequence.frames else sequence.frames[0] + return { + "env_ids": _metadata_value(sequence.env_ids), + "frame_count": sequence.frame_count, + "setpoints": [ + { + "endpoint": list(setpoint.endpoint_key), + "channel_id": setpoint.binding.channel_id, + "state_type": ( + f"{type(setpoint.desired).__module__}." + f"{type(setpoint.desired).__qualname__}" + ), + "feedback_source": { + "provider_id": setpoint.binding.source.provider_id, + "revision": setpoint.binding.source.revision, + "address_fingerprint": _metadata_value( + setpoint.binding.source.address.address_fingerprint + ), + }, + "projector": { + "projector_id": setpoint.binding.projector.projector_id, + "revision": setpoint.binding.projector.revision, + }, + "route_fingerprint": _metadata_value( + setpoint.binding.route_fingerprint + ), + } + for setpoint in (() if first_frame is None else first_frame.setpoints) + ], + } + + @dataclass(frozen=True, slots=True) class ResolvedCorePolicyTrace: """Resolved preset, core policies, and execution binding for one plan.""" @@ -349,6 +546,7 @@ class ResolvedCorePolicyTrace: preset_id: str preset_schema_version: int motion_policy: MotionPolicy + tracking_policy: TrackingPolicy recovery_policy: RecoveryPolicy endpoints: tuple[SkillEndpointBindingTrace, ...] @@ -364,6 +562,8 @@ def __post_init__(self) -> None: raise ValueError("preset_schema_version must be a positive integer.") if not isinstance(self.motion_policy, MotionPolicy): raise TypeError("motion_policy must be a MotionPolicy.") + if not isinstance(self.tracking_policy, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") if not isinstance(self.recovery_policy, RecoveryPolicy): raise TypeError("recovery_policy must be a RecoveryPolicy.") endpoints = tuple(self.endpoints) @@ -375,6 +575,7 @@ def __post_init__(self) -> None: if len(set(keys)) != len(keys): raise ValueError("endpoints must use unique slot/endpoint keys.") object.__setattr__(self, "motion_policy", replace(self.motion_policy)) + object.__setattr__(self, "tracking_policy", self.tracking_policy.snapshot()) object.__setattr__(self, "recovery_policy", replace(self.recovery_policy)) object.__setattr__(self, "endpoints", endpoints) @@ -386,6 +587,7 @@ def from_resolved_binding( preset_id: str, preset_schema_version: int, motion_policy: MotionPolicy, + tracking_policy: TrackingPolicy, recovery_policy: RecoveryPolicy, endpoints: Iterable[EndpointBinding], ) -> ResolvedCorePolicyTrace: @@ -395,6 +597,7 @@ def from_resolved_binding( preset_id=preset_id, preset_schema_version=preset_schema_version, motion_policy=motion_policy, + tracking_policy=tracking_policy, recovery_policy=recovery_policy, endpoints=tuple( SkillEndpointBindingTrace.from_binding(endpoint) @@ -409,6 +612,7 @@ def snapshot(self) -> ResolvedCorePolicyTrace: preset_id=self.preset_id, preset_schema_version=self.preset_schema_version, motion_policy=self.motion_policy, + tracking_policy=self.tracking_policy, recovery_policy=self.recovery_policy, endpoints=self.endpoints, ) @@ -422,6 +626,7 @@ def to_metadata(self) -> dict[str, object]: "schema_version": self.preset_schema_version, }, "motion_policy": _motion_policy_to_metadata(self.motion_policy), + "tracking_policy": _tracking_policy_to_metadata(self.tracking_policy), "recovery_policy": _recovery_policy_to_metadata(self.recovery_policy), "endpoints": [endpoint.to_metadata() for endpoint in self.endpoints], } @@ -455,7 +660,8 @@ class SkillPlanAttemptTrace: scene_dependency_monitor_until: Mapping[str, int] collision_world_sensitive: bool replannable: bool - feedback_mode: ExecutionFeedbackMode + tracking_policy: TrackingPolicy + tracking: TimedTrackingSequence | None effect_verification_kind: str | None resolved_core_policy: ResolvedCorePolicyTrace planner_backend: str @@ -545,8 +751,12 @@ def __post_init__(self) -> None: raise TypeError("collision_world_sensitive must be a bool.") if type(self.replannable) is not bool: raise TypeError("replannable must be a bool.") - if not isinstance(self.feedback_mode, ExecutionFeedbackMode): - raise TypeError("feedback_mode must be an ExecutionFeedbackMode.") + if not isinstance(self.tracking_policy, TrackingPolicy): + raise TypeError("tracking_policy must be a TrackingPolicy.") + if self.tracking is not None and not isinstance( + self.tracking, TimedTrackingSequence + ): + raise TypeError("tracking must be a TimedTrackingSequence or None.") if self.effect_verification_kind is not None and ( type(self.effect_verification_kind) is not str or not self.effect_verification_kind @@ -577,6 +787,9 @@ def __post_init__(self) -> None: "scene_dependency_monitor_until", MappingProxyType(monitor_until), ) + object.__setattr__(self, "tracking_policy", self.tracking_policy.snapshot()) + if self.tracking is not None: + object.__setattr__(self, "tracking", self.tracking.snapshot()) object.__setattr__( self, "resolved_core_policy", @@ -623,7 +836,8 @@ def from_execution_attempt( scene_dependency_monitor_until=plan.scene_dependency_monitor_until, collision_world_sensitive=plan.collision_world_sensitive, replannable=plan.replannable, - feedback_mode=plan.feedback_mode, + tracking_policy=plan.tracking_policy, + tracking=plan.tracking, effect_verification_kind=( None if plan.effect_verification is None @@ -634,6 +848,7 @@ def from_execution_attempt( preset_id=preset_id, preset_schema_version=preset_schema_version, motion_policy=request.motion_policy, + tracking_policy=request.tracking_policy, recovery_policy=request.recovery_policy, endpoints=request.binding.endpoints, ), @@ -664,7 +879,8 @@ def snapshot(self) -> SkillPlanAttemptTrace: scene_dependency_monitor_until=self.scene_dependency_monitor_until, collision_world_sensitive=self.collision_world_sensitive, replannable=self.replannable, - feedback_mode=self.feedback_mode, + tracking_policy=self.tracking_policy, + tracking=self.tracking, effect_verification_kind=self.effect_verification_kind, resolved_core_policy=self.resolved_core_policy, planner_backend=self.planner_backend, @@ -709,7 +925,8 @@ def to_metadata(self) -> dict[str, object]: }, "collision_world_sensitive": self.collision_world_sensitive, "replannable": self.replannable, - "feedback_mode": self.feedback_mode.value, + "tracking_policy": _tracking_policy_to_metadata(self.tracking_policy), + "tracking_contract": _tracking_sequence_to_metadata(self.tracking), "effect_verification_kind": self.effect_verification_kind, "resolved_core_policy": self.resolved_core_policy.to_metadata(), "planner_diagnostics": { @@ -2042,6 +2259,11 @@ def _append_preparation_failure_trace( if invocation is None else invocation.motion_policy ), + tracking_policy=( + preset.tracking_policy + if invocation is None + else invocation.tracking_policy + ), recovery_policy=( preset.recovery_policy if invocation is None @@ -2282,6 +2504,7 @@ def cancel( "ResolvedCorePolicyTrace", "SkillCallTrace", "SkillEndpointBindingTrace", + "SkillEndpointTrackingChannelTrace", "SkillEffectTrace", "SkillFailure", "SkillPlanAttemptTrace", diff --git a/tests/gym/envs/expert_program/test_completion_metadata.py b/tests/gym/envs/expert_program/test_completion_metadata.py index 1befa0ac6..ce54ce7d3 100644 --- a/tests/gym/envs/expert_program/test_completion_metadata.py +++ b/tests/gym/envs/expert_program/test_completion_metadata.py @@ -165,7 +165,7 @@ def _plan( class _TraceObservationProvider: - """Move one scene dependency after the first installed command frame.""" + """Move the scene once and report accepted commands as observed state.""" def __init__(self, clock: EnvironmentStepClock) -> None: self.clock = clock @@ -177,7 +177,10 @@ def observe(self, task_state: TaskState) -> PlanningContext: pose = torch.eye(4).repeat(BATCH_SIZE, 1, 1) if replanned_scene: pose[:, 0, 3] = 0.25 - qpos = torch.zeros(BATCH_SIZE, ROBOT_DOF) + qpos = torch.full( + (BATCH_SIZE, ROBOT_DOF), + float(min(max(self.calls - 1, 0), 3)), + ) timestamp = self.clock.now() return PlanningContext( robot=RobotObservation( diff --git a/tests/sim/atomic_actions/test_core.py b/tests/sim/atomic_actions/test_core.py index 171ebd054..8fe203f87 100644 --- a/tests/sim/atomic_actions/test_core.py +++ b/tests/sim/atomic_actions/test_core.py @@ -36,10 +36,11 @@ DynamicCollisionMode, EndpointBinding, EndpointCommand, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, EndEffectorPoseGoal, EntityState, EffectVerificationRequirement, - ExecutionFeedbackMode, HeldObjectState, JointPositionPayload, JointPositionTarget, @@ -58,7 +59,15 @@ StateDelta, TaskState, TimedCommandSequence, + TimedTerminalAcceptance, + TimedTrackingSequence, TimedTrajectory, + TrackingFeedbackSourceRef, + TrackingFrame, + TrackingPolicy, + TrackingProjectorRef, + TrackingSetpoint, + JointPositionTrackingState, ) from embodichain.lab.sim.atomic_actions.goals import ( _resolve_object_pose, @@ -196,7 +205,8 @@ def _action_plan( *, plan_success: torch.Tensor | None = None, joint_trajectory: TimedTrajectory | None = None, - feedback_mode: ExecutionFeedbackMode = ExecutionFeedbackMode.TIMED, + tracking_policy: TrackingPolicy | None = None, + tracking: TimedTrackingSequence | None = None, expected_effects: StateDelta | None = None, effect_verification: EffectVerificationRequirement | None = None, diagnostics: PlannerDiagnostics | None = None, @@ -214,12 +224,15 @@ def _action_plan( plan_success=plan_success, commands=commands, recovery_policy=RecoveryPolicy(), + tracking_policy=( + TrackingPolicy.timed() if tracking_policy is None else tracking_policy + ), planned_scene_version=0, planned_collision_world_revision=(0,) * commands.batch_size, diagnostics=( PlannerDiagnostics(backend="test") if diagnostics is None else diagnostics ), - feedback_mode=feedback_mode, + tracking=tracking, joint_trajectory=joint_trajectory, scene_dependencies=scene_dependencies, scene_dependency_monitor_until=( @@ -232,6 +245,41 @@ def _action_plan( ) +def _joint_tracking_sequence( + commands: TimedCommandSequence, +) -> TimedTrackingSequence: + frames: list[TrackingFrame] = [] + for command_frame in commands.frames: + setpoints: list[TrackingSetpoint] = [] + for command in command_frame.commands: + assert isinstance(command.target, JointPositionTarget) + assert isinstance(command.payload, JointPositionPayload) + channel = EndpointTrackingChannelBinding( + channel_id="joint.position", + source=TrackingFeedbackSourceRef( + provider_id="planning_context.robot", + revision="1", + address=EndpointTrackingFeedbackAddress( + target=command.target, + channel_id="joint.position", + ), + ), + projector=TrackingProjectorRef( + projector_id="joint_position_payload", + revision="1", + ), + ) + setpoints.append( + TrackingSetpoint( + endpoint_key=("primary", "motion"), + binding=channel, + desired=JointPositionTrackingState(command.payload.positions), + ) + ) + frames.append(TrackingFrame(tuple(setpoints))) + return TimedTrackingSequence(commands.env_ids, tuple(frames)) + + @pytest.mark.parametrize("kind", ("", " physical", "physical ", 1, True)) def test_effect_verification_requirement_rejects_invalid_kind(kind: object) -> None: with pytest.raises(ValueError, match="kind"): @@ -353,6 +401,7 @@ def _plan( frame_count=1, ), recovery_policy=RecoveryPolicy(), + tracking_policy=TrackingPolicy.timed(), planned_scene_version=context.scene.version, planned_collision_world_revision=(0,) * context.batch_size, diagnostics=PlannerDiagnostics(backend="test"), @@ -756,6 +805,7 @@ def test_build_plan_uses_action_scene_dependency_hook() -> None: goal=EndEffectorPoseGoal(SceneEntityPose("tracked")), binding=ActionBinding(owner_id=engine.binding_owner_id), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -807,6 +857,7 @@ def test_build_command_plan_rejects_unbound_runtime_destination() -> None: goal=EndEffectorPoseGoal(SceneEntityPose("tracked")), binding=ActionBinding(owner_id=engine.binding_owner_id), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -835,6 +886,7 @@ def test_public_plan_authorizes_raw_action_plan_destinations() -> None: goal=EndEffectorPoseGoal(SceneEntityPose("tracked")), binding=ActionBinding(owner_id=engine.binding_owner_id), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -861,6 +913,7 @@ def test_command_target_authorization_rejects_altered_joint_claims() -> None: ), ), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -901,6 +954,7 @@ def test_command_target_authorization_rejects_custom_claim_conflicts() -> None: goal=EndEffectorPoseGoal(SceneEntityPose("tracked")), binding=ActionBinding(owner_id="test-engine", endpoints=endpoints), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -945,7 +999,6 @@ def test_action_plan_owns_commands_and_optional_joint_trajectory() -> None: commands, plan_success=plan_success, joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, ) payload = commands.frames[0].commands[0].payload assert isinstance(payload, JointPositionPayload) @@ -1061,7 +1114,8 @@ def test_action_plan_allows_timed_commands_without_joint_trajectory() -> None: assert plan.commands.frame_count == 1 assert plan.joint_trajectory is None - assert plan.feedback_mode is ExecutionFeedbackMode.TIMED + assert isinstance(plan.tracking_policy.terminal, TimedTerminalAcceptance) + assert plan.tracking is None def test_action_plan_rejects_command_device_mismatch() -> None: @@ -1103,7 +1157,6 @@ def test_action_plan_validates_joint_trajectory_against_commands( _action_plan( commands, joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, ) @@ -1122,7 +1175,54 @@ def test_joint_position_plan_rejects_empty_commands_for_successful_rows() -> Non commands, plan_success=torch.tensor([True]), joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, + tracking_policy=TrackingPolicy.joint_position(), + tracking=_joint_tracking_sequence(commands), + ) + + +@pytest.mark.parametrize("changed_route", ["source", "projector"]) +def test_tracking_plan_rejects_route_changes_between_frames( + changed_route: str, +) -> None: + commands = _command_sequence( + env_ids=torch.tensor([4], dtype=torch.long), + frame_count=2, + ) + tracking = _joint_tracking_sequence(commands) + first_frame, second_frame = tracking.frames + original = second_frame.setpoints[0] + source = original.binding.source + projector = original.binding.projector + if changed_route == "source": + source = TrackingFeedbackSourceRef( + provider_id=source.provider_id, + revision="alternate", + address=source.address, + ) + else: + projector = TrackingProjectorRef( + projector_id=projector.projector_id, + revision="alternate", + ) + changed = TrackingSetpoint( + endpoint_key=original.endpoint_key, + binding=EndpointTrackingChannelBinding( + channel_id=original.binding.channel_id, + source=source, + projector=projector, + ), + desired=original.desired, + ) + changed_tracking = TimedTrackingSequence( + commands.env_ids, + (first_frame, TrackingFrame((changed,))), + ) + + with pytest.raises(ValueError, match="source fingerprint and projector route"): + _action_plan( + commands, + tracking_policy=TrackingPolicy.joint_position(), + tracking=changed_tracking, ) @@ -1140,19 +1240,14 @@ def test_joint_position_plan_allows_empty_commands_when_all_rows_fail() -> None: commands, plan_success=torch.tensor([False]), joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, + tracking_policy=TrackingPolicy.joint_position(), + tracking=_joint_tracking_sequence(commands), ) assert plan.commands.frame_count == 0 -@pytest.mark.parametrize( - "feedback_mode", - [ExecutionFeedbackMode.TIMED, ExecutionFeedbackMode.JOINT_POSITION], -) -def test_action_plan_requires_stable_destination_set( - feedback_mode: ExecutionFeedbackMode, -) -> None: +def test_action_plan_requires_stable_destination_set() -> None: env_ids = torch.tensor([4], dtype=torch.long) commands = _command_sequence( env_ids=env_ids, @@ -1162,22 +1257,8 @@ def test_action_plan_requires_stable_destination_set( JointPositionTarget("other_arm", (0, 1)), ), ) - trajectory = ( - TimedTrajectory.from_positions( - torch.tensor([[[1.0, 1.0], [2.0, 2.0]]]), - env_ids=env_ids, - control_dt=0.1, - ) - if feedback_mode is ExecutionFeedbackMode.JOINT_POSITION - else None - ) - with pytest.raises(ValueError, match="same destination set"): - _action_plan( - commands, - joint_trajectory=trajectory, - feedback_mode=feedback_mode, - ) + _action_plan(commands) def test_action_plan_requires_stable_exact_target_type() -> None: @@ -1210,87 +1291,6 @@ def test_action_plan_requires_stable_target_address_fingerprint() -> None: _action_plan(commands) -def test_joint_position_plan_rejects_joint_ids_outside_trajectory() -> None: - env_ids = torch.tensor([4], dtype=torch.long) - commands = _command_sequence( - env_ids=env_ids, - frame_count=1, - targets=(JointPositionTarget("arm", (0, 2)),), - ) - trajectory = TimedTrajectory.from_positions( - torch.ones(1, 1, 2), - env_ids=env_ids, - control_dt=0.1, - ) - - with pytest.raises(ValueError, match="outside joint_trajectory robot_dof"): - _action_plan( - commands, - joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, - ) - - -def test_joint_position_plan_rejects_payload_position_mismatch() -> None: - env_ids = torch.tensor([4], dtype=torch.long) - commands = _command_sequence(env_ids=env_ids, frame_count=1) - trajectory = TimedTrajectory.from_positions( - torch.zeros(1, 1, 2), - env_ids=env_ids, - control_dt=0.1, - ) - - with pytest.raises(ValueError, match="positions.*exactly match"): - _action_plan( - commands, - joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, - ) - - -def test_joint_position_plan_rejects_payload_velocity_presence_mismatch() -> None: - env_ids = torch.tensor([4], dtype=torch.long) - commands = _command_sequence( - env_ids=env_ids, - frame_count=1, - velocities=(torch.zeros(1, 2),), - ) - trajectory = TimedTrajectory.from_positions( - torch.ones(1, 1, 2), - env_ids=env_ids, - control_dt=0.1, - ) - - with pytest.raises(ValueError, match="same presence"): - _action_plan( - commands, - joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, - ) - - -def test_joint_position_plan_rejects_payload_velocity_value_mismatch() -> None: - env_ids = torch.tensor([4], dtype=torch.long) - commands = _command_sequence( - env_ids=env_ids, - frame_count=1, - velocities=(torch.zeros(1, 2),), - ) - trajectory = TimedTrajectory.from_positions( - torch.ones(1, 1, 2), - velocities=torch.ones(1, 1, 2), - env_ids=env_ids, - control_dt=0.1, - ) - - with pytest.raises(ValueError, match="velocities.*exactly match"): - _action_plan( - commands, - joint_trajectory=trajectory, - feedback_mode=ExecutionFeedbackMode.JOINT_POSITION, - ) - - def test_scene_snapshot_expands_global_collision_world_revision() -> None: pose = torch.eye(4).repeat(2, 1, 1) snapshot = SceneSnapshot( diff --git a/tests/sim/atomic_actions/test_endpoint_runtime_e2e.py b/tests/sim/atomic_actions/test_endpoint_runtime_e2e.py index 7c6dd3688..2e15c5ea5 100644 --- a/tests/sim/atomic_actions/test_endpoint_runtime_e2e.py +++ b/tests/sim/atomic_actions/test_endpoint_runtime_e2e.py @@ -52,6 +52,7 @@ SkillResourceSlot, TaskState, TimedCommandSequence, + TrackingPolicy, ) from embodichain.lab.sim.atomic_actions.invocation import ResolvedActionRequest from embodichain.lab.sim.planners import PlanResult @@ -485,6 +486,7 @@ def test_custom_planar_velocity_endpoint_runs_from_profile_through_router() -> N skill_id="drive_velocity", goal=_DriveGoal(goal_twist), binding=binding, + tracking_policy=TrackingPolicy.timed(), ) clock = _Clock() provider = _Provider(robot, clock) diff --git a/tests/sim/atomic_actions/test_engine_per_env.py b/tests/sim/atomic_actions/test_engine_per_env.py index f71f2179f..b81c96e33 100644 --- a/tests/sim/atomic_actions/test_engine_per_env.py +++ b/tests/sim/atomic_actions/test_engine_per_env.py @@ -39,6 +39,8 @@ EndEffectorPoseGoal, EndpointBinding, EndpointCommand, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, EntityState, ExecutionEventKind, ExecutionSession, @@ -66,7 +68,13 @@ StateDelta, TaskState, TimedCommandSequence, + TimedTrackingSequence, TimedTrajectory, + TrackingFeedbackSourceRef, + TrackingFrame, + TrackingPolicy, + TrackingProjectorRef, + TrackingSetpoint, ) from embodichain.lab.sim.common import BatchEntity from embodichain.lab.sim.atomic_actions.goals import resolve_pose_goal @@ -310,9 +318,14 @@ class DestinationSequenceAction(AtomicAction[EndEffectorPoseGoal, ActionOptions] ) ) - def __init__(self, destinations: tuple[str | None, ...]) -> None: + def __init__( + self, + destinations: tuple[str | None, ...], + tracking_provider_revisions: tuple[str | None, ...] | None = None, + ) -> None: super().__init__() self.destinations = destinations + self.tracking_provider_revisions = tracking_provider_revisions self.plan_count = 0 def _plan( @@ -358,7 +371,7 @@ def _plan( device=context.robot.qpos.device, ), ) - return self.build_command_plan( + plan = self.build_command_plan( request, context, success=True, @@ -367,6 +380,35 @@ def _plan( env_ids=context.env_ids, ), ) + if self.tracking_provider_revisions is None: + return plan + provider_revision = self.tracking_provider_revisions[index] + if provider_revision is None or plan.tracking is None: + return plan + original = plan.tracking.frames[0].setpoints[0] + changed = TrackingSetpoint( + endpoint_key=original.endpoint_key, + binding=EndpointTrackingChannelBinding( + channel_id=original.binding.channel_id, + source=TrackingFeedbackSourceRef( + provider_id=original.binding.source.provider_id, + revision=provider_revision, + address=original.binding.source.address, + ), + projector=TrackingProjectorRef( + projector_id=original.binding.projector.projector_id, + revision=original.binding.projector.revision, + ), + ), + desired=original.desired, + ) + return replace( + plan, + tracking=TimedTrackingSequence( + plan.tracking.env_ids, + (TrackingFrame((changed,)),), + ), + ) class UncopyableEntity(BatchEntity): @@ -413,6 +455,7 @@ def _engine(batch_size: int = 1) -> tuple[AtomicActionEngine, DynamicAction]: def _destination_engine( destinations: tuple[str | None, ...], + tracking_provider_revisions: tuple[str | None, ...] | None = None, ) -> tuple[AtomicActionEngine, DestinationSequenceAction]: robot = Mock() robot.device = torch.device("cpu") @@ -430,7 +473,7 @@ def _destination_engine( generator.planner.cfg.planner_type = "stub" generator.supports_dynamic_collision_world = False engine = AtomicActionEngine(generator, load_builtins=False) - action = DestinationSequenceAction(destinations) + action = DestinationSequenceAction(destinations, tracking_provider_revisions) engine.register(action) return engine, action @@ -560,7 +603,6 @@ def _invocation( recovery_policy=RecoveryPolicy( max_replans=max_replans, max_action_retries=max_action_retries, - tracking_error_threshold=0.05, goal_translation_threshold=0.02, action_timeout=action_timeout, ), @@ -908,6 +950,7 @@ def test_request_snapshot_preserves_live_entity_identity() -> None: goal=goal, binding=ActionBinding(owner_id="snapshot-test"), motion_policy=MotionPolicy(), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy(), skill_options=ActionOptions(), ) @@ -1076,6 +1119,23 @@ def test_empty_failed_replan_preserves_destination_for_same_target_retry() -> No assert resumed.command.commands[0].target.target_id == "arm_a" +def test_empty_failed_replan_does_not_erase_active_tracking_route() -> None: + engine, action = _destination_engine( + ("first", None, "first"), + tracking_provider_revisions=("1", None, "alternate"), + ) + invocation = _destination_invocation(engine) + initial = _context(0.0, 0.0, 0.1, 0) + session = engine.start((invocation,), initial) + + session.tick(initial) + + with pytest.raises(ValueError, match="tracking source fingerprints"): + session.tick(_context(0.1, 0.0, 0.3, 1)) + + assert action.plan_count == 3 + + def test_collision_world_change_replans_with_latest_obstacle_pose() -> None: engine, action = _engine() generator = engine.motion_generator @@ -1535,6 +1595,7 @@ def test_session_revision_rejects_changed_target_address_fingerprint() -> None: owner_id=invocation.binding.owner_id, endpoints=(changed_endpoint,), ), + tracking_policy=TrackingPolicy.timed(), revision=1, ) @@ -1549,6 +1610,36 @@ def test_session_revision_rejects_changed_target_address_fingerprint() -> None: assert target.joint_ids == (0, 1) +def test_tracking_continuity_rejection_leaves_revision_state_transactional( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine, _ = _engine() + invocation = _invocation(engine) + initial = _context(0.0, 0.0, 0.1, 0) + replacement_context = _context(0.5, 0.2, 0.1, 0) + session = engine.start((invocation,), initial) + revised = replace(invocation, revision=1) + attempt_count = len(session.plan_attempts) + + with monkeypatch.context() as scoped: + scoped.setattr( + session, + "_validate_tracking_continuity", + Mock(side_effect=ValueError("tracking route changed")), + ) + with pytest.raises(ValueError, match="tracking route changed"): + session.revise_current(revised, context=replacement_context) + + assert len(session.plan_attempts) == attempt_count + assert session.active_plan.invocation_revision == 0 + assert session.latest_context.robot.timestamp == pytest.approx(0.0) + + session.revise_current(revised, context=replacement_context) + + assert session.active_plan.invocation_revision == 1 + assert session.latest_context.robot.timestamp == pytest.approx(0.5) + + def test_tracking_error_fails_when_replan_budget_is_zero() -> None: engine, _ = _engine() session = engine.start( @@ -1560,7 +1651,7 @@ def test_tracking_error_fails_when_replan_budget_is_zero() -> None: tick = session.tick(_context(0.1, 1.0, 0.2, 0)) kinds = {event.kind for event in tick.events} - assert ExecutionEventKind.TRACKING_ERROR in kinds + assert ExecutionEventKind.TRACKING_DIVERGED in kinds assert ExecutionEventKind.RECOVERY_EXHAUSTED in kinds assert tick.status is ExecutionStatus.FAILED assert tick.eligible_mask.tolist() == [False] diff --git a/tests/sim/atomic_actions/test_runner.py b/tests/sim/atomic_actions/test_runner.py index 58dda7380..12bec40ac 100644 --- a/tests/sim/atomic_actions/test_runner.py +++ b/tests/sim/atomic_actions/test_runner.py @@ -45,9 +45,11 @@ HeldObjectState, JOINT_POSITION_CAPABILITY, JointPositionPayload, + JointPositionTrackingMetric, JointPositionTarget, MotionPolicy, ObjectSemantics, + PlanningContextTrackingFeedbackProvider, PlanningContext, RecoveryPolicy, ResolvedActionRequest, @@ -62,6 +64,14 @@ StateDelta, TaskState, TimedTrajectory, + TrackingEvaluation, + TrackingEvaluatorRegistry, + TrackingFeedbackBatch, + TrackingFeedbackProviderRegistry, + TrackingFeedbackSourceRef, + TrackingMetricCfg, + TrackingRuntime, + TrackingState, ) BATCH_SIZE = 1 @@ -185,6 +195,66 @@ def cancel( return CommandAcknowledgement.accepted_ack() +class RaisingFeedbackProvider: + """Built-in-source replacement that simulates a provider failure.""" + + provider_id = "planning_context.robot" + revision = "1" + + def observe( + self, + source: TrackingFeedbackSourceRef, + context: PlanningContext, + ) -> TrackingFeedbackBatch: + """Raise instead of returning required feedback.""" + del source, context + raise RuntimeError("provider unavailable") + + +class RaisingJointTrackingEvaluator: + """Joint evaluator replacement that simulates an evaluation failure.""" + + metric_id = JointPositionTrackingMetric.metric_id + revision = JointPositionTrackingMetric.revision + metric_type = JointPositionTrackingMetric + + def evaluate( + self, + desired: TrackingState, + observed: TrackingState, + valid_mask: torch.Tensor, + metric: TrackingMetricCfg, + ) -> TrackingEvaluation: + """Raise instead of evaluating required feedback.""" + del desired, observed, valid_mask, metric + raise RuntimeError("evaluator unavailable") + + +class MaskedFeedbackProvider(PlanningContextTrackingFeedbackProvider): + """Context provider exposing a deterministic per-row validity mask.""" + + def __init__(self, valid_mask: tuple[bool, ...]) -> None: + self.valid_mask = valid_mask + + def observe( + self, + source: TrackingFeedbackSourceRef, + context: PlanningContext, + ) -> TrackingFeedbackBatch: + """Return built-in feedback with selected rows marked invalid.""" + feedback = super().observe(source, context) + return TrackingFeedbackBatch( + source=feedback.source, + state=feedback.state, + valid_mask=torch.tensor( + self.valid_mask, + dtype=torch.bool, + device=feedback.state.device, + ), + timestamp=feedback.timestamp, + ) + + class TimedAction(AtomicAction[EndEffectorPoseGoal, ActionOptions]): """Test action with explicit non-uniform command intervals.""" @@ -269,6 +339,7 @@ def _make_runner( control_joint_ids: tuple[int, ...] | None = None, max_action_retries: int = 2, action_timeout: float = 10.0, + tracking_runtime: TrackingRuntime | None = None, ) -> tuple[ ExecutionRunner, FakeClock, @@ -292,7 +363,7 @@ def _make_runner( generator.device = torch.device("cpu") generator.planner.cfg.planner_type = "stub" action = TimedAction(with_effect=with_effect) - engine = AtomicActionEngine(generator) + engine = AtomicActionEngine(generator, tracking_runtime=tracking_runtime) engine.register(action) initial_task = TaskState.empty(batch_size, "cpu") initial_context = provider.observe(initial_task) @@ -306,7 +377,6 @@ def _make_runner( recovery_policy=RecoveryPolicy( max_replans=2, max_action_retries=max_action_retries, - tracking_error_threshold=0.05, action_timeout=action_timeout, ), ) @@ -356,7 +426,7 @@ def test_joint_feedback_ignores_motion_outside_bound_endpoint() -> None: assert action.plan_count == 1 assert len(sink.sent) == 3 assert not any( - event.kind is ExecutionEventKind.TRACKING_ERROR + event.kind is ExecutionEventKind.TRACKING_DIVERGED for step in (second, completed) if step.tick is not None for event in step.tick.events @@ -507,11 +577,131 @@ def test_runner_replans_from_observation_after_tracking_error() -> None: assert action.plan_count == 2 assert recovered.tick is not None event_kinds = {event.kind for event in recovered.tick.events} - assert ExecutionEventKind.TRACKING_ERROR in event_kinds + assert ExecutionEventKind.TRACKING_DIVERGED in event_kinds assert ExecutionEventKind.REPLANNED in event_kinds assert recovered.status is RunnerStatus.RUNNING +@pytest.mark.parametrize("failure_kind", ["provider", "evaluator"]) +def test_runner_fails_closed_when_required_tracking_runtime_raises( + failure_kind: str, +) -> None: + builtins = TrackingRuntime.with_builtins() + if failure_kind == "provider": + tracking_runtime = TrackingRuntime( + TrackingFeedbackProviderRegistry((RaisingFeedbackProvider(),)), + builtins.projectors, + builtins.evaluators, + ) + else: + tracking_runtime = TrackingRuntime( + builtins.providers, + builtins.projectors, + TrackingEvaluatorRegistry((RaisingJointTrackingEvaluator(),)), + ) + runner, clock, _, sink, action = _make_runner(tracking_runtime=tracking_runtime) + + runner.step() + clock.advance(FIRST_INTERVAL) + failed = runner.step() + + assert failed.status is RunnerStatus.FAILED + assert failed.tick is not None + event_kinds = {event.kind for event in failed.tick.events} + assert ExecutionEventKind.TRACKING_FEEDBACK_FAILED in event_kinds + assert ExecutionEventKind.REPLANNED not in event_kinds + assert action.plan_count == 1 + assert sink.cancel_count == 1 + + +def test_runner_deactivates_only_rows_with_invalid_required_feedback() -> None: + builtins = TrackingRuntime.with_builtins() + tracking_runtime = TrackingRuntime( + TrackingFeedbackProviderRegistry((MaskedFeedbackProvider((True, False)),)), + builtins.projectors, + builtins.evaluators, + ) + runner, clock, _, _, _ = _make_runner( + batch_size=2, + tracking_runtime=tracking_runtime, + ) + + runner.step() + clock.advance(2.0 * FIRST_INTERVAL) + partial = runner.step() + + assert partial.status is RunnerStatus.RUNNING + assert partial.tick is not None + assert partial.tick.command is not None + assert partial.tick.command.active_mask.tolist() == [True, False] + feedback_failure = next( + event + for event in partial.tick.events + if event.kind is ExecutionEventKind.TRACKING_FEEDBACK_FAILED + ) + assert feedback_failure.env_mask.tolist() == [False, True] + + +def test_runner_maintains_final_target_while_terminal_acceptance_is_pending() -> None: + runner, clock, _, sink, action = _make_runner() + sink.follow_commands.extend([True, True, False]) + + runner.step() + clock.advance(FIRST_INTERVAL) + runner.step() + clock.advance(SECOND_INTERVAL) + runner.step() + final_command = sink.sent[-1] + + clock.advance(SECOND_INTERVAL) + settling = runner.step() + + assert action.plan_count == 1 + assert settling.status is RunnerStatus.RUNNING + assert settling.tick is not None + assert settling.tick.command is not None + assert len(sink.sent) == 4 + assert sink.sent[-1] is settling.tick.command + assert torch.equal(sink.sent[-1].active_mask, final_command.active_mask) + final_payload = final_command.commands[0].payload + settling_payload = sink.sent[-1].commands[0].payload + assert isinstance(final_payload, JointPositionPayload) + assert isinstance(settling_payload, JointPositionPayload) + assert torch.equal(settling_payload.positions, final_payload.positions) + event_kinds = {event.kind for event in settling.tick.events} + assert ExecutionEventKind.TERMINAL_ACCEPTANCE_PENDING in event_kinds + assert ExecutionEventKind.REPLANNED not in event_kinds + + +def test_terminal_settle_reemits_final_target_only_for_pending_rows() -> None: + runner, clock, provider, sink, action = _make_runner(batch_size=2) + + runner.step() + clock.advance(2.0 * FIRST_INTERVAL) + runner.step() + clock.advance(2.0 * SECOND_INTERVAL) + runner.step() + provider.qpos[1].zero_() + + clock.advance(2.0 * SECOND_INTERVAL) + settling = runner.step() + + assert action.plan_count == 1 + assert settling.status is RunnerStatus.RUNNING + assert settling.tick is not None + assert settling.tick.command is not None + assert settling.tick.command.active_mask.tolist() == [False, True] + pending = next( + event + for event in settling.tick.events + if event.kind is ExecutionEventKind.TERMINAL_ACCEPTANCE_PENDING + ) + assert pending.env_mask.tolist() == [False, True] + assert not any( + event.kind is ExecutionEventKind.REPLANNED for event in settling.tick.events + ) + + def test_runner_revision_waits_for_deadline_and_plans_from_fresh_observation() -> None: runner, clock, provider, sink, action = _make_runner() first = runner.step() @@ -526,7 +716,6 @@ def test_runner_revision_waits_for_deadline_and_plans_from_fresh_observation() - motion_policy=MotionPolicy(sample_count=3, control_dt=FIRST_INTERVAL), recovery_policy=RecoveryPolicy( max_replans=2, - tracking_error_threshold=0.05, action_timeout=10.0, ), revision=1, @@ -573,7 +762,6 @@ def test_runner_revision_rejects_pending_effect_verification() -> None: motion_policy=MotionPolicy(sample_count=3, control_dt=FIRST_INTERVAL), recovery_policy=RecoveryPolicy( max_replans=2, - tracking_error_threshold=0.05, action_timeout=10.0, ), revision=1, diff --git a/tests/sim/atomic_actions/test_tracking.py b/tests/sim/atomic_actions/test_tracking.py new file mode 100644 index 000000000..3bff07f3e --- /dev/null +++ b/tests/sim/atomic_actions/test_tracking.py @@ -0,0 +1,263 @@ +# ---------------------------------------------------------------------------- +# 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 dataclasses import dataclass +from typing import ClassVar + +import pytest +import torch + +from embodichain.lab.sim.atomic_actions.bindings import JointPositionTarget +from embodichain.lab.sim.atomic_actions.runtime_commands import ( + EndpointCommand, + JointPositionPayload, +) +from embodichain.lab.sim.atomic_actions.state import ( + PlanningContext, + RobotObservation, + SceneSnapshot, + TaskState, +) +from embodichain.lab.sim.atomic_actions.tracking import ( + BASE_POSE_CHANNEL, + JOINT_POSITION_CHANNEL, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, + FeedbackTerminalAcceptance, + InFlightTrackingPolicy, + JointPositionTrackingMetric, + JointPositionTrackingState, + PoseTrackingEvaluator, + PoseTrackingMetric, + PoseTrackingState, + TimedTerminalAcceptance, + TimedTrackingSequence, + TrackingFeedbackSourceRef, + TrackingFrame, + TrackingMetricCfg, + TrackingPolicy, + TrackingProjectorRef, + TrackingRuntime, + TrackingSetpoint, + WholeBodyPoseTrackingEvaluator, + WholeBodyPoseTrackingMetric, + WholeBodyPoseTrackingState, +) + + +@dataclass(frozen=True, slots=True) +class _AlternateJointMetric(TrackingMetricCfg): + """Different metric identity deliberately sharing the joint channel.""" + + metric_id: ClassVar[str] = "joint.alternate" + channel_id: ClassVar[str] = JOINT_POSITION_CHANNEL + + +def _joint_binding(target: JointPositionTarget) -> EndpointTrackingChannelBinding: + return EndpointTrackingChannelBinding( + channel_id=JOINT_POSITION_CHANNEL, + source=TrackingFeedbackSourceRef( + provider_id="planning_context.robot", + revision="1", + address=EndpointTrackingFeedbackAddress( + target=target, + channel_id=JOINT_POSITION_CHANNEL, + ), + ), + projector=TrackingProjectorRef( + projector_id="joint_position_payload", + revision="1", + ), + ) + + +def _context(qpos: torch.Tensor) -> PlanningContext: + batch_size = qpos.shape[0] + device = qpos.device + return PlanningContext( + robot=RobotObservation( + timestamp=1.0, + qpos=qpos, + qvel=torch.zeros_like(qpos), + root_pose=torch.eye(4, device=device).repeat(batch_size, 1, 1), + ), + task=TaskState.empty(batch_size=batch_size, device=device), + scene=SceneSnapshot.empty(), + env_ids=torch.arange(batch_size, dtype=torch.long, device=device), + ) + + +def test_joint_policy_factory_separates_in_flight_and_terminal_contracts() -> None: + policy = TrackingPolicy.joint_position( + in_flight_max_abs_error=0.1, + terminal_max_abs_error=0.08, + terminal_settle_timeout=0.25, + ) + + assert policy.in_flight is not None + assert policy.in_flight.metrics == (JointPositionTrackingMetric(0.1),) + assert isinstance(policy.terminal, FeedbackTerminalAcceptance) + assert policy.terminal.metrics == (JointPositionTrackingMetric(0.08),) + assert policy.terminal.settle_timeout == pytest.approx(0.25) + + +def test_policy_rejects_ambiguous_metric_id_for_a_shared_channel() -> None: + with pytest.raises(ValueError, match="same exact metric ID"): + TrackingPolicy( + in_flight=InFlightTrackingPolicy( + metrics=(JointPositionTrackingMetric(0.1),) + ), + terminal=FeedbackTerminalAcceptance(metrics=(_AlternateJointMetric(),)), + ) + + +def test_timed_policy_is_an_explicit_no_feedback_contract() -> None: + policy = TrackingPolicy.timed(settle_duration=0.2) + + assert policy.in_flight is None + assert isinstance(policy.terminal, TimedTerminalAcceptance) + assert policy.terminal.settle_duration == pytest.approx(0.2) + + +def test_tracking_values_and_routes_own_tensor_and_target_snapshots() -> None: + positions = torch.tensor([[0.1, 0.2]]) + target = JointPositionTarget(control_part="arm", joint_ids=(0, 1)) + setpoint = TrackingSetpoint( + endpoint_key=("arm", "controller"), + binding=_joint_binding(target), + desired=JointPositionTrackingState(positions), + ) + + positions.add_(1.0) + + assert torch.equal(setpoint.desired.positions, torch.tensor([[0.1, 0.2]])) + assert setpoint.binding.source.address.target is not target + assert setpoint.key == ("arm", "controller", JOINT_POSITION_CHANNEL) + + +def test_timed_tracking_sequence_owns_env_ids_and_validates_batches() -> None: + env_ids = torch.tensor([2, 5], dtype=torch.long) + target = JointPositionTarget(control_part="arm", joint_ids=(0, 1)) + frame = TrackingFrame( + ( + TrackingSetpoint( + endpoint_key=("arm", "controller"), + binding=_joint_binding(target), + desired=JointPositionTrackingState(torch.zeros(2, 2)), + ), + ) + ) + sequence = TimedTrackingSequence(env_ids=env_ids, frames=(frame,)) + + env_ids[0] = 99 + + assert sequence.env_ids.tolist() == [2, 5] + assert sequence.batch_size == 2 + assert sequence.frame_count == 1 + + +def test_timed_tracking_sequence_rejects_mismatched_setpoint_batch() -> None: + target = JointPositionTarget(control_part="arm", joint_ids=(0, 1)) + frame = TrackingFrame( + ( + TrackingSetpoint( + endpoint_key=("arm", "controller"), + binding=_joint_binding(target), + desired=JointPositionTrackingState(torch.zeros(1, 2)), + ), + ) + ) + + with pytest.raises(ValueError, match="setpoint batch"): + TimedTrackingSequence( + env_ids=torch.tensor([0, 1], dtype=torch.long), + frames=(frame,), + ) + + +def test_builtin_runtime_projects_observes_and_evaluates_joint_positions() -> None: + target = JointPositionTarget(control_part="arm", joint_ids=(1, 3)) + binding = _joint_binding(target) + command = EndpointCommand( + target=target, + payload=JointPositionPayload(positions=torch.tensor([[0.3, 0.5], [0.1, 0.2]])), + ) + runtime = TrackingRuntime.with_builtins() + desired = runtime.project(command, binding) + setpoint = TrackingSetpoint(("arm", "controller"), binding, desired) + context = _context( + torch.tensor( + [ + [0.0, 0.32, 0.0, 0.49], + [0.0, 0.25, 0.0, 0.2], + ] + ) + ) + + feedback = runtime.observe(setpoint, context) + evaluation = runtime.evaluate( + setpoint, + feedback, + JointPositionTrackingMetric(tolerance=0.05), + ) + + assert torch.equal(evaluation.accepted_mask, torch.tensor([True, False])) + assert torch.allclose( + evaluation.component_errors["joint_max_abs"], + torch.tensor([0.02, 0.15]), + ) + + +def test_pose_metric_preserves_translation_and_rotation_components() -> None: + desired = torch.eye(4).repeat(2, 1, 1) + observed = desired.clone() + observed[0, 0, 3] = 0.01 + observed[1, :2, :2] = torch.tensor([[0.0, -1.0], [1.0, 0.0]]) + evaluator = PoseTrackingEvaluator() + + evaluation = evaluator.evaluate( + PoseTrackingState(desired), + PoseTrackingState(observed), + torch.ones(2, dtype=torch.bool), + PoseTrackingMetric(translation_tolerance=0.02, rotation_tolerance=0.1), + ) + + assert evaluation.channel_id == BASE_POSE_CHANNEL + assert evaluation.accepted_mask.tolist() == [True, False] + assert set(evaluation.component_errors) == {"translation", "rotation"} + + +def test_whole_body_metric_requires_pose_and_joint_acceptance() -> None: + root = torch.eye(4).repeat(2, 1, 1) + desired = WholeBodyPoseTrackingState(root, torch.zeros(2, 2)) + observed = WholeBodyPoseTrackingState( + root, + torch.tensor([[0.01, 0.0], [0.0, 0.2]]), + ) + + evaluation = WholeBodyPoseTrackingEvaluator().evaluate( + desired, + observed, + torch.ones(2, dtype=torch.bool), + WholeBodyPoseTrackingMetric(joint_position_tolerance=0.05), + ) + + assert evaluation.accepted_mask.tolist() == [True, False] + assert torch.allclose( + evaluation.component_errors["joint_max_abs"], torch.tensor([0.01, 0.2]) + ) diff --git a/tests/sim/skills/test_compiler.py b/tests/sim/skills/test_compiler.py index 5ad526016..8e5eb6d2e 100644 --- a/tests/sim/skills/test_compiler.py +++ b/tests/sim/skills/test_compiler.py @@ -50,6 +50,10 @@ SkillDescriptor, TaskState, ) +from embodichain.lab.sim.atomic_actions.tracking import ( + JointPositionTrackingMetric, + TrackingPolicy, +) from embodichain.lab.sim.skills.calls import ( HandOver, Pick, @@ -932,6 +936,10 @@ def test_grounded_safe_invocation_requires_registered_dynamic_collision() -> Non preset=SkillPolicyPreset( "safe", motion_policy=MotionPolicy(strategy="motion_gen"), + tracking_policy=TrackingPolicy.joint_position( + in_flight_max_abs_error=0.125, + terminal_max_abs_error=0.125, + ), ) ) compiler, engine = _compiler( @@ -951,6 +959,14 @@ def test_grounded_safe_invocation_requires_registered_dynamic_collision() -> Non engine.resolve(grounded.invocation).motion_policy.dynamic_collision_mode is DynamicCollisionMode.REQUIRED ) + invocation_tracking = grounded.invocation.tracking_policy.in_flight + resolved_tracking = engine.resolve(grounded.invocation).tracking_policy.in_flight + assert invocation_tracking is not None + assert resolved_tracking is not None + assert isinstance(invocation_tracking.metrics[0], JointPositionTrackingMetric) + assert isinstance(resolved_tracking.metrics[0], JointPositionTrackingMetric) + assert invocation_tracking.metrics[0].tolerance == 0.125 + assert resolved_tracking.metrics[0].tolerance == 0.125 def test_pick_relation_lookahead_stays_late_bound_scene_dependency() -> None: diff --git a/tests/sim/skills/test_curobo_semantic_runtime_dynamic_recovery_gpu.py b/tests/sim/skills/test_curobo_semantic_runtime_dynamic_recovery_gpu.py index cfdb3ca3a..0aa252c9f 100644 --- a/tests/sim/skills/test_curobo_semantic_runtime_dynamic_recovery_gpu.py +++ b/tests/sim/skills/test_curobo_semantic_runtime_dynamic_recovery_gpu.py @@ -46,6 +46,7 @@ SimulationExecutionAdapter, SkillDescriptor, ) +from embodichain.lab.sim.atomic_actions.tracking import TrackingPolicy # noqa: E402 from embodichain.lab.sim.cfg import RigidBodyAttributesCfg # noqa: E402 from embodichain.lab.sim.objects import RigidObjectCfg # noqa: E402 from embodichain.lab.sim.planners import MotionGenCfg, MotionGenerator # noqa: E402 @@ -185,9 +186,12 @@ def _profile() -> RobotSkillProfile: sample_count=SAMPLE_COUNT, control_dt=COMMAND_CYCLE_TIME, ), + tracking_policy=TrackingPolicy.joint_position( + in_flight_max_abs_error=0.1, + terminal_max_abs_error=0.1, + ), recovery_policy=RecoveryPolicy( max_replans=2, - tracking_error_threshold=0.1, action_timeout=30.0, ), runner_cfg=ExecutionRunnerCfg(minimum_cycle_time=COMMAND_CYCLE_TIME), diff --git a/tests/sim/skills/test_profiles.py b/tests/sim/skills/test_profiles.py index f50666d3b..5400a54ad 100644 --- a/tests/sim/skills/test_profiles.py +++ b/tests/sim/skills/test_profiles.py @@ -55,6 +55,12 @@ RuntimeEndpointTarget, ) from embodichain.lab.sim.atomic_actions.state import PlanningContext +from embodichain.lab.sim.atomic_actions.tracking import ( + JOINT_POSITION_CHANNEL, + EndpointTrackingFeedbackAddress, + JointPositionTrackingMetric, + TrackingPolicy, +) from embodichain.lab.sim.skills import ( AmbiguousSkillBindingError, COMPOSITE_EFFECT_MONITOR_ID, @@ -1150,6 +1156,19 @@ def test_unique_capability_binding_lowers_to_exact_action_binding() -> None: assert grasp.require_target(JointPositionTarget).control_part == "left_hand" assert motion.task_state_key == "left_actor" assert grasp.task_state_key == "left_actor" + motion_tracking = motion.tracking_channel(JOINT_POSITION_CHANNEL) + assert motion_tracking.source.provider_id == "planning_context.robot" + assert motion_tracking.source.revision == "1" + assert motion_tracking.projector.projector_id == "joint_position_payload" + assert motion_tracking.projector.revision == "1" + assert isinstance( + motion_tracking.source.address, + EndpointTrackingFeedbackAddress, + ) + assert ( + motion_tracking.source.address.target.address_fingerprint + == motion.target.address_fingerprint + ) resource = resolved.resources["primary"] motion_sources = resource.endpoints["motion"].effect_sources grasp_sources = resource.endpoints["grasp"].effect_sources @@ -1362,6 +1381,10 @@ def test_presets_are_versioned_snapshots_and_validate_planner() -> None: preset = SkillPolicyPreset( "safe", motion_policy=MotionPolicy(planner="stub_planner", sample_count=80), + tracking_policy=TrackingPolicy.joint_position( + in_flight_max_abs_error=0.125, + terminal_max_abs_error=0.125, + ), ) profile = RobotSkillProfile( "presets", @@ -1379,6 +1402,11 @@ def test_presets_are_versioned_snapshots_and_validate_planner() -> None: assert first is not second assert first.schema_version == 1 assert first.motion_policy.sample_count == 80 + assert first.tracking_policy is not second.tracking_policy + first_tracking = first.tracking_policy.in_flight + assert first_tracking is not None + assert isinstance(first_tracking.metrics[0], JointPositionTrackingMetric) + assert first_tracking.metrics[0].tolerance == 0.125 mutable_runner = first.runner_cfg mutable_runner.command_timeout = 99.0 assert bound.preset().runner_cfg.command_timeout == 1.0 diff --git a/tests/sim/skills/test_runtime.py b/tests/sim/skills/test_runtime.py index 098f0d7cf..128f23d13 100644 --- a/tests/sim/skills/test_runtime.py +++ b/tests/sim/skills/test_runtime.py @@ -39,6 +39,8 @@ EffectVerificationRequirement, EffectVerificationRequest, EndpointBinding, + EndpointTrackingChannelBinding, + EndpointTrackingFeedbackAddress, JointPositionTarget, MotionPolicy, PlanningContext, @@ -50,7 +52,10 @@ StateDelta, TaskState, TimedCommandSequence, + TrackingFeedbackSourceRef, + TrackingProjectorRef, ) +from embodichain.lab.sim.atomic_actions.tracking import TrackingPolicy from embodichain.lab.sim.skills.calls import RegisteredSemanticCall from embodichain.lab.sim.skills.compiler import SemanticSkillCompiler from embodichain.lab.sim.skills.effects import ( @@ -345,6 +350,7 @@ def ground( velocity_limit=0.4, acceleration_limit=0.8, ), + tracking_policy=TrackingPolicy.timed(), recovery_policy=RecoveryPolicy( max_replans=0, max_action_retries=0, @@ -627,9 +633,16 @@ def test_result_metadata_is_json_safe_and_contains_typed_runtime_trace() -> None assert resolved["motion_policy"]["strategy"] == "ik_interp" assert resolved["motion_policy"]["planner"] == "runtime_test" assert resolved["motion_policy"]["sample_count"] == 7 + assert resolved["tracking_policy"] == { + "in_flight": None, + "terminal": {"mode": "timed", "settle_duration": 0.0}, + } assert resolved["recovery_policy"]["max_replans"] == 0 assert resolved["endpoints"] == [] assert attempt["resolved_core_policy"] == resolved + assert attempt["tracking_policy"] == resolved["tracking_policy"] + assert attempt["tracking_contract"] is None + assert "feedback_mode" not in attempt assert result.calls[0].resolved_core_policy.preset_id == "runtime_test_preset" effect = call["effects"][0] assert effect["effect_spec"]["semantic_id"] == "test.metadata" @@ -670,16 +683,34 @@ def test_plan_attempt_trace_rejects_monitor_cutoff_for_non_dependency() -> None: def test_endpoint_binding_trace_records_only_stable_binding_choices() -> None: + target = JointPositionTarget("left_arm_control", (3, 1)) binding = EndpointBinding( slot_id="primary", endpoint_id="motion", resource_id="left_arm", adapter_id="control_part", - target=JointPositionTarget("left_arm_control", (3, 1)), + target=target, task_state_key="left_arm_state", capabilities=frozenset({"cartesian_pose", "joint_position"}), claim_tokens=frozenset({"arm_workspace", "left_side"}), joint_ids=(3, 1), + tracking_channels={ + "joint.position": EndpointTrackingChannelBinding( + channel_id="joint.position", + source=TrackingFeedbackSourceRef( + provider_id="planning_context.robot", + revision="1", + address=EndpointTrackingFeedbackAddress( + target=target, + channel_id="joint.position", + ), + ), + projector=TrackingProjectorRef( + projector_id="joint_position_payload", + revision="1", + ), + ) + }, ) trace = SkillEndpointBindingTrace.from_binding(binding) @@ -694,6 +725,32 @@ def test_endpoint_binding_trace_records_only_stable_binding_choices() -> None: assert metadata["claim_tokens"] == ["arm_workspace", "left_side"] assert metadata["joint_ids"] == [3, 1] assert "target" not in metadata + tracking = metadata["tracking_channels"][0] + target_fingerprint = [ + { + "__type__": ( + "embodichain.lab.sim.atomic_actions.bindings." "JointPositionTarget" + ) + }, + "robot.joint_position", + "left_arm_control", + [3, 1], + ] + address_fingerprint = [target_fingerprint, "joint.position"] + assert tracking["feedback_source"]["address_fingerprint"] == address_fingerprint + assert tracking["route_fingerprint"] == [ + "joint.position", + ["planning_context.robot", "1", address_fingerprint], + "joint_position_payload", + "1", + ] + tracking["feedback_source"]["address_fingerprint"][0][1] = "mutated" + assert ( + trace.to_metadata()["tracking_channels"][0]["feedback_source"][ + "address_fingerprint" + ][0][1] + == "robot.joint_position" + ) def test_preparation_failure_keeps_resolved_policy_without_plan_attempt(