Skip to content

Commit a2a362d

Browse files
committed
style: optimize typing and linting
1 parent 3c9a786 commit a2a362d

2 files changed

Lines changed: 28 additions & 53 deletions

File tree

‎python/rcs/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -126,7 +126,7 @@ class RobotMetaConfig:
126126
}
127127

128128
GRIPPER_OFFSETS: dict[common.GripperType, common.Pose] = {
129-
common.GripperType.FrankaHand: common.FrankaHandTCPOffset(),
129+
common.GripperType.FrankaHand: common.Pose(pose_matrix=common.FrankaHandTCPOffset()),
130130
common.GripperType("Robotiq2F85"): common.Pose(translation=np.array([0.0, 0.0, 0.1628])),
131131
}
132132

‎python/rcs/lerobot_joint_converter.py‎

Lines changed: 27 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -2,15 +2,15 @@
22

33
from dataclasses import dataclass
44
from pathlib import Path
5-
from typing import Iterable
5+
from typing import Any, Iterable
66

77
import duckdb
88
import numpy as np
99
import pandas as pd
1010
import pyarrow as pa
1111
import torch
1212
from lerobot.datasets.lerobot_dataset import LeRobotDataset
13-
from rcs._core.common import RobotType
13+
from rcs._core.common import GripperType, RobotType
1414
from torchvision.io import decode_jpeg
1515
from torchvision.transforms import v2
1616

@@ -56,35 +56,6 @@ def image_column(self) -> str:
5656
DEFAULT_PER_ROBOT_ARM_DIM = 7
5757

5858

59-
def _resolve_robot_type(robot_type: str | RobotType) -> RobotType:
60-
candidates = list(rcs.ROBOTS.keys())
61-
if isinstance(robot_type, RobotType):
62-
for candidate in candidates:
63-
if candidate == robot_type or str(candidate).lower() == str(robot_type).lower():
64-
return candidate
65-
return robot_type
66-
67-
normalized = robot_type.lower()
68-
for candidate in candidates:
69-
candidate_text = str(candidate).lower()
70-
if candidate_text == normalized or normalized in candidate_text:
71-
return candidate
72-
73-
msg = f"Unknown robot type '{robot_type}'"
74-
raise ValueError(msg)
75-
76-
77-
def _resolve_gripper_type(gripper_type: str) -> rcs.common.GripperType:
78-
try:
79-
return rcs.common.GripperType(gripper_type)
80-
except Exception:
81-
attr = getattr(rcs.common.GripperType, gripper_type, None)
82-
if attr is None:
83-
msg = f"Unknown gripper type '{gripper_type}'"
84-
raise ValueError(msg)
85-
return attr
86-
87-
8859
def parse_camera_spec(spec: str) -> CamConversionConfig:
8960
name_source, _, resolution_spec = spec.partition("@")
9061
name, sep, source_name = name_source.partition(":")
@@ -116,13 +87,13 @@ class JointDatasetConverter:
11687
def __init__(
11788
self,
11889
root: str | Path,
119-
dataset_paths: list[str | Path] | None = None,
90+
robot_type: RobotType,
91+
gripper_type: GripperType,
92+
dataset_paths: list[str] | None = None,
12093
repo_id: str = DEFAULT_REPO_ID,
121-
robot_type: str = DEFAULT_ROBOT_TYPE,
12294
fps: int = DEFAULT_FPS,
12395
robot_keys: list[str] | None = None,
12496
joints: bool = DEFAULT_JOINTS,
125-
gripper_type: str = DEFAULT_GRIPPER_TYPE,
12697
cameras: list[CamConversionConfig] | None = None,
12798
image_batch_size: int = DEFAULT_IMAGE_BATCH_SIZE,
12899
per_robot_arm_dim: int = DEFAULT_PER_ROBOT_ARM_DIM,
@@ -135,18 +106,18 @@ def __init__(
135106
self.fps = fps
136107
self.robot_keys = robot_keys or list(DEFAULT_ROBOT_KEYS)
137108
self.joints = joints
138-
self.gripper_type = _resolve_gripper_type(gripper_type)
109+
self.gripper_type = gripper_type
139110
self.cameras = cameras or list(DEFAULT_CAMERAS)
140111
self.image_batch_size = image_batch_size
141112
self.per_robot_arm_dim = per_robot_arm_dim
142113
self.per_robot_state_dim = self.per_robot_arm_dim + 1
143114
self.state_dim = len(self.robot_keys) * self.per_robot_state_dim
144115
self.source_sql = self._build_source_sql(self.dataset_paths)
145-
resolved_robot_type = _resolve_robot_type(self.robot_type)
116+
146117
self.tcp_offset = rcs.GRIPPER_OFFSETS[self.gripper_type]
147118
self.ik = rcs.common.Pin(
148-
rcs.ROBOTS[resolved_robot_type].mjcf_model_path,
149-
rcs.ROBOTS[resolved_robot_type].attachment_site,
119+
rcs.ROBOTS[robot_type].mjcf_model_path,
120+
rcs.ROBOTS[robot_type].attachment_site,
150121
)
151122
self.camera_resizers = {camera.name: v2.Resize(camera.resolution) for camera in self.cameras}
152123

@@ -161,7 +132,7 @@ def __init__(
161132
image_writer_processes=0,
162133
)
163134

164-
def _build_features(self) -> dict[str, dict[str, object]]:
135+
def _build_features(self) -> dict[str, dict[str, Any]]:
165136
state_names = []
166137
for robot_key in self.robot_keys:
167138
state_names.extend([f"{robot_key}_joint_{i}" for i in range(self.per_robot_arm_dim)])
@@ -187,7 +158,7 @@ def _build_features(self) -> dict[str, dict[str, object]]:
187158
}
188159
return features
189160

190-
def _build_source_sql(self, dataset_paths: list[str | Path]) -> str:
161+
def _build_source_sql(self, dataset_paths: list[str]) -> str:
191162
queries = []
192163
for path in dataset_paths:
193164
escaped = str(path).replace("'", "''")
@@ -235,12 +206,12 @@ def _fetch_transition_table(self, episode_id: str) -> pd.DataFrame:
235206
).df()
236207

