se3labs-interface 0.0.1a1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (35) hide show
  1. se3labs_interface-0.0.1a1/PKG-INFO +14 -0
  2. se3labs_interface-0.0.1a1/pyproject.toml +37 -0
  3. se3labs_interface-0.0.1a1/setup.cfg +4 -0
  4. se3labs_interface-0.0.1a1/src/se3labs/interface/__init__.py +0 -0
  5. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/__init__.py +23 -0
  6. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/arm.py +169 -0
  7. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/camera.py +65 -0
  8. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/job.py +22 -0
  9. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/policy.py +32 -0
  10. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/repository.py +27 -0
  11. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/station.py +179 -0
  12. se3labs_interface-0.0.1a1/src/se3labs/interface/eval/task.py +177 -0
  13. se3labs_interface-0.0.1a1/src/se3labs/interface/util/__init__.py +0 -0
  14. se3labs_interface-0.0.1a1/src/se3labs/interface/util/validation.py +16 -0
  15. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/__init__.py +21 -0
  16. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/codec.py +267 -0
  17. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status.proto +28 -0
  18. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2.py +40 -0
  19. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2.pyi +29 -0
  20. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/job_status_pb2_grpc.py +102 -0
  21. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session.proto +243 -0
  22. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2.py +120 -0
  23. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2.pyi +353 -0
  24. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/policy_session_pb2_grpc.py +117 -0
  25. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link.proto +49 -0
  26. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2.py +48 -0
  27. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2.pyi +49 -0
  28. se3labs_interface-0.0.1a1/src/se3labs/interface/wire/station_link_pb2_grpc.py +107 -0
  29. se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/PKG-INFO +14 -0
  30. se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/SOURCES.txt +33 -0
  31. se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/dependency_links.txt +1 -0
  32. se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/requires.txt +10 -0
  33. se3labs_interface-0.0.1a1/src/se3labs_interface.egg-info/top_level.txt +1 -0
  34. se3labs_interface-0.0.1a1/tests/test_spec.py +239 -0
  35. se3labs_interface-0.0.1a1/tests/test_wire.py +148 -0
