Team Ai
Apppublic

IPEC-COMMUNITY/openx_lerobot_visualizer

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
25likes
visualize_dataset_html.py454 linesDownload Raw Back to root
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