@@ -153,22 +153,40 @@ def export_episode_videos(
153153 conn = duckdb .connect ()
154154 relation = conn .sql (f"SELECT * FROM read_parquet('{ source_escaped } ')" )
155155 frame_struct = relation .select ("obs.frames" ).types [0 ]
156+ action_struct = relation .select ("action" ).types [0 ]
156157 camera_names = [name for name , _ in frame_struct .children ]
157158 robot_names = [name for name , _ in relation .select ("obs" ).types [0 ].children if name != "frames" ]
159+ action_fields_by_robot = {
160+ robot : {field_name for field_name , _ in robot_struct .children }
161+ for robot , robot_struct in action_struct .children
162+ }
158163
159164 uuids = conn .execute (f"SELECT DISTINCT uuid FROM read_parquet('{ source_escaped } ') ORDER BY uuid" ).fetchall ()
160165 for index , (episode_id ,) in enumerate (uuids ):
161166 if n != - 1 and index >= n :
162167 break
163168
164169 image_selects = ", " .join (f"obs.frames.{ camera } .rgb.data AS { camera } " for camera in camera_names )
170+ joint_selects = ", " .join (
171+ (
172+ f"COALESCE(action.{ robot } .joints, obs.{ robot } .joints) AS joints_{ robot } "
173+ if "joints" in action_fields_by_robot .get (robot , set ())
174+ else f"obs.{ robot } .joints AS joints_{ robot } "
175+ )
176+ for robot in robot_names
177+ )
178+ gripper_selects = ", " .join (
179+ (
180+ f"COALESCE(CAST(action.{ robot } .gripper[1] AS DOUBLE), obs.{ robot } .gripper[1]) AS gripper_{ robot } "
181+ if "gripper" in action_fields_by_robot .get (robot , set ())
182+ else f"obs.{ robot } .gripper[1] AS gripper_{ robot } "
183+ )
184+ for robot in robot_names
185+ )
165186 state_selects = ", " .join (
166187 [
167- * (f"obs.{ robot } .joints AS joints_{ robot } " for robot in robot_names ),
168- * (
169- f"COALESCE(CAST(action.{ robot } .gripper[1] AS DOUBLE), obs.{ robot } .gripper[1]) AS gripper_{ robot } "
170- for robot in robot_names
171- ),
188+ joint_selects ,
189+ gripper_selects ,
172190 ]
173191 )
174192 not_null_checks = " " .join (f"AND obs.frames.{ camera } .rgb.data IS NOT NULL" for camera in camera_names )
0 commit comments