237208
def _fetch_episode_success(self, episode_id: str) -> bool:
238-
return bool(
239-
self.conn.execute(
240-
f"SELECT COALESCE(MAX(success), FALSE) FROM ({self.source_sql}) AS src WHERE uuid = ?",
241-
[episode_id],
242-
).fetchone()[0]
243-
)
209+
success = self.conn.execute(
210+
f"SELECT COALESCE(MAX(success), FALSE) FROM ({self.source_sql}) AS src WHERE uuid = ?",
211+
[episode_id],
212+
).fetchone()
213+
assert success is not None
214+
return bool(success[0])
244215

245216
def _image_query(self) -> str:
246217
image_selects = ",\n ".join(
@@ -332,7 +303,9 @@ def _convert_action_to_joint_space(self, row: pd.Series) -> np.ndarray:
332303
translation=absolute_action_vec[:3],
333304
quaternion=absolute_action_vec[3:7],
334305
)
335-
ik_joints = self.ik.inverse(target_pose, observation_joints_vec, tcp_offset=self.tcp_offset)
306+
ik_joints: np.ndarray | None = self.ik.inverse(
307+
target_pose, observation_joints_vec, tcp_offset=self.tcp_offset
308+
)
336309
if ik_joints is None:
337310
msg = f"IK failed for robot '{robot_key}' at step {row['step']}"
338311
raise ValueError(msg)
@@ -350,7 +323,7 @@ def _prepare_transition_table(self, table: pd.DataFrame) -> pd.DataFrame:
350323
if len(table) == 0:
351324
return table
352325

353-
df = table.copy()
326+
df = table.copy() # noqa: PD901
354327
df["observation_state"] = df.apply(self._build_observation_state, axis=1)
355328
df["action_vector"] = df.apply(self._convert_action_to_joint_space, axis=1)
356329

@@ -389,7 +362,7 @@ def parse_episode(self, episode_id: str, table: pd.DataFrame, success: bool):
389362
continue
390363
images = frames_by_step[step]
391364

392-
frame = {camera.dataset_key: images[camera.name] for camera in self.cameras}
365+
frame: dict[str, Any] = {camera.dataset_key: images[camera.name] for camera in self.cameras}
393366
frame["observation.state"] = curr["observation_state"]
394367
frame["action"] = curr["action_vector"]
395368
frame["task"] = str(curr["instruction"])
@@ -424,7 +397,7 @@ def _decode_image_batch(self, batch: pa.RecordBatch, frames_by_step: dict[int, d
424397

425398
def run_conversion(
426399
root: str | Path = DEFAULT_HF_DATA_DIR,
427-
dataset_paths: list[str | Path] | None = None,
400+
dataset_paths: list[str] | None = None,
428401
repo_id: str = DEFAULT_REPO_ID,
429402
robot_type: str = DEFAULT_ROBOT_TYPE,
430403
fps: int = DEFAULT_FPS,
@@ -437,15 +410,17 @@ def run_conversion(
437410
success: bool = True,
438411
n: int = -1,
439412
) -> None:
413+
robot_type_converted = RobotType(robot_type)
414+
gripper_type_converted = GripperType(gripper_type)
440415
converter = JointDatasetConverter(
441416
root=root,
417+
robot_type=robot_type_converted,
418+
gripper_type=gripper_type_converted,
442419
dataset_paths=dataset_paths,
443420
repo_id=repo_id,
444-
robot_type=robot_type,
445421
fps=fps,
446422
robot_keys=robot_keys,
447423
joints=joints,
448-
gripper_type=gripper_type,
449424
cameras=cameras,
450425
image_batch_size=image_batch_size,
451426
per_robot_arm_dim=per_robot_arm_dim,

0 commit comments

Comments
 (0)