@@ -133,6 +133,24 @@ def _render_action_panel(
133133 return image
134134
135135
136+ def _episode_starts (conn : duckdb .DuckDBPyConnection , source_escaped : str ) -> list [tuple [str , float ]]:
137+ """Return episodes in recording order, with a stable tie-breaker."""
138+ return conn .execute (
139+ f"""
140+ SELECT uuid, MIN(timestamp) AS start_timestamp
141+ FROM read_parquet('{ source_escaped } ')
142+ GROUP BY uuid
143+ ORDER BY start_timestamp, uuid
144+ """
145+ ).fetchall ()
146+
147+
148+ def _episode_filename (timestamp : float , episode_number : int , camera_name : str | None = None ) -> str :
149+ timestamp_text = datetime .datetime .fromtimestamp (timestamp ).strftime ("%Y-%m-%d-%H-%M-%S" )
150+ filename = f"{ timestamp_text } _episode-{ episode_number :06d} "
151+ return f"{ filename } _{ camera_name } .mp4" if camera_name is not None else f"{ filename } .mp4"
152+
153+
136154def export_episode_videos (
137155 dataset : str | Path ,
138156 output : str | Path ,
@@ -162,8 +180,8 @@ def export_episode_videos(
162180 for robot , robot_struct in action_struct .children
163181 }
164182
165- uuids = conn . execute ( f"SELECT DISTINCT uuid FROM read_parquet(' { source_escaped } ') ORDER BY uuid" ). fetchall ( )
166- for index , (episode_id ,) in enumerate (uuids ):
183+ episodes = _episode_starts ( conn , source_escaped )
184+ for index , (episode_id , _ ) in enumerate (episodes ):
167185 if n != - 1 and index >= n :
168186 break
169187
@@ -204,7 +222,6 @@ def export_episode_videos(
204222 if not rows :
205223 continue
206224
207- timestamp = datetime .datetime .fromtimestamp (float (rows [0 ][0 ])).strftime ("%Y-%m-%d-%H-%M-%S" )
208225 frames = []
209226 joint_history = {
210227 robot : np .asarray ([row [1 + len (camera_names ) + robot_idx ] for row in rows ], dtype = np .float32 )
@@ -233,4 +250,70 @@ def export_episode_videos(
233250 tiled [top : top + height , left : left + width ] = image
234251 frames .append (tiled )
235252
236- _write_mp4 (frames , output / f"{ timestamp } .mp4" , fps = fps )
253+ _write_mp4 (frames , output / _episode_filename (float (rows [0 ][0 ]), index ), fps = fps )
254+
255+
256+ def export_camera_episode_videos (
257+ dataset : str | Path ,
258+ output : str | Path ,
259+ fps : int = 30 ,
260+ camera : str | None = None ,
261+ episode : int | None = None ,
262+ ) -> None :
263+ """Export raw camera frames as one MP4 for every selected camera and episode.
264+
265+ ``episode`` is the zero-based recording-order index used in the filenames.
266+ """
267+ import torch
268+ from torchvision .io import decode_jpeg
269+
270+ dataset = Path (dataset )
271+ output = Path (output )
272+ output .mkdir (parents = True , exist_ok = True )
273+
274+ source = str (dataset / "*.parquet" ) if dataset .is_dir () else str (dataset )
275+ source_escaped = source .replace ("'" , "''" )
276+ conn = duckdb .connect ()
277+ relation = conn .sql (f"SELECT * FROM read_parquet('{ source_escaped } ')" )
278+ frame_struct = relation .select ("obs.frames" ).types [0 ]
279+ camera_names = [name for name , _ in frame_struct .children ]
280+ if camera is not None :
281+ if camera not in camera_names :
282+ available = ", " .join (camera_names )
283+ message = f"Unknown camera { camera !r} . Available cameras: { available } "
284+ raise ValueError (message )
285+ camera_names = [camera ]
286+
287+ episodes = _episode_starts (conn , source_escaped )
288+ if episode is not None :
289+ if episode < 0 or episode >= len (episodes ):
290+ message = f"Episode { episode } is out of range (dataset has { len (episodes )} episodes)."
291+ raise ValueError (message )
292+ selected_episodes = [(episode , episodes [episode ])]
293+ else :
294+ selected_episodes = list (enumerate (episodes ))
295+
296+ for episode_number , (episode_id , _ ) in selected_episodes :
297+ for camera_name in camera_names :
298+ rows = conn .execute (
299+ f"""
300+ SELECT timestamp, obs.frames.{ camera_name } .rgb.data
301+ FROM read_parquet('{ source_escaped } ')
302+ WHERE uuid = ?
303+ AND obs.frames.{ camera_name } .rgb.data IS NOT NULL
304+ ORDER BY step
305+ """ ,
306+ [episode_id ],
307+ ).fetchall ()
308+ if not rows :
309+ continue
310+
311+ frames = [
312+ decode_jpeg (torch .frombuffer (bytearray (image_bytes ), dtype = torch .uint8 )).permute (1 , 2 , 0 ).cpu ().numpy ()
313+ for _ , image_bytes in rows
314+ ]
315+ _write_mp4 (
316+ frames ,
317+ output / _episode_filename (float (rows [0 ][0 ]), episode_number , camera_name ),
318+ fps = fps ,
319+ )
0 commit comments