#!/usr/bin/env python # Copyright 2024 The HuggingFace Inc. team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. from dataclasses import dataclass, field import abc import draccus from src.dataset.transforms import ImageTransformsConfig from src.transforms.core import TransformGroup from src.dataset.video_utils import get_safe_default_codec @dataclass class DatasetConfig(draccus.ChoiceRegistry, abc.ABC): # You may provide a list of datasets here. `train.py` creates them all and concatenates them. Note: only data # keys common between the datasets are kept. Each dataset gets and additional transform that inserts the # "dataset_index" into the returned item. The index mapping is made according to the order in which the # datasets are provided. repo_id: str # Root directory where the dataset will be stored (e.g. 'dataset/path'). root: str | None = None episodes: list[int] | None = None image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig) revision: str | None = None use_imagenet_stats: bool = True use_external_stats: bool = False external_stats_path: str | None = None weight_rules_path: str | None = None video_backend: str = field(default_factory=get_safe_default_codec) streaming: bool = False dist_loading: bool = False buffer_size: int = 1024 action_mode: str = "abs" # abs | delta repack_transforms: TransformGroup = field(default_factory=TransformGroup) data_transforms: TransformGroup = field(default_factory=TransformGroup) model_transforms: TransformGroup = field(default_factory=TransformGroup) def __post_init__(self): assert self.action_mode in ['abs', 'delta'], "Either abs or delta for 'action_type'. " # The two external-stats knobs must agree, else stats silently go # missing or a real custom-stats file is silently ignored. if self.use_external_stats and not self.external_stats_path: raise ValueError( "use_external_stats=True but external_stats_path is None. " "Either set external_stats_path, or set use_external_stats=False." ) if (not self.use_external_stats) and self.external_stats_path: raise ValueError( f"external_stats_path={self.external_stats_path!r} is set but " f"use_external_stats=False. The path would be silently ignored. " f"Set use_external_stats=True to consume it, or clear the path." )