Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
96 changes: 96 additions & 0 deletions src/spatialcf/adapters/ai2thor/capture.py
Original file line number Diff line number Diff line change
Expand Up @@ -694,6 +694,100 @@ def _stable_instance_pixel_counts(
counts[obj.object_id] = int(np.count_nonzero(mask))
return counts

def _stable_instance_colors(
self,
scene: Scene,
event: Any,
) -> tuple[tuple[str, tuple[int, int, int]], ...]:
raw_objects = event.metadata.get("objects")
if not isinstance(raw_objects, list):
raise AI2ThorNativeReturnError("observation returned no object metadata")
all_native_ids = {
item.get("objectId")
for item in raw_objects
if isinstance(item, dict) and type(item.get("objectId")) is str
}
domain_objects = _domain_object_metadata(raw_objects)
native_by_name = {
str(item.get("name")): item for item in domain_objects if isinstance(item, dict)
}
if len(native_by_name) != len(domain_objects):
raise AI2ThorNativeReturnError("observation returned duplicate object names")
native_colors = getattr(event, "object_id_to_color", None)
if not isinstance(native_colors, Mapping):
raise AI2ThorNativeReturnError("observation returned invalid instance colors")
colors: dict[str, tuple[int, int, int]] = {}
for native_id, color in native_colors.items():
if (
type(native_id) is not str
or type(color) is not tuple
or len(color) != 3
or any(
type(channel) is not int or not 0 <= channel <= 255
for channel in color
)
):
raise AI2ThorNativeReturnError(
"observation returned invalid instance colors"
)
colors[native_id] = color
stable_by_native: dict[str, str] = {}
for obj in scene.objects:
metadata = native_by_name.get(obj.name)
if metadata is None or type(metadata.get("objectId")) is not str:
raise AI2ThorNativeReturnError(
"observation returned invalid instance colors"
)
stable_by_native[metadata["objectId"]] = obj.object_id
counts = self._stable_instance_pixel_counts(scene, event)
masks = event.instance_masks
if any(
native_id not in all_native_ids and np.count_nonzero(mask) > 0
for native_id, mask in masks.items()
):
raise AI2ThorNativeReturnError(
"observation returned unknown instance colors"
)
stable_colors: dict[str, tuple[int, int, int]] = {}
native_by_stable = {
stable_id: native_id for native_id, stable_id in stable_by_native.items()
}
for stable_id, count in counts.items():
native_id = native_by_stable[stable_id]
color = colors.get(native_id)
if count > 0 and color is None:
raise AI2ThorNativeReturnError(
"observation returned incomplete instance colors"
)
if color is None:
continue
mask = masks.get(native_id)
if mask is None:
mask = np.zeros((self.height, self.width), dtype=bool)
instance = np.asarray(
getattr(event, "instance_segmentation_frame", None)
)
png_mask = np.all(
instance == np.asarray(color, dtype=np.uint8), axis=2
)
if not np.array_equal(png_mask, mask):
raise AI2ThorNativeReturnError(
"observation instance colors disagree with PNG"
)
if count > 0:
stable_colors[stable_id] = color
if set(stable_colors) != {
object_id for object_id, count in counts.items() if count > 0
}:
raise AI2ThorNativeReturnError(
"observation returned incomplete instance colors"
)
if len(set(stable_colors.values())) != len(stable_colors):
raise AI2ThorNativeReturnError(
"observation returned duplicate instance colors"
)
return tuple(sorted(stable_colors.items()))

def _observation_from_event(
self,
scene: Scene,
Expand All @@ -709,6 +803,7 @@ def _observation_from_event(
pointcloud_ply=self._pointcloud_bytes(camera, depth, rgb),
instance_pixel_counts=self._stable_instance_pixel_counts(scene, event),
is_scene_at_rest=self._native_scene_at_rest(event),
instance_colors=self._stable_instance_colors(scene, event),
)

def capture_current_observation(self, scene: Scene) -> AI2ThorObservation:
Expand Down Expand Up @@ -862,6 +957,7 @@ class AI2ThorCaptureMixin:
_png_bytes = staticmethod(_png_bytes)
_npy_bytes = staticmethod(_npy_bytes)
_stable_instance_pixel_counts = _stable_instance_pixel_counts
_stable_instance_colors = _stable_instance_colors
_observation_from_event = _observation_from_event
capture_current_observation = capture_current_observation
_validated_frames = _validated_frames
Expand Down
53 changes: 53 additions & 0 deletions src/spatialcf/adapters/ai2thor/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
AppliedCertifiedEdit,
CapturedSource,
CertifiedEditApplication,
InstanceEvidenceProvenance,
SettledReadback,
)
from spatialcf.domain.scene import (
Expand Down Expand Up @@ -613,6 +614,7 @@ class AI2ThorObservation:
pointcloud_ply_sha256: str
instance_pixel_counts: Mapping[str, int]
is_scene_at_rest: bool
instance_colors: tuple[tuple[str, tuple[int, int, int]], ...]

@classmethod
def create(
Expand All @@ -625,6 +627,7 @@ def create(
pointcloud_ply: bytes,
instance_pixel_counts: Mapping[str, int],
is_scene_at_rest: bool,
instance_colors: tuple[tuple[str, tuple[int, int, int]], ...],
) -> AI2ThorObservation:
return cls(
scene=scene,
Expand All @@ -640,8 +643,54 @@ def create(
dict(sorted(instance_pixel_counts.items()))
),
is_scene_at_rest=is_scene_at_rest,
instance_colors=tuple(
(object_id, tuple(color)) for object_id, color in instance_colors
),
)

def __post_init__(self) -> None:
if not isinstance(self.instance_pixel_counts, Mapping) or any(
type(object_id) is not str
or type(count) is not int
or count < 0
for object_id, count in self.instance_pixel_counts.items()
):
raise TypeError("AI2-THOR instance pixel counts must be exact")
if set(self.instance_pixel_counts) != {
obj.object_id for obj in self.scene.objects
}:
raise ValueError(
"AI2-THOR instance pixel counts must cover every scene object"
)
if type(self.instance_colors) is not tuple or any(
type(item) is not tuple
or len(item) != 2
or type(item[0]) is not str
or type(item[1]) is not tuple
or len(item[1]) != 3
or any(
type(channel) is not int or not 0 <= channel <= 255
for channel in item[1]
)
for item in self.instance_colors
):
raise TypeError("AI2-THOR instance colors must be exact RGB pairs")
if (
self.instance_colors != tuple(sorted(self.instance_colors))
or len({item[0] for item in self.instance_colors})
!= len(self.instance_colors)
or len({item[1] for item in self.instance_colors})
!= len(self.instance_colors)
):
raise ValueError("AI2-THOR instance colors must be sorted and unique")
positive_ids = {
object_id
for object_id, count in self.instance_pixel_counts.items()
if count > 0
}
if {item[0] for item in self.instance_colors} != positive_ids:
raise ValueError("AI2-THOR instance colors must cover positive pixels")


@dataclass(frozen=True)
class AI2ThorCameraApplication:
Expand Down Expand Up @@ -1014,6 +1063,10 @@ def adapter_observation_from_native(value: AI2ThorObservation) -> AdapterObserva
pointcloud_ply=value.pointcloud_ply,
instance_pixel_counts=tuple(sorted(value.instance_pixel_counts.items())),
is_settled=value.is_scene_at_rest,
instance_colors=value.instance_colors,
instance_evidence_provenance=(
InstanceEvidenceProvenance.SAME_EVENT_INSTANCE_SEGMENTATION
),
)


Expand Down
97 changes: 96 additions & 1 deletion src/spatialcf/adapters/base.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import json
import math
from dataclasses import asdict, dataclass
from dataclasses import asdict, dataclass, field
from enum import StrEnum
from hashlib import sha256
from pathlib import Path
from types import TracebackType
Expand Down Expand Up @@ -651,6 +652,14 @@ def __post_init__(self) -> None:
)


class InstanceEvidenceProvenance(StrEnum):
"""Whether instance sidecars are historical-neutral or one same-event readback."""

LEGACY_NEUTRAL = "LEGACY_NEUTRAL"
SAME_EVENT_INSTANCE_SEGMENTATION = "SAME_EVENT_INSTANCE_SEGMENTATION"
PUBLICATION_REPLAY_COUNTS_ONLY = "PUBLICATION_REPLAY_COUNTS_ONLY"


@dataclass(frozen=True)
class AdapterObservation:
scene: Scene
Expand All @@ -664,6 +673,15 @@ class AdapterObservation:
pointcloud_ply_sha256: str
instance_pixel_counts: tuple[tuple[str, int], ...]
is_settled: bool
# Derived sidecar: preserve legacy neutral/public equality identity. AI2-THOR
# capture must still provide and validate it; this is not a fallback signal.
instance_colors: tuple[tuple[str, tuple[int, int, int]], ...] = field(
default=(), compare=False
)
instance_evidence_provenance: InstanceEvidenceProvenance = field(
default=InstanceEvidenceProvenance.LEGACY_NEUTRAL,
compare=False,
)

@classmethod
def create(
Expand All @@ -676,6 +694,10 @@ def create(
pointcloud_ply: bytes,
instance_pixel_counts: tuple[tuple[str, int], ...],
is_settled: bool,
instance_colors: tuple[tuple[str, tuple[int, int, int]], ...] = (),
instance_evidence_provenance: InstanceEvidenceProvenance = (
InstanceEvidenceProvenance.LEGACY_NEUTRAL
),
) -> Self:
return cls(
scene=scene,
Expand All @@ -689,6 +711,10 @@ def create(
pointcloud_ply_sha256=sha256(pointcloud_ply).hexdigest(),
instance_pixel_counts=instance_pixel_counts,
is_settled=is_settled,
instance_colors=tuple(
(object_id, tuple(color)) for object_id, color in instance_colors
),
instance_evidence_provenance=instance_evidence_provenance,
)

def __post_init__(self) -> None:
Expand Down Expand Up @@ -739,8 +765,77 @@ def __post_init__(self) -> None:
raise TypeError(
"adapter observation pixel counts must be sorted exact pairs"
)
if self.instance_pixel_counts and {
item[0] for item in self.instance_pixel_counts
} != {obj.object_id for obj in self.scene.objects}:
raise ValueError(
"adapter observation pixel counts must cover every scene object"
)
if type(self.is_settled) is not bool:
raise TypeError("adapter observation settled flag must be exact")
if (
type(self.instance_colors) is not tuple
or any(
type(item) is not tuple
or len(item) != 2
or type(item[0]) is not str
or type(item[1]) is not tuple
or len(item[1]) != 3
or any(
type(channel) is not int or not 0 <= channel <= 255
for channel in item[1]
)
for item in self.instance_colors
)
or self.instance_colors != tuple(sorted(self.instance_colors))
or len({item[0] for item in self.instance_colors})
!= len(self.instance_colors)
or len({item[1] for item in self.instance_colors})
!= len(self.instance_colors)
):
raise TypeError(
"adapter observation instance colors must be sorted exact RGB pairs"
)
positive_ids = {
object_id for object_id, count in self.instance_pixel_counts if count > 0
}
if (
self.instance_colors
and {item[0] for item in self.instance_colors} != positive_ids
):
raise ValueError(
"adapter observation instance colors must cover positive pixels"
)
if type(self.instance_evidence_provenance) is not InstanceEvidenceProvenance:
raise TypeError("adapter observation instance evidence provenance is invalid")
if (
self.instance_evidence_provenance
is InstanceEvidenceProvenance.LEGACY_NEUTRAL
):
if self.instance_pixel_counts or self.instance_colors:
raise ValueError(
"adapter observation legacy neutral evidence requires empty sidecars"
)
elif (
self.instance_evidence_provenance
is InstanceEvidenceProvenance.PUBLICATION_REPLAY_COUNTS_ONLY
):
if (
tuple(item[0] for item in self.instance_pixel_counts)
!= tuple(sorted(obj.object_id for obj in self.scene.objects))
or self.instance_colors
):
raise ValueError(
"adapter observation publication replay requires full counts only"
)
elif (
tuple(item[0] for item in self.instance_pixel_counts)
!= tuple(sorted(obj.object_id for obj in self.scene.objects))
or {item[0] for item in self.instance_colors} != positive_ids
):
raise ValueError(
"adapter observation same-event evidence does not close scene sidecars"
)


@dataclass(frozen=True)
Expand Down
5 changes: 4 additions & 1 deletion src/spatialcf/generation/capture/compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
SourceSurfaceEvidence,
build_competition_native_camera_pose_bank_v2_9_3,
competition_native_camera_pose_bank_sha256_v2_9_3,
competition_native_roster_selection_identity_v2_9,
score_competition_native_camera_scene_v2_9_3,
score_competition_native_editable_camera_scene_v2_9_4,
select_competition_native_camera_score_index_v2_9_3,
Expand Down Expand Up @@ -276,7 +277,9 @@ def _candidate_identity(
"reference_id": reference_id,
"relation_before": relation.value,
"scene_id": source.scene_id,
"source_capture_sha256": capture.source_capture_sha256,
"source_capture_sha256": competition_native_roster_selection_identity_v2_9(
capture
),
"source_id": source.source_id,
"split": source.split,
"subject_id": subject_id,
Expand Down
Loading
Loading