@@ -0,0 +1,14 @@
1
+ Metadata-Version: 2.4
2
+ Name: se3labs-interface
3
+ Version: 0.0.1a1
4
+ Summary: SE3 Labs SDK shared interface
5
+ Requires-Python: >=3.10
6
+ Requires-Dist: numpy>=2.2
7
+ Requires-Dist: jaxtyping>=0.3
8
+ Requires-Dist: beartype>=0.22
9
+ Requires-Dist: grpcio>=1.83.1
10
+ Requires-Dist: protobuf>=7.35.1
11
+ Requires-Dist: pillow>=10
12
+ Provides-Extra: dev
13
+ Requires-Dist: pytest; extra == "dev"
14
+ Requires-Dist: pyright; extra == "dev"
@@ -0,0 +1,37 @@
1
+ [build-system]
2
+ requires = ["setuptools>=61"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "se3labs-interface"
7
+ version = "0.0.1a1"
8
+ description = "SE3 Labs SDK shared interface"
9
+ requires-python = ">=3.10"
10
+ dependencies = [
11
+ "numpy>=2.2",
12
+ "jaxtyping>=0.3",
13
+ "beartype>=0.22",
14
+ "grpcio>=1.83.1",
15
+ "protobuf>=7.35.1",
16
+ "pillow>=10",
17
+ ]
18
+
19
+ [project.optional-dependencies]
20
+ dev = [
21
+ "pytest",
22
+ "pyright",
23
+ ]
24
+
25
+ [tool.setuptools.packages.find]
26
+ where = ["src"]
27
+
28
+ [tool.setuptools.package-data]
29
+ "se3labs.interface.wire" = ["*.proto", "*.pyi"]
30
+
31
+ [tool.pytest.ini_options]
32
+ testpaths = ["tests"]
33
+
34
+ [tool.pyright]
35
+ include = ["src", "tests"]
36
+ exclude = ["**/__pycache__", "**/*_pb2.py", "**/*_pb2_grpc.py"]
37
+ typeCheckingMode = "standard"
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1,23 @@
1
+ from . import arm, camera, job, policy, station, task
2
+ from .arm import (
3
+ MotorizedArm_ArmType,
4
+ MotorizedArm_GripperType,
5
+ MotorizedArmSpec,
6
+ MotorizedArm_Observation,
7
+ MotorizedArm_JointPose,
8
+ MotorizedArm_ActionChunk,
9
+ )
10
+ from .camera import Camera_DeviceType, Camera_DeviceSpec, Camera_Observation
11
+ from .station import (
12
+ Station_BaseSpec,
13
+ Station_Spec,
14
+ Bimanual_Station_Spec,
15
+ Single_Arm_Station_Spec,
16
+ Station_Observation,
17
+ Station_JointPose,
18
+ Station_ActionChunk,
19
+ )
20
+ from .task import Task_ScoreKind, Task_ScoreChoice, Task_ScoreItem, Task_ScoreAnswers, Task_ScoreSheet, Task_Spec
21
+ from .policy import Policy
22
+ from .job import JobRequest
23
+ from .repository import Bimanual_YAM_Station
@@ -0,0 +1,169 @@
1
+ """Motorized arm hardware spec and the per-arm observation / action containers.
2
+
3
+ Conventions (all arrays are ``np.float32``):
4
+
5
+ * joint angles in radians, joint velocities in rad/s, joint torques in N*m,
6
+ joint temperatures in degrees Celsius.
7
+ * ``gripper_pos`` is a normalized open fraction in ``[0, 1]`` where ``0`` is
8
+ fully closed and ``1`` is fully open. The station maps this to the physical
9
+ gripper travel, so policies never see device units.
10
+ * Action targets are **absolute** joint positions.
11
+ """
12
+ import numpy as np
13
+ import jaxtyping as jt
14
+ from beartype import beartype
15
+ from enum import Enum
16
+ from dataclasses import dataclass
17
+
18
+ from ..util.validation import check_hint
19
+
20
+
21
+ class MotorizedArm_ArmType(Enum):
22
+ YAM_6DoF = "yam_6dof"
23
+ """A YAM 6-DoF arm.
24
+
25
+ Reference: https://i2rt.com/products/yam-6-dof-arm
26
+ """
27
+
28
+ @property
29
+ def n_dof(self) -> int:
30
+ return _ARM_N_DOF[self]
31
+
32
+
33
+ _ARM_N_DOF: dict[MotorizedArm_ArmType, int] = {
34
+ MotorizedArm_ArmType.YAM_6DoF: 6,
35
+ }
36
+
37
+
38
+ class MotorizedArm_GripperType(Enum):
39
+ I2RT_linear_gripper_1DoF = "i2rt_linear_gripper_1dof"
40
+ """I2RT linear gripper with 1 DoF.
41
+
42
+ Reference: https://i2rt.com/products/yam-gripper-1
43
+ """
44
+
45
+ @property
46
+ def n_dof(self) -> int:
47
+ return _GRIPPER_N_DOF[self]
48
+
49
+
50
+ _GRIPPER_N_DOF: dict[MotorizedArm_GripperType, int] = {
51
+ MotorizedArm_GripperType.I2RT_linear_gripper_1DoF: 1,
52
+ }
53
+
54
+
55
+ @dataclass(frozen=True)
56
+ class MotorizedArmSpec:
57
+ arm : MotorizedArm_ArmType
58
+ gripper : MotorizedArm_GripperType
59
+
60
+ refresh_rate: int
61
+ """Rate at which observations are published, unit: [Hz]."""
62
+
63
+ control_rate: int
64
+ """Rate at which joint position targets are sent to the motors, unit: [Hz].
65
+
66
+ Action chunks emitted at a lower rate are interpolated up to this rate by the station.
67
+ """
68
+
69
+ @property
70
+ def n_arm_dof(self) -> int:
71
+ return self.arm.n_dof
72
+
73
+ @property
74
+ def n_gripper_dof(self) -> int:
75
+ return self.gripper.n_dof
76
+
77
+
78
+ @jt.jaxtyped(typechecker=beartype)
79
+ @dataclass(frozen=True)
80
+ class MotorizedArm_Observation:
81
+ """One observation tick of a single arm. Every channel is always populated.
82
+
83
+ Construction enforces dtype / rank and that every joint channel shares ``n_arm_dof``;
84
+ ``validate`` binds ``n_arm_dof`` / ``n_gripper_dof`` to a concrete ``MotorizedArmSpec``.
85
+ """
86
+
87
+ timestamp : float
88
+ """Monotonic capture time on the station clock, unit: [s]."""
89
+
90
+ joint_angle_q : jt.Float32[np.ndarray, "n_arm_dof"]
91
+ """Joint positions, unit: [rad]."""
92
+ joint_angle_vel_dq: jt.Float32[np.ndarray, "n_arm_dof"]
93
+ """Joint velocities, unit: [rad/s]."""
94
+ joint_torque : jt.Float32[np.ndarray, "n_arm_dof"]
95
+ """Joint torques, unit: [N*m]."""
96
+ joint_temperature : jt.Float32[np.ndarray, "n_arm_dof"]
97
+ """Motor temperatures, unit: [degC]."""
98
+ gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
99
+ """Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
100
+
101
+ def validate(self, spec: MotorizedArmSpec) -> None:
102
+ joint = jt.Float32[np.ndarray, f"{spec.n_arm_dof}"]
103
+ check_hint("joint_angle_q" , self.joint_angle_q , joint)
104
+ check_hint("joint_angle_vel_dq", self.joint_angle_vel_dq, joint)
105
+ check_hint("joint_torque" , self.joint_torque , joint)
106
+ check_hint("joint_temperature" , self.joint_temperature , joint)
107
+ check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"])
108
+
109
+
110
+ @jt.jaxtyped(typechecker=beartype)
111
+ @dataclass(frozen=True)
112
+ class MotorizedArm_JointPose:
113
+ """One absolute joint configuration of a single arm, e.g. the pose the arm walks to before an episode."""
114
+
115
+ joint_angle_q: jt.Float32[np.ndarray, "n_arm_dof"]
116
+ """Absolute joint positions, unit: [rad]."""
117
+
118
+ gripper_pos : jt.Float32[np.ndarray, "n_gripper_dof"]
119
+ """Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
120
+
121
+ def __post_init__(self) -> None:
122
+ if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
123
+ raise ValueError("gripper_pos must lie in [0, 1]")
124
+
125
+ def validate(self, spec: MotorizedArmSpec) -> None:
126
+ """Bind the dims to a concrete arm spec."""
127
+ check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{spec.n_arm_dof}"])
128
+ check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{spec.n_gripper_dof}"])
129
+
130
+
131
+ @jt.jaxtyped(typechecker=beartype)
132
+ @dataclass(frozen=True)
133
+ class MotorizedArm_ActionChunk:
134
+ """A horizon of absolute joint position targets for a single arm.
135
+
136
+ Construction enforces dtype / rank and a shared ``horizon`` between joints and gripper;
137
+ ``validate`` binds the dims to a concrete ``MotorizedArmSpec`` and horizon.
138
+
139
+ Row ``i`` is the target ``i / rate`` seconds after the chunk's reference timestamp
140
+ (both carried by the enclosing ``station.Station_ActionChunk``).
141
+ """
142
+
143
+ joint_angle_q: jt.Float32[np.ndarray, "horizon n_arm_dof"]
144
+ """Absolute joint position targets, unit: [rad]."""
145
+
146
+ gripper_pos : jt.Float32[np.ndarray, "horizon n_gripper_dof"]
147
+ """Normalized gripper open fraction in [0, 1]; 0 = closed, 1 = open."""
148
+
149
+ def __post_init__(self) -> None:
150
+ if np.any(self.gripper_pos < 0.0) or np.any(self.gripper_pos > 1.0):
151
+ raise ValueError("gripper_pos must lie in [0, 1]")
152
+
153
+ @property
154
+ def horizon(self) -> int:
155
+ return int(self.joint_angle_q.shape[0])
156
+
157
+ def __len__(self) -> int:
158
+ return self.horizon
159
+
160
+ def __getitem__(self, index: int) -> MotorizedArm_JointPose:
161
+ """Row ``index`` of the chunk as a single joint pose. Negative indices count from the end."""
162
+ if not -self.horizon <= index < self.horizon:
163
+ raise IndexError(f"chunk index {index} out of range for horizon {self.horizon}")
164
+ return MotorizedArm_JointPose(self.joint_angle_q[index], self.gripper_pos[index])
165
+
166
+ def validate(self, spec: MotorizedArmSpec, horizon: int) -> None:
167
+ """Bind the dims to a concrete arm spec and horizon."""
168
+ check_hint("joint_angle_q", self.joint_angle_q, jt.Float32[np.ndarray, f"{horizon} {spec.n_arm_dof}"])
169
+ check_hint("gripper_pos" , self.gripper_pos , jt.Float32[np.ndarray, f"{horizon} {spec.n_gripper_dof}"])
@@ -0,0 +1,65 @@
1
+ import numpy as np
2
+ import jaxtyping as jt
3
+ from enum import Enum
4
+ from dataclasses import dataclass
5
+ from beartype import beartype
6
+
7
+ from ..util.validation import check_hint
8
+
9
+
10
+ class Camera_DeviceType(Enum):
11
+ Realsense_D405 = "realsense_d405"
12
+ """Realsense D405 Stereo Camera. Both imagers are RGB, so one observation carries 2 views (left, right)."""
13
+
14
+ @property
15
+ def n_views(self) -> int:
16
+ """Number of frames delivered per observation (1 for mono, 2 for a stereo pair)."""
17
+ return _CAMERA_N_VIEWS[self]
18
+
19
+
20
+ _CAMERA_N_VIEWS: dict[Camera_DeviceType, int] = {
21
+ Camera_DeviceType.Realsense_D405: 2,
22
+ }
23
+
24
+
25
+ @dataclass(frozen=True)
26
+ class Camera_DeviceSpec:
27
+ device_type : Camera_DeviceType
28
+ """Device type of the camera"""
29
+
30
+ resolution : tuple[int, int]
31
+ """Resolution of each view of the camera stream, in tuple of (height, width), unit: [px]"""
32
+
33
+ refresh_rate : int
34
+ """Refresh rate of the camera, unit: [fps]"""
35
+
36
+ jpeg_compression: int | None
37
+ """jpeg compression quality value if specified, if None then no compression will be made."""
38
+
39
+ @property
40
+ def n_views(self) -> int:
41
+ return self.device_type.n_views
42
+
43
+
44
+ @jt.jaxtyped(typechecker=beartype)
45
+ @dataclass(frozen=True)
46
+ class Camera_Observation:
47
+ """One decoded observation of a camera: ``b`` RGB views captured together (``b = 2`` for a stereo pair).
48
+
49
+ Frames are delivered already decoded regardless of ``jpeg_compression``. Construction enforces
50
+ ``uint8`` BHWC with 3 channels; ``validate`` binds (b, h, w) to the camera spec.
51
+ """
52
+
53
+ timestamp: float
54
+ """Monotonic capture time on the station clock, unit: [s]. Shared by all views."""
55
+
56
+ rgb : jt.UInt8[np.ndarray, "b h w 3"]
57
+ """Views in device order (stereo: index 0 = left, 1 = right)."""
58
+
59
+ @property
60
+ def n_views(self) -> int:
61
+ return int(self.rgb.shape[0])
62
+
63
+ def validate(self, spec: Camera_DeviceSpec) -> None:
64
+ h, w = spec.resolution
65
+ check_hint("rgb", self.rgb, jt.UInt8[np.ndarray, f"{spec.n_views} {h} {w} 3"])
@@ -0,0 +1,22 @@
1
+ """What a customer submits to a station: the task, the policy, and how many
2
+ episodes to run. The station executes the request as uploaded; it does not
3
+ get to say what is being evaluated."""
4
+
5
+ from __future__ import annotations
6
+
7
+ from dataclasses import dataclass
8
+
9
+ from .task import Task_Spec
10
+
11
+ from .policy import Policy
12
+
13
+
14
+ @dataclass(frozen=True)
15
+ class JobRequest:
16
+ task : Task_Spec
17
+ policy : Policy
18
+ episodes: int
19
+
20
+ def __post_init__(self) -> None:
21
+ if self.episodes <= 0:
22
+ raise ValueError("a job needs at least one episode")
@@ -0,0 +1,32 @@
1
+ """The policy contract: what a caller hands to ``evaluate(station, task, policy)``.
2
+
3
+ Lifecycle, driven by the station, once per episode:
4
+
5
+ 1. ``initialize(station, task)``: return the joint pose the arms walk to before the episode starts.
6
+ Load weights and read the task instruction here.
7
+ 2. ``infer(observation)``, repeated on every observation tick until the episode ends. Return a new
8
+ chunk, or ``None`` to keep executing the current one. A new chunk replaces the remainder of the
9
+ current one.
10
+ 3. ``reset()`` after the episode ends. Clear observation history and any unexecuted chunk so the
11
+ next ``initialize`` starts clean.
12
+
13
+ Anything about how the policy runs (its rate, chunk length, replanning, image preprocessing) lives
14
+ inside these methods and is not declared to the SDK.
15
+ """
16
+ from abc import ABC, abstractmethod
17
+
18
+ from .station import Station_Spec, Station_Observation, Station_ActionChunk, Station_JointPose
19
+ from .task import Task_Spec
20
+
21
+
22
+ class Policy(ABC):
23
+ @abstractmethod
24
+ def initialize(self, station: Station_Spec, task: Task_Spec) -> Station_JointPose:
25
+ """Prepare for one episode of ``task`` on ``station`` and return the pose the arms start from."""
26
+
27
+ @abstractmethod
28
+ def infer(self, observation: Station_Observation) -> Station_ActionChunk | None:
29
+ """A new action chunk for this tick, or ``None`` to continue executing the current chunk."""
30
+
31
+ def reset(self) -> None:
32
+ """Clear per-episode state. Called once after the episode ends. No-op by default."""
@@ -0,0 +1,27 @@
1
+ """Registry of known station specs."""
2
+ from . import arm, camera
3
+ from .station import Bimanual_Station_Spec, Station_BaseSpec
4
+
5
+
6
+ Bimanual_YAM_Station = Bimanual_Station_Spec(
7
+ base=Station_BaseSpec.Static,
8
+ arm =arm.MotorizedArmSpec(
9
+ arm =arm.MotorizedArm_ArmType.YAM_6DoF,
10
+ gripper =arm.MotorizedArm_GripperType.I2RT_linear_gripper_1DoF,
11
+ refresh_rate=100,
12
+ control_rate=100,
13
+ ),
14
+ wrist_camera=camera.Camera_DeviceSpec(
15
+ device_type =camera.Camera_DeviceType.Realsense_D405,
16
+ resolution =(360, 640), # (h, w); 2x downsample for 720p.
17
+ refresh_rate =30,
18
+ jpeg_compression=80
19
+ ),
20
+ topdown_camera=camera.Camera_DeviceSpec(
21
+ device_type =camera.Camera_DeviceType.Realsense_D405,
22
+ resolution =(360, 640), # (h, w); 2x downsample for 720p.
23
+ refresh_rate =30,
24
+ jpeg_compression=80
25
+ )
26
+ )
27
+ """Static bimanual YAM station: two YAM 6-DoF arms with I2RT grippers, two wrist D405s, one top-down D405."""
@@ -0,0 +1,179 @@
1
+ """Station specs (hardware capability) and the station-level observation / action containers.
2
+
3
+ A station is a named collection of arms and cameras. Observations and actions are keyed by
4
+ those names so that a bimanual and a single-arm station share one container type.
5
+
6
+ The policy side of the contract is ``policy.Policy``: it consumes ``Station_Observation`` ticks and
7
+ produces ``Station_ActionChunk``s. Anything about how the policy runs (its rate, chunk length,
8
+ replanning, image preprocessing) is decided inside the policy, not declared to the SDK.
9
+ """
10
+ from abc import ABC, abstractmethod
11
+ from enum import Enum
12
+ from dataclasses import dataclass
13
+
14
+ from . import arm as arm_types
15
+ from . import camera
16
+
17
+
18
+ class Station_BaseSpec(Enum):
19
+ Static = "static"
20
+ """Static manipulation base (the arm is mounted on a fixed, non-mobile base)."""
21
+
22
+
23
+ class Station_Spec(ABC):
24
+ base: Station_BaseSpec
25
+
26
+ @property
27
+ @abstractmethod
28
+ def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
29
+ """Arm name -> spec. Names are the keys of ``Station_Observation.arms`` and ``Station_ActionChunk.arms``."""
30
+
31
+ @property
32
+ @abstractmethod
33
+ def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
34
+ """Camera name -> spec. Names are the keys of ``Station_Observation.cameras``."""
35
+
36
+
37
+ @dataclass(frozen=True)
38
+ class Bimanual_Station_Spec(Station_Spec):
39
+ base : Station_BaseSpec
40
+
41
+ arm : arm_types.MotorizedArmSpec
42
+ """Spec shared by the ``left`` and ``right`` arms."""
43
+
44
+ wrist_camera : camera.Camera_DeviceSpec
45
+ """Spec shared by the ``wrist_left`` and ``wrist_right`` cameras."""
46
+
47
+ topdown_camera: camera.Camera_DeviceSpec
48
+
49
+ @property
50
+ def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
51
+ return {"left": self.arm, "right": self.arm}
52
+
53
+ @property
54
+ def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
55
+ return {
56
+ "wrist_left" : self.wrist_camera,
57
+ "wrist_right": self.wrist_camera,
58
+ "topdown" : self.topdown_camera,
59
+ }
60
+
61
+
62
+ @dataclass(frozen=True)
63
+ class Single_Arm_Station_Spec(Station_Spec):
64
+ base : Station_BaseSpec
65
+ arm : arm_types.MotorizedArmSpec
66
+ wrist_camera : camera.Camera_DeviceSpec
67
+
68
+ @property
69
+ def arms(self) -> dict[str, arm_types.MotorizedArmSpec]:
70
+ return {"arm": self.arm}
71
+
72
+ @property
73
+ def cameras(self) -> dict[str, camera.Camera_DeviceSpec]:
74
+ return {"wrist": self.wrist_camera}
75
+
76
+
77
+ @dataclass(frozen=True)
78
+ class Station_Observation:
79
+ """One policy-side observation: the latest tick of every arm and camera.
80
+
81
+ Arms and cameras run at different rates, so each entry carries its own timestamp; consumers that
82
+ need temporal alignment (world models, chunked policies) should use those rather than assume sync.
83
+ """
84
+
85
+ arms : dict[str, arm_types.MotorizedArm_Observation]
86
+ cameras: dict[str, camera.Camera_Observation]
87
+
88
+ def validate(self, station: Station_Spec) -> None:
89
+ """Check every arm / camera of ``station`` is present with the shapes its spec promises."""
90
+ missing_arms = set(station.arms) - set(self.arms)
91
+ if missing_arms:
92
+ raise ValueError(f"observation missing arms: {sorted(missing_arms)}")
93
+ for name, spec in station.arms.items():
94
+ try:
95
+ self.arms[name].validate(spec)
96
+ except (TypeError, ValueError) as e:
97
+ raise type(e)(f"arm '{name}': {e}") from None
98
+
99
+ missing_cams = set(station.cameras) - set(self.cameras)
100
+ if missing_cams:
101
+ raise ValueError(f"observation missing cameras: {sorted(missing_cams)}")
102
+ for name, spec in station.cameras.items():
103
+ try:
104
+ self.cameras[name].validate(spec)
105
+ except (TypeError, ValueError) as e:
106
+ raise type(e)(f"camera '{name}': {e}") from None
107
+
108
+
109
+ @dataclass(frozen=True)
110
+ class Station_JointPose:
111
+ """One absolute joint configuration for every arm of a station."""
112
+
113
+ arms: dict[str, arm_types.MotorizedArm_JointPose]
114
+
115
+ def __post_init__(self) -> None:
116
+ if not self.arms:
117
+ raise ValueError("joint pose has no arms")
118
+
119
+ def validate(self, station: Station_Spec) -> None:
120
+ """Check the pose names exactly the station's arms with the right shapes."""
121
+ if set(self.arms) != set(station.arms):
122
+ raise ValueError(f"pose arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
123
+ for name, spec in station.arms.items():
124
+ try:
125
+ self.arms[name].validate(spec)
126
+ except (TypeError, ValueError) as e:
127
+ raise type(e)(f"arm '{name}': {e}") from None
128
+
129
+
130
+ @dataclass(frozen=True)
131
+ class Station_ActionChunk:
132
+ """A chunk of ``horizon`` absolute joint-position targets for every arm of a station.
133
+
134
+ Row ``i`` of every arm is the target at ``timestamp + i / rate``. The station interpolates rows
135
+ up to each arm's ``control_rate`` and executes them until the policy returns a new chunk, which
136
+ replaces the remainder of the current one.
137
+ """
138
+
139
+ timestamp: float
140
+ """Station-clock time of the observation this chunk was predicted from, unit: [s]."""
141
+
142
+ rate : float
143
+ """Spacing between consecutive rows, unit: [Hz]. Chosen by the policy per chunk."""
144
+
145
+ arms : dict[str, arm_types.MotorizedArm_ActionChunk]
146
+
147
+ def __post_init__(self) -> None:
148
+ if self.rate <= 0:
149
+ raise ValueError(f"rate must be positive, got {self.rate}")
150
+ if not self.arms:
151
+ raise ValueError("action chunk has no arms")
152
+ horizons = {a.horizon for a in self.arms.values()}
153
+ if len(horizons) != 1:
154
+ raise ValueError(f"inconsistent horizons across arms: {horizons}")
155
+
156
+ @property
157
+ def horizon(self) -> int:
158
+ return next(iter(self.arms.values())).horizon
159
+
160
+ def __len__(self) -> int:
161
+ return self.horizon
162
+
163
+ def __getitem__(self, index: int) -> Station_JointPose:
164
+ """Row ``index`` of the chunk, for every arm, as one station pose. Negative indices count from the end."""
165
+ return Station_JointPose({name: chunk[index] for name, chunk in self.arms.items()})
166
+
167
+ def validate(self, station: Station_Spec) -> None:
168
+ """Check the chunk names exactly the station's arms, with the right shapes and a feasible rate."""
169
+ if set(self.arms) != set(station.arms):
170
+ raise ValueError(f"action arms {sorted(self.arms)} != station arms {sorted(station.arms)}")
171
+ h = self.horizon
172
+ for name, spec in station.arms.items():
173
+ if self.rate > spec.control_rate:
174
+ raise ValueError(f"arm '{name}': chunk rate {self.rate} Hz exceeds control_rate {spec.control_rate} Hz")
175
+ try:
176
+ self.arms[name].validate(spec, h)
177
+ except (TypeError, ValueError) as e:
178
+ raise type(e)(f"arm '{name}': {e}") from None
179
+