22
33from dataclasses import dataclass
44from pathlib import Path
5- from typing import Iterable
5+ from typing import Any , Iterable
66
77import duckdb
88import numpy as np
99import pandas as pd
1010import pyarrow as pa
1111import torch
1212from lerobot .datasets .lerobot_dataset import LeRobotDataset
13- from rcs ._core .common import RobotType
13+ from rcs ._core .common import GripperType , RobotType
1414from torchvision .io import decode_jpeg
1515from torchvision .transforms import v2
1616
@@ -56,35 +56,6 @@ def image_column(self) -> str:
5656DEFAULT_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-
8859def 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
425398def 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