diff --git a/embodichain/lab/sim/skills/__init__.py b/embodichain/lab/sim/skills/__init__.py index 95381c478..2b192950e 100644 --- a/embodichain/lab/sim/skills/__init__.py +++ b/embodichain/lab/sim/skills/__init__.py @@ -69,6 +69,7 @@ EffectEvidenceAddress, EffectEvidenceBatch, EffectEvidenceSourceRef, + EffectExpectationDecision, EffectMonitor, EffectMonitorDecision, EffectMonitorFactory, @@ -260,6 +261,7 @@ "EffectEvidenceQuery", "EffectEvidenceQueryValue", "EffectEvidenceSourceRef", + "EffectExpectationDecision", "EffectMonitor", "EffectMonitorDecision", "EffectMonitorFactory", diff --git a/embodichain/lab/sim/skills/effects.py b/embodichain/lab/sim/skills/effects.py index 852cb99a0..7bcbb4636 100644 --- a/embodichain/lab/sim/skills/effects.py +++ b/embodichain/lab/sim/skills/effects.py @@ -1509,12 +1509,90 @@ def _evidence_metadata( ) +@dataclass(frozen=True, slots=True, eq=False) +class EffectExpectationDecision: + """Per-row outcome for one physical state expectation. + + Rows absent from both ``satisfied_mask`` and ``contradicted_mask`` remain + unresolved. ``inverse_satisfied_mask`` is deliberately stronger than + contradiction: it requires every clause in the expectation group to have + reached its explicit inverse band for the configured consecutive-sample + window. This distinction lets failure reconciliation retain a relation + only from complete inverse evidence rather than from one contradictory + clause. + """ + + expectation_id: str + satisfied_mask: torch.Tensor + contradicted_mask: torch.Tensor + inverse_satisfied_mask: torch.Tensor + + def __post_init__(self) -> None: + _validate_identifier( + self.expectation_id, + field_name="EffectExpectationDecision.expectation_id", + ) + for field_name in ( + "satisfied_mask", + "contradicted_mask", + "inverse_satisfied_mask", + ): + value = getattr(self, field_name) + if not isinstance(value, torch.Tensor): + raise TypeError(f"{field_name} must be a torch.Tensor.") + if value.dtype != torch.bool or value.dim() != 1: + raise ValueError(f"{field_name} must be a one-dimensional bool tensor.") + masks = ( + self.satisfied_mask, + self.contradicted_mask, + self.inverse_satisfied_mask, + ) + if any(value.shape != masks[0].shape for value in masks[1:]): + raise ValueError("Expectation decision masks must have equal shapes.") + if any(value.device != masks[0].device for value in masks[1:]): + raise ValueError("Expectation decision masks must use the same device.") + if (self.satisfied_mask & self.contradicted_mask).any(): + raise ValueError("satisfied_mask and contradicted_mask must not overlap.") + if (self.inverse_satisfied_mask & ~self.contradicted_mask).any(): + raise ValueError( + "inverse_satisfied_mask must be a subset of contradicted_mask." + ) + object.__setattr__(self, "satisfied_mask", self.satisfied_mask.clone()) + object.__setattr__( + self, + "contradicted_mask", + self.contradicted_mask.clone(), + ) + object.__setattr__( + self, + "inverse_satisfied_mask", + self.inverse_satisfied_mask.clone(), + ) + + def snapshot(self) -> EffectExpectationDecision: + """Return an independently owned expectation outcome.""" + return EffectExpectationDecision( + expectation_id=self.expectation_id, + satisfied_mask=self.satisfied_mask, + contradicted_mask=self.contradicted_mask, + inverse_satisfied_mask=self.inverse_satisfied_mask, + ) + + @dataclass(frozen=True, slots=True, eq=False) class EffectMonitorDecision: - """Uncorrelated per-row decision; runtime adds the verification ID.""" + """Uncorrelated aggregate and per-expectation monitor decision. + + When ``expectation_decisions`` is non-empty, the aggregate masks are + authoritative reductions of that current observation: success is the + conjunction of every satisfied mask and failure is the union of every + contradicted mask. This prevents callers from combining expectation + outcomes observed on different ticks. + """ success_mask: torch.Tensor failure_mask: torch.Tensor + expectation_decisions: tuple[EffectExpectationDecision, ...] = () def __post_init__(self) -> None: for field_name in ("success_mask", "failure_mask"): @@ -1529,8 +1607,51 @@ def __post_init__(self) -> None: raise ValueError("Decision masks must use the same device.") if (self.success_mask & self.failure_mask).any(): raise ValueError("Decision masks must not overlap.") + expectation_decisions = tuple(self.expectation_decisions) + if not all( + type(value) is EffectExpectationDecision for value in expectation_decisions + ): + raise TypeError( + "expectation_decisions must contain exact " + "EffectExpectationDecision values." + ) + expectation_ids = [value.expectation_id for value in expectation_decisions] + if len(set(expectation_ids)) != len(expectation_ids): + raise ValueError("Expectation decision IDs must be unique.") + if expectation_decisions: + for value in expectation_decisions: + if value.satisfied_mask.shape != self.success_mask.shape: + raise ValueError( + "Expectation and aggregate decision masks must have " + "equal shapes." + ) + if value.satisfied_mask.device != self.success_mask.device: + raise ValueError( + "Expectation and aggregate decision masks must use the " + "same device." + ) + expected_success = torch.ones_like(self.success_mask) + expected_failure = torch.zeros_like(self.failure_mask) + for value in expectation_decisions: + expected_success &= value.satisfied_mask + expected_failure |= value.contradicted_mask + if not torch.equal(self.success_mask, expected_success): + raise ValueError( + "success_mask must equal the conjunction of expectation " + "satisfied masks." + ) + if not torch.equal(self.failure_mask, expected_failure): + raise ValueError( + "failure_mask must equal the union of expectation " + "contradicted masks." + ) object.__setattr__(self, "success_mask", self.success_mask.clone()) object.__setattr__(self, "failure_mask", self.failure_mask.clone()) + object.__setattr__( + self, + "expectation_decisions", + tuple(value.snapshot() for value in expectation_decisions), + ) class EffectMonitor(ABC): @@ -1876,8 +1997,9 @@ def __init__( self._cfg = cfg self._attempt_generation: int | None = None self._active_env_ids: frozenset[int] = frozenset() - self._success_counts: dict[int, int] = {} - self._failure_counts: dict[int, int] = {} + self._success_counts: dict[tuple[str, int], int] = {} + self._failure_counts: dict[tuple[str, int], int] = {} + self._inverse_success_counts: dict[tuple[str, int], int] = {} self._last_observations: dict[int, tuple[float, int]] = {} @property @@ -1901,6 +2023,7 @@ def _prepare_request(self, request: EffectVerificationRequest) -> None: self._active_env_ids = active_env_ids self._success_counts.clear() self._failure_counts.clear() + self._inverse_success_counts.clear() self._last_observations.clear() return if not active_env_ids.issubset(self._active_env_ids): @@ -1910,14 +2033,19 @@ def _prepare_request(self, request: EffectVerificationRequest) -> None: ) self._active_env_ids = active_env_ids self._success_counts = { - env_id: count - for env_id, count in self._success_counts.items() - if env_id in active_env_ids + key: count + for key, count in self._success_counts.items() + if key[1] in active_env_ids } self._failure_counts = { - env_id: count - for env_id, count in self._failure_counts.items() - if env_id in active_env_ids + key: count + for key, count in self._failure_counts.items() + if key[1] in active_env_ids + } + self._inverse_success_counts = { + key: count + for key, count in self._inverse_success_counts.items() + if key[1] in active_env_ids } self._last_observations = { env_id: observation @@ -1996,9 +2124,13 @@ def _normalize_evidence( raise ValueError("Evidence contains env_ids outside the effect spec.") missing = self._active_env_ids.difference(observed_env_ids) if missing: + expectation_ids = {clause.expectation_id for clause in self._spec.clauses} for env_id in missing: - self._success_counts[env_id] = 0 - self._failure_counts[env_id] = 0 + for expectation_id in expectation_ids: + key = (expectation_id, env_id) + self._success_counts[key] = 0 + self._failure_counts[key] = 0 + self._inverse_success_counts[key] = 0 raise ValueError( "Evidence must cover every active request env_id exactly once; " f"missing {sorted(missing)}. Acquisition failures must be explicit " @@ -2095,8 +2227,6 @@ def observe( requested_at=request.requested_at, deadline=request.deadline, ) - success_mask = torch.zeros_like(request.env_mask) - failure_mask = torch.zeros_like(request.env_mask) spec_rows = { int(env_id): row for row, env_id in enumerate(self._spec.env_ids.detach().cpu().tolist()) @@ -2128,7 +2258,23 @@ def observe( clauses_by_expectation: dict[str, list[EffectClause]] = {} for clause in self._spec.clauses: clauses_by_expectation.setdefault(clause.expectation_id, []).append(clause) - physical_expectation_ids = set(clauses_by_expectation) + physical_expectation_ids = tuple( + expectation.expectation_id + for expectation in self._spec.state_expectations + if expectation.expectation_id in clauses_by_expectation + ) + satisfied_masks = { + expectation_id: torch.zeros_like(request.env_mask) + for expectation_id in physical_expectation_ids + } + contradicted_masks = { + expectation_id: torch.zeros_like(request.env_mask) + for expectation_id in physical_expectation_ids + } + inverse_satisfied_masks = { + expectation_id: torch.zeros_like(request.env_mask) + for expectation_id in physical_expectation_ids + } for evidence_row, env_id in enumerate(observed_env_ids): request_row = request_rows.get(env_id) @@ -2138,8 +2284,6 @@ def observe( continue self._last_observations[env_id] = observation_token spec_row = spec_rows[env_id] - expected_groups = True - contradicted_group = False for expectation_id in physical_expectation_ids: classifications = [ self._classify_clause( @@ -2153,24 +2297,55 @@ def observe( ] group_expected = all(value == 1 for value in classifications) group_contradicted = any(value == -1 for value in classifications) - expected_groups = expected_groups and group_expected - contradicted_group = contradicted_group or group_contradicted - if expected_groups: - self._success_counts[env_id] = self._success_counts.get(env_id, 0) + 1 - self._failure_counts[env_id] = 0 - elif contradicted_group: - self._failure_counts[env_id] = self._failure_counts.get(env_id, 0) + 1 - self._success_counts[env_id] = 0 - else: - self._success_counts[env_id] = 0 - self._failure_counts[env_id] = 0 - if self._success_counts.get(env_id, 0) >= self._cfg.consecutive_samples: - success_mask[request_row] = True - elif self._failure_counts.get(env_id, 0) >= self._cfg.consecutive_samples: - failure_mask[request_row] = True - success_mask &= request.env_mask - failure_mask &= request.env_mask - return EffectMonitorDecision(success_mask, failure_mask) + group_inverse_satisfied = all(value == -1 for value in classifications) + key = (expectation_id, env_id) + if group_expected: + self._success_counts[key] = self._success_counts.get(key, 0) + 1 + else: + self._success_counts[key] = 0 + if group_contradicted: + self._failure_counts[key] = self._failure_counts.get(key, 0) + 1 + else: + self._failure_counts[key] = 0 + if group_inverse_satisfied: + self._inverse_success_counts[key] = ( + self._inverse_success_counts.get(key, 0) + 1 + ) + else: + self._inverse_success_counts[key] = 0 + if self._success_counts.get(key, 0) >= self._cfg.consecutive_samples: + satisfied_masks[expectation_id][request_row] = True + if self._failure_counts.get(key, 0) >= self._cfg.consecutive_samples: + contradicted_masks[expectation_id][request_row] = True + if ( + self._inverse_success_counts.get(key, 0) + >= self._cfg.consecutive_samples + ): + inverse_satisfied_masks[expectation_id][request_row] = True + + expectation_decisions = tuple( + EffectExpectationDecision( + expectation_id=expectation_id, + satisfied_mask=satisfied_masks[expectation_id] & request.env_mask, + contradicted_mask=( + contradicted_masks[expectation_id] & request.env_mask + ), + inverse_satisfied_mask=( + inverse_satisfied_masks[expectation_id] & request.env_mask + ), + ) + for expectation_id in physical_expectation_ids + ) + success_mask = request.env_mask.clone() + failure_mask = torch.zeros_like(request.env_mask) + for decision in expectation_decisions: + success_mask &= decision.satisfied_mask + failure_mask |= decision.contradicted_mask + return EffectMonitorDecision( + success_mask, + failure_mask, + expectation_decisions, + ) class CompositeEffectMonitorFactory(EffectMonitorFactory): @@ -2222,6 +2397,7 @@ def create( "EffectEvidenceAddress", "EffectEvidenceBatch", "EffectEvidenceSourceRef", + "EffectExpectationDecision", "EffectMonitor", "EffectMonitorDecision", "EffectMonitorFactory", diff --git a/tests/gym/envs/tasks/test_multi_segments_cube_pick_place.py b/tests/gym/envs/tasks/test_multi_segments_cube_pick_place.py index 6e6fe9495..5f9dd99f8 100644 --- a/tests/gym/envs/tasks/test_multi_segments_cube_pick_place.py +++ b/tests/gym/envs/tasks/test_multi_segments_cube_pick_place.py @@ -185,7 +185,8 @@ class FakeRobot: uid = "UR5" @staticmethod - def get_qpos() -> torch.Tensor: + def get_qpos(*, target: bool = False) -> torch.Tensor: + del target return torch.zeros((1, 8), dtype=torch.float32) class FakeCube: diff --git a/tests/gym/envs/tasks/test_open_drawer.py b/tests/gym/envs/tasks/test_open_drawer.py index 725d3c77d..acb4293d1 100644 --- a/tests/gym/envs/tasks/test_open_drawer.py +++ b/tests/gym/envs/tasks/test_open_drawer.py @@ -162,7 +162,8 @@ class FakeRobot: uid = "CobotMagic" @staticmethod - def get_qpos() -> torch.Tensor: + def get_qpos(*, target: bool = False) -> torch.Tensor: + del target return torch.zeros((1, 16), dtype=torch.float32) class FakeDrawer: diff --git a/tests/sim/skills/test_effects.py b/tests/sim/skills/test_effects.py index 3794dad03..72e479ea4 100644 --- a/tests/sim/skills/test_effects.py +++ b/tests/sim/skills/test_effects.py @@ -49,6 +49,7 @@ EffectEvidenceAddress, EffectEvidenceBatch, EffectEvidenceSourceRef, + EffectExpectationDecision, EffectMonitor, EffectMonitorDecision, EffectMonitorFactory, @@ -72,8 +73,13 @@ _ENV_IDS = torch.tensor([101, 205, 309], dtype=torch.long) _OBJECT_ID = "scene/cube" _STATE_KEY = "left_actor" +_SOURCE_STATE_KEY = "source_actor" +_DESTINATION_STATE_KEY = "destination_actor" _SKILL_ID = "pick_up" _INVOCATION_ID = "call-7" +_ATTACHED_OFFSET = 0.0 +_DETACHED_OFFSET = 0.1 # Above the built-in 0.06 translation threshold. +_UNRESOLVED_OFFSET = 0.04 # Between the attached and detached thresholds. @dataclass(frozen=True, slots=True) @@ -178,6 +184,120 @@ def _attach_spec() -> SemanticEffectSpec: ) +def _transfer_spec() -> SemanticEffectSpec: + source = _expectation( + HeldObjectRelation.DETACHED, + expectation_id="source", + state_key=_SOURCE_STATE_KEY, + ) + destination = _expectation( + expectation_id="destination", + state_key=_DESTINATION_STATE_KEY, + ) + return SemanticEffectSpec( + semantic_id="hand_over", + effect_kind=SemanticEffectKind.TRANSFER, + skill_id=_SKILL_ID, + invocation_id=_INVOCATION_ID, + invocation_revision=2, + env_ids=_ENV_IDS, + state_expectations=(source, destination), + clauses=( + PoseRelationClause( + "source.pose", + "source", + _source("source_pose_relation"), + PoseRelationExpectation.SEPARATED, + baseline_object_to_endpoint=_poses(0.0, 0.0, 0.0), + ), + BinaryEffectClause( + "source.constraint", + "source", + _source("source_constraint"), + BinaryEvidenceKind.CONSTRAINT, + False, + ), + PoseRelationClause( + "destination.pose", + "destination", + _source("destination_pose_relation"), + PoseRelationExpectation.MATCHED, + ), + BinaryEffectClause( + "destination.constraint", + "destination", + _source("destination_constraint"), + BinaryEvidenceKind.CONSTRAINT, + True, + ), + ), + ) + + +def _transfer_request() -> EffectVerificationRequest: + return _request( + effects=StateDelta( + held_object_updates={ + _SOURCE_STATE_KEY: None, + _DESTINATION_STATE_KEY: _held(), + } + ) + ) + + +def _transfer_evidence( + *, + source_offsets: tuple[float, ...], + source_constraints: tuple[bool, ...], + destination_offsets: tuple[float, ...], + destination_constraints: tuple[bool, ...], + timestamp: float, + revision: int, +) -> Mapping[str, EffectEvidenceBatch]: + valid = torch.ones(len(source_offsets), dtype=torch.bool) + errors = tuple(None for _ in source_offsets) + return { + "source.pose": PoseRelationEvidenceBatch( + "source.pose", + _poses(*source_offsets), + valid, + errors, + timestamp, + _ENV_IDS, + revision, + ), + "source.constraint": BinaryEffectEvidenceBatch( + "source.constraint", + BinaryEvidenceKind.CONSTRAINT, + torch.tensor(source_constraints, dtype=torch.bool), + valid, + errors, + timestamp, + _ENV_IDS, + revision, + ), + "destination.pose": PoseRelationEvidenceBatch( + "destination.pose", + _poses(*destination_offsets), + valid, + errors, + timestamp, + _ENV_IDS, + revision, + ), + "destination.constraint": BinaryEffectEvidenceBatch( + "destination.constraint", + BinaryEvidenceKind.CONSTRAINT, + torch.tensor(destination_constraints, dtype=torch.bool), + valid, + errors, + timestamp, + _ENV_IDS, + revision, + ), + } + + def _request( *, env_mask: torch.Tensor | None = None, @@ -544,6 +664,183 @@ def test_valid_raw_evidence_rejects_nonfinite_payload() -> None: ) +def test_expectation_decision_owns_all_outcome_masks() -> None: + satisfied = torch.tensor([True, False, False]) + contradicted = torch.tensor([False, True, False]) + inverse_satisfied = torch.tensor([False, True, False]) + + decision = EffectExpectationDecision( + "source", + satisfied, + contradicted, + inverse_satisfied, + ) + satisfied.zero_() + contradicted.zero_() + inverse_satisfied.zero_() + + assert decision.satisfied_mask.tolist() == [True, False, False] + assert decision.contradicted_mask.tolist() == [False, True, False] + assert decision.inverse_satisfied_mask.tolist() == [False, True, False] + + aggregate = EffectMonitorDecision( + decision.satisfied_mask, + decision.contradicted_mask, + (decision,), + ) + decision.satisfied_mask.zero_() + assert aggregate.expectation_decisions[0].satisfied_mask.tolist() == [ + True, + False, + False, + ] + + +def test_expectation_decision_requires_complete_inverse_to_be_contradicted() -> None: + with pytest.raises(ValueError, match="subset of contradicted_mask"): + EffectExpectationDecision( + "source", + torch.tensor([False]), + torch.tensor([False]), + torch.tensor([True]), + ) + + +def test_transfer_monitor_reports_each_expectation_and_strong_inverse() -> None: + monitor = CompositeEffectMonitor( + _transfer_spec(), + CompositeEffectMonitorCfg(consecutive_samples=1), + ) + + decision = monitor.observe( + _transfer_request(), + _transfer_evidence( + source_offsets=( + _DETACHED_OFFSET, + _ATTACHED_OFFSET, + _DETACHED_OFFSET, + ), + source_constraints=(False, True, False), + destination_offsets=( + _ATTACHED_OFFSET, + _ATTACHED_OFFSET, + _DETACHED_OFFSET, + ), + destination_constraints=(True, True, False), + timestamp=2.0, + revision=4, + ), + ) + + outcomes = { + outcome.expectation_id: outcome for outcome in decision.expectation_decisions + } + assert tuple(outcomes) == ("source", "destination") + assert outcomes["source"].satisfied_mask.tolist() == [True, False, True] + assert outcomes["source"].contradicted_mask.tolist() == [False, True, False] + assert outcomes["source"].inverse_satisfied_mask.tolist() == [False, True, False] + assert outcomes["destination"].satisfied_mask.tolist() == [True, True, False] + assert outcomes["destination"].contradicted_mask.tolist() == [False, False, True] + assert outcomes["destination"].inverse_satisfied_mask.tolist() == [ + False, + False, + True, + ] + assert decision.success_mask.tolist() == [True, False, False] + assert decision.failure_mask.tolist() == [False, True, True] + + +def test_transfer_contradictions_are_counted_per_expectation() -> None: + monitor = CompositeEffectMonitor( + _transfer_spec(), + CompositeEffectMonitorCfg(consecutive_samples=2), + ) + request = _transfer_request() + monitor.observe( + request, + _transfer_evidence( + source_offsets=(_ATTACHED_OFFSET,) * 3, + source_constraints=(True,) * 3, + destination_offsets=(_ATTACHED_OFFSET,) * 3, + destination_constraints=(True,) * 3, + timestamp=2.0, + revision=4, + ), + ) + + alternating = monitor.observe( + request, + _transfer_evidence( + source_offsets=(_DETACHED_OFFSET,) * 3, + source_constraints=(False,) * 3, + destination_offsets=(_DETACHED_OFFSET,) * 3, + destination_constraints=(False,) * 3, + timestamp=3.0, + revision=5, + ), + ) + persistent = monitor.observe( + request, + _transfer_evidence( + source_offsets=(_DETACHED_OFFSET,) * 3, + source_constraints=(False,) * 3, + destination_offsets=(_DETACHED_OFFSET,) * 3, + destination_constraints=(False,) * 3, + timestamp=4.0, + revision=6, + ), + ) + + assert not alternating.failure_mask.any() + assert persistent.failure_mask.all() + persistent_outcomes = { + outcome.expectation_id: outcome for outcome in persistent.expectation_decisions + } + assert persistent_outcomes["source"].satisfied_mask.all() + assert persistent_outcomes["destination"].contradicted_mask.all() + + +def test_transfer_success_never_stitches_expectations_across_ticks() -> None: + monitor = CompositeEffectMonitor( + _transfer_spec(), + CompositeEffectMonitorCfg(consecutive_samples=1), + ) + request = _transfer_request() + source_only = monitor.observe( + request, + _transfer_evidence( + source_offsets=(_DETACHED_OFFSET,) * 3, + source_constraints=(False,) * 3, + destination_offsets=(_UNRESOLVED_OFFSET,) * 3, + destination_constraints=(True,) * 3, + timestamp=2.0, + revision=4, + ), + ) + destination_only = monitor.observe( + request, + _transfer_evidence( + source_offsets=(_UNRESOLVED_OFFSET,) * 3, + source_constraints=(False,) * 3, + destination_offsets=(_ATTACHED_OFFSET,) * 3, + destination_constraints=(True,) * 3, + timestamp=3.0, + revision=5, + ), + ) + + assert not source_only.success_mask.any() + assert not source_only.failure_mask.any() + assert not destination_only.success_mask.any() + assert not destination_only.failure_mask.any() + destination_outcomes = { + outcome.expectation_id: outcome + for outcome in destination_only.expectation_decisions + } + assert not destination_outcomes["source"].satisfied_mask.any() + assert destination_outcomes["destination"].satisfied_mask.all() + + def test_monitor_requires_pose_and_binary_physical_evidence() -> None: monitor = CompositeEffectMonitor( _attach_spec(), @@ -561,6 +858,9 @@ def test_monitor_requires_pose_and_binary_physical_evidence() -> None: assert not pose_only.success_mask.any() assert pose_only.failure_mask.all() + outcome = pose_only.expectation_decisions[0] + assert outcome.contradicted_mask.all() + assert not outcome.inverse_satisfied_mask.any() def test_monitor_reports_success_only_for_complete_consecutive_evidence() -> None: