IPEC-COMMUNITY/openx_lerobot_visualizer
25
1#!/usr/bin/env python2 3# Copyright 2024 The HuggingFace Inc. team. All rights reserved.4#5# Licensed under the Apache License, Version 2.0 (the "License");6# you may not use this file except in compliance with the License.7# You may obtain a copy of the License at8#9# http://www.apache.org/licenses/LICENSE-2.010#11# Unless required by applicable law or agreed to in writing, software12# distributed under the License is distributed on an "AS IS" BASIS,13# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.14# See the License for the specific language governing permissions and15# limitations under the License.16""" Visualize data of **all** frames of any episode of a dataset of type LeRobotDataset.17copy from https://github.com/huggingface/lerobot/blob/main/lerobot/scripts/visualize_dataset_html.py18 19Note: The last frame of the episode doesnt always correspond to a final state.20That's because our datasets are composed of transition from state to state up to21the antepenultimate state associated to the ultimate action to arrive in the final state.22However, there might not be a transition from a final state to another state.23 24Note: This script aims to visualize the data used to train the neural networks.25~What you see is what you get~. When visualizing image modality, it is often expected to observe26lossly compression artifacts since these images have been decoded from compressed mp4 videos to27save disk space. The compression factor applied has been tuned to not affect success rate.28 29Example of usage:30 31- Visualize data stored on a local machine:32```bash33local$ python lerobot/scripts/visualize_dataset_html.py \34 --repo-id lerobot/pusht35 36local$ open http://localhost:909037```38 39- Visualize data stored on a distant machine with a local viewer:40```bash41distant$ python lerobot/scripts/visualize_dataset_html.py \42 --repo-id lerobot/pusht43 44local$ ssh -L 9090:localhost:9090 distant # create a ssh tunnel45local$ open http://localhost:909046```47 48- Select episodes to visualize:49```bash50python lerobot/scripts/visualize_dataset_html.py \51 --repo-id lerobot/pusht \52 --episodes 7 3 5 1 453```54"""55 56import argparse57import csv58import json59import logging60import re61import shutil62import tempfile63from io import StringIO64from pathlib import Path65 66import numpy as np67import pandas as pd68import requests69from flask import Flask, redirect, render_template, request, url_for70from huggingface_hub import HfApi71from lerobot.common.datasets.lerobot_dataset import LeRobotDataset72from lerobot.common.datasets.utils import IterableNamespace73from lerobot.common.utils.utils import init_logging74 75 76def available_datasets():77 api = HfApi()78 datasets = api.list_datasets(author="IPEC-COMMUNITY", tags=["LeRobot"], filter="modality:video")79 datasets = [dataset.id for dataset in datasets]80 return datasets81 82 83 84def run_server(85 dataset: LeRobotDataset | IterableNamespace | None,86 episodes: list[int] | None,87 host: str,88 port: str,89 static_folder: Path,90 template_folder: Path,91):92 app = Flask(__name__, static_folder=static_folder.resolve(), template_folder=template_folder.resolve())93 app.config["SEND_FILE_MAX_AGE_DEFAULT"] = 0 # specifying not to cache94 95 @app.route("/")96 def hommepage(dataset=dataset):97 if dataset:98 dataset_namespace, dataset_name = dataset.repo_id.split("/")99 return redirect(100 url_for(101 "show_episode",102 dataset_namespace=dataset_namespace,103 dataset_name=dataset_name,104 episode_id=0,105 )106 )107 108 dataset_param, episode_param = None, None109 all_params = request.args110 if "dataset" in all_params:111 dataset_param = all_params["dataset"]112 if "episode" in all_params:113 episode_param = int(all_params["episode"])114 115 if dataset_param:116 dataset_namespace, dataset_name = dataset_param.split("/")117 return redirect(118 url_for(119 "show_episode",120 dataset_namespace=dataset_namespace,121 dataset_name=dataset_name,122 episode_id=episode_param if episode_param is not None else 0,123 )124 )125 126 featured_datasets = [127 "IPEC-COMMUNITY/roboturk_lerobot",128 "IPEC-COMMUNITY/cmu_play_fusion_lerobot",129 "IPEC-COMMUNITY/fractal20220817_data_lerobot",130 ]131 return render_template(132 "visualize_dataset_homepage.html",133 featured_datasets=featured_datasets,134 lerobot_datasets=available_datasets(),135 )136 137 @app.route("/<string:dataset_namespace>/<string:dataset_name>")138 def show_first_episode(dataset_namespace, dataset_name):139 first_episode_id = 0140 return redirect(141 url_for(142 "show_episode",143 dataset_namespace=dataset_namespace,144 dataset_name=dataset_name,145 episode_id=first_episode_id,146 )147 )148 149 @app.route("/<string:dataset_namespace>/<string:dataset_name>/episode_<int:episode_id>")150 def show_episode(dataset_namespace, dataset_name, episode_id, dataset=dataset, episodes=episodes):151 repo_id = f"{dataset_namespace}/{dataset_name}"152 try:153 if dataset is None:154 dataset = get_dataset_info(repo_id)155 except FileNotFoundError:156 return (157 "Make sure to convert your LeRobotDataset to v2 & above. See how to convert your dataset at https://github.com/huggingface/lerobot/pull/461",158 400,159 )160 dataset_version = dataset.meta._version if isinstance(dataset, LeRobotDataset) else dataset.codebase_version161 match = re.search(r"v(\d+)\.", dataset_version)162 if match:163 major_version = int(match.group(1))164 if major_version < 2:165 return "Make sure to convert your LeRobotDataset to v2 & above."166 167 episode_data_csv_str, columns = get_episode_data(dataset, episode_id)168 dataset_info = {169 "repo_id": f"{dataset_namespace}/{dataset_name}",170 "num_samples": dataset.num_frames if isinstance(dataset, LeRobotDataset) else dataset.total_frames,171 "num_episodes": dataset.num_episodes if isinstance(dataset, LeRobotDataset) else dataset.total_episodes,172 "fps": dataset.fps,173 }174 if isinstance(dataset, LeRobotDataset):175 video_paths = [dataset.meta.get_video_file_path(episode_id, key) for key in dataset.meta.video_keys]176 videos_info = [177 {"url": url_for("static", filename=video_path), "filename": video_path.parent.name}178 for video_path in video_paths179 ]180 tasks = dataset.meta.episodes[episode_id]["tasks"]181 else:182 video_keys = [key for key, ft in dataset.features.items() if ft["dtype"] == "video"]183 videos_info = [184 {185 "url": f"https://huggingface.co/datasets/{repo_id}/resolve/main/"186 + dataset.video_path.format(187 episode_chunk=int(episode_id) // dataset.chunks_size,188 video_key=video_key,189 episode_index=episode_id,190 ),191 "filename": video_key,192 }193 for video_key in video_keys194 ]195 196 response = requests.get(f"https://huggingface.co/datasets/{repo_id}/resolve/main/meta/episodes.jsonl")197 response.raise_for_status()198 # Split into lines and parse each line as JSON199 tasks_jsonl = [json.loads(line) for line in response.text.splitlines() if line.strip()]200 201 filtered_tasks_jsonl = [row for row in tasks_jsonl if row["episode_index"] == episode_id]202 tasks = filtered_tasks_jsonl[0]["tasks"]203 204 videos_info[0]["language_instruction"] = tasks205 206 if episodes is None:207 episodes = list(208 range(dataset.num_episodes if isinstance(dataset, LeRobotDataset) else dataset.total_episodes)209 )210 211 return render_template(212 "visualize_dataset_template.html",213 episode_id=episode_id,214 episodes=episodes,215 dataset_info=dataset_info,216 videos_info=videos_info,217 episode_data_csv_str=episode_data_csv_str,218 columns=columns,219 )220 221 app.run(host=host, port=port)222 223 224def get_ep_csv_fname(episode_id: int):225 ep_csv_fname = f"episode_{episode_id}.csv"226 return ep_csv_fname227 228 229def get_episode_data(dataset: LeRobotDataset | IterableNamespace, episode_index):230 """Get a csv str containing timeseries data of an episode (e.g. state and action).231 This file will be loaded by Dygraph javascript to plot data in real time."""232 columns = []233 234 selected_columns = [col for col, ft in dataset.features.items() if ft["dtype"] == "float32"]235 selected_columns.remove("timestamp")236 237 # init header of csv with state and action names238 header = ["timestamp"]239 240 for column_name in selected_columns:241 dim_state = (242 dataset.meta.shapes[column_name][0]243 if isinstance(dataset, LeRobotDataset)244 else dataset.features[column_name].shape[0]245 )246 header += [f"{column_name}_{i}" for i in range(dim_state)]247 248 if "names" in dataset.features[column_name] and dataset.features[column_name]["names"]:249 column_names = dataset.features[column_name]["names"]250 while not isinstance(column_names, list):251 column_names = list(column_names.values())[0]252 else:253 column_names = [f"motor_{i}" for i in range(dim_state)]254 columns.append({"key": column_name, "value": column_names})255 256 selected_columns.insert(0, "timestamp")257 258 if isinstance(dataset, LeRobotDataset):259 from_idx = dataset.episode_data_index["from"][episode_index]260 to_idx = dataset.episode_data_index["to"][episode_index]261 data = dataset.hf_dataset.select(range(from_idx, to_idx)).select_columns(selected_columns).with_format("pandas")262 else:263 repo_id = dataset.repo_id264 265 url = f"https://huggingface.co/datasets/{repo_id}/resolve/main/" + dataset.data_path.format(266 episode_chunk=int(episode_index) // dataset.chunks_size, episode_index=episode_index267 )268 df = pd.read_parquet(url)269 data = df[selected_columns] # Select specific columns270 271 rows = np.hstack(272 (273 np.expand_dims(data["timestamp"], axis=1),274 *[np.vstack(data[col]) for col in selected_columns[1:]],275 )276 ).tolist()277 278 # Convert data to CSV string279 csv_buffer = StringIO()280 csv_writer = csv.writer(csv_buffer)281 # Write header282 csv_writer.writerow(header)283 # Write data rows284 csv_writer.writerows(rows)285 csv_string = csv_buffer.getvalue()286 287 return csv_string, columns288 289 290def get_episode_video_paths(dataset: LeRobotDataset, ep_index: int) -> list[str]:291 # get first frame of episode (hack to get video_path of the episode)292 first_frame_idx = dataset.episode_data_index["from"][ep_index].item()293 return [dataset.hf_dataset.select_columns(key)[first_frame_idx][key]["path"] for key in dataset.meta.video_keys]294 295 296def get_episode_language_instruction(dataset: LeRobotDataset, ep_index: int) -> list[str]:297 # check if the dataset has language instructions298 if "language_instruction" not in dataset.features:299 return None300 301 # get first frame index302 first_frame_idx = dataset.episode_data_index["from"][ep_index].item()303 304 language_instruction = dataset.hf_dataset[first_frame_idx]["language_instruction"]305 # TODO (michel-aractingi) hack to get the sentence, some strings in openx are badly stored306 # with the tf.tensor appearing in the string307 return language_instruction.removeprefix("tf.Tensor(b'").removesuffix("', shape=(), dtype=string)")308 309 310def get_dataset_info(repo_id: str) -> IterableNamespace:311 response = requests.get(f"https://huggingface.co/datasets/{repo_id}/resolve/main/meta/info.json")312 response.raise_for_status() # Raises an HTTPError for bad responses313 dataset_info = response.json()314 dataset_info["repo_id"] = repo_id315 return IterableNamespace(dataset_info)316 317 318def visualize_dataset_html(319 dataset: LeRobotDataset | None,320 episodes: list[int] | None = None,321 output_dir: Path | None = None,322 serve: bool = True,323 host: str = "127.0.0.1",324 port: int = 9090,325 force_override: bool = False,326) -> Path | None:327 init_logging()328 329 template_dir = Path(__file__).resolve().parent / "templates"330 331 if output_dir is None:332 # Create a temporary directory that will be automatically cleaned up333 output_dir = tempfile.mkdtemp(prefix="lerobot_visualize_dataset_")334 335 output_dir = Path(output_dir)336 if output_dir.exists():337 if force_override:338 shutil.rmtree(output_dir)339 else:340 logging.info(f"Output directory already exists. Loading from it: '{output_dir}'")341 342 output_dir.mkdir(parents=True, exist_ok=True)343 344 static_dir = output_dir / "static"345 static_dir.mkdir(parents=True, exist_ok=True)346 347 if dataset is None:348 if serve:349 run_server(350 dataset=None,351 episodes=None,352 host=host,353 port=port,354 static_folder=static_dir,355 template_folder=template_dir,356 )357 else:358 # Create a simlink from the dataset video folder containg mp4 files to the output directory359 # so that the http server can get access to the mp4 files.360 if isinstance(dataset, LeRobotDataset):361 ln_videos_dir = static_dir / "videos"362 if not ln_videos_dir.exists():363 ln_videos_dir.symlink_to((dataset.root / "videos").resolve())364 365 if serve:366 run_server(dataset, episodes, host, port, static_dir, template_dir)367 368 369def main():370 parser = argparse.ArgumentParser()371 372 parser.add_argument(373 "--repo-id",374 type=str,375 default=None,376 help="Name of hugging face repositery containing a LeRobotDataset dataset (e.g. `lerobot/pusht` for https://huggingface.co/datasets/lerobot/pusht).",377 )378 parser.add_argument(379 "--local-files-only",380 type=int,381 default=0,382 help="Use local files only. By default, this script will try to fetch the dataset from the hub if it exists.",383 )384 parser.add_argument(385 "--root",386 type=Path,387 default=None,388 help="Root directory for a dataset stored locally (e.g. `--root data`). By default, the dataset will be loaded from hugging face cache folder, or downloaded from the hub if available.",389 )390 parser.add_argument(391 "--load-from-hf-hub",392 type=int,393 default=0,394 help="Load videos and parquet files from HF Hub rather than local system.",395 )396 parser.add_argument(397 "--episodes",398 type=int,399 nargs="*",400 default=None,401 help="Episode indices to visualize (e.g. `0 1 5 6` to load episodes of index 0, 1, 5 and 6). By default loads all episodes.",402 )403 parser.add_argument(404 "--output-dir",405 type=Path,406 default=None,407 help="Directory path to write html files and kickoff a web server. By default write them to 'outputs/visualize_dataset/REPO_ID'.",408 )409 parser.add_argument(410 "--serve",411 type=int,412 default=1,413 help="Launch web server.",414 )415 parser.add_argument(416 "--host",417 type=str,418 default="127.0.0.1",419 help="Web host used by the http server.",420 )421 parser.add_argument(422 "--port",423 type=int,424 default=9090,425 help="Web port used by the http server.",426 )427 parser.add_argument(428 "--force-override",429 type=int,430 default=0,431 help="Delete the output directory if it exists already.",432 )433 434 args = parser.parse_args()435 kwargs = vars(args)436 repo_id = kwargs.pop("repo_id")437 load_from_hf_hub = kwargs.pop("load_from_hf_hub")438 root = kwargs.pop("root")439 local_files_only = kwargs.pop("local_files_only")440 441 dataset = None442 if repo_id:443 dataset = (444 LeRobotDataset(repo_id, root=root, local_files_only=local_files_only)445 if not load_from_hf_hub446 else get_dataset_info(repo_id)447 )448 449 visualize_dataset_html(dataset, **vars(args))450 451 452if __name__ == "__main__":453 main()454 