eval_stack: pi0.5 negmesh8 160-episode replay eval harness, job scripts, protocol docs
51c2c72 verified Download eval_stack/harness/config.py from Ronaldo-GOAT/transfer: direct link, hf CLI and curl.
- Browser
- Download file 65.7 kB
-
https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/eval_stack/harness/config.py
- Command line
-
hf download hf://Ronaldo-GOAT/transfer/eval_stack/harness/config.py
-
curl -L -o config.py https://huggingface.co/Ronaldo-GOAT/transfer/resolve/main/eval_stack/harness/config.py
65.7 kB
| """See _CONFIGS for the list of available configs.""" | |
| import abc | |
| import json | |
| import os | |
| from collections.abc import Sequence | |
| import dataclasses | |
| import difflib | |
| import logging | |
| import pathlib | |
| from typing import Any, Protocol, TypeAlias | |
| import etils.epath as epath | |
| import flax.nnx as nnx | |
| from typing_extensions import override | |
| import tyro | |
| import openpi.models.model as _model | |
| import openpi.models.pi0 as pi0 | |
| import openpi.models.pi0_fast as pi0_fast | |
| import openpi.models.tokenizer as _tokenizer | |
| import openpi.policies.aloha_policy as aloha_policy | |
| import openpi.policies.droid_policy as droid_policy | |
| import openpi.policies.libero_policy as libero_policy | |
| import openpi.policies.robocasa_policy as robocasa_policy | |
| import openpi.shared.download as _download | |
| import openpi.shared.normalize as _normalize | |
| import openpi.training.droid_rlds_dataset as droid_rlds_dataset | |
| import openpi.training.optimizer as _optimizer | |
| import openpi.training.weight_loaders as weight_loaders | |
| import openpi.transforms as _transforms | |
| import numpy as np | |
| import openpi.groot_utils.groot_openpi_dataset as _groot_openpi_dataset | |
| from robocasa.macros import DATASET_BASE_PATH | |
| from robocasa.utils.dataset_registry import DATASET_SOUP_REGISTRY | |
| from robocasa.utils.dataset_registry_utils import get_ds_meta | |
| ModelType: TypeAlias = _model.ModelType | |
| # Work around a tyro issue with using nnx.filterlib.Filter directly. | |
| Filter: TypeAlias = nnx.filterlib.Filter | |
| class AssetsConfig: | |
| """Determines the location of assets (e.g., norm stats) that will be used to set up the data pipeline. | |
| These assets will be replicated inside the checkpoint under the `assets/asset_id` directory. | |
| This can be used to load assets from a different checkpoint (e.g., base model checkpoint) or some other | |
| centralized location. For example, to load the norm stats for the Trossen robot from the base model checkpoint | |
| during fine-tuning, use: | |
| ``` | |
| AssetsConfig( | |
| assets_dir="gs://openpi-assets/checkpoints/pi0_base/assets", | |
| asset_id="trossen", | |
| ) | |
| ``` | |
| """ | |
| # Assets directory. If not provided, the config assets_dirs will be used. This is useful to load assets from | |
| # a different checkpoint (e.g., base model checkpoint) or some other centralized location. | |
| assets_dir: str | None = None | |
| # Asset id. If not provided, the repo id will be used. This allows users to reference assets that describe | |
| # different robot platforms. | |
| asset_id: str | None = None | |
| class DataConfig: | |
| # LeRobot repo id. If None, fake data will be created. | |
| repo_id: str | None = None | |
| # Directory within the assets directory containing the data assets. | |
| asset_id: str | None = None | |
| # Contains precomputed normalization stats. If None, normalization will not be performed. | |
| norm_stats: dict[str, _transforms.NormStats] | None = None | |
| # Used to adopt the inputs from a dataset specific format to a common format | |
| # which is expected by the data transforms. | |
| repack_transforms: _transforms.Group = dataclasses.field(default_factory=_transforms.Group) | |
| # Data transforms, typically include robot specific transformations. Will be applied | |
| # before the data is normalized. See `model.Observation` and `model.Actions` to learn about the | |
| # normalized data. | |
| data_transforms: _transforms.Group = dataclasses.field(default_factory=_transforms.Group) | |
| # Model specific transforms. Will be applied after the data is normalized. | |
| model_transforms: _transforms.Group = dataclasses.field(default_factory=_transforms.Group) | |
| # If true, will use quantile normalization. Otherwise, normal z-score normalization will be used. | |
| use_quantile_norm: bool = False | |
| # Names of keys that will be used by the data loader to generate the action sequence. The length of the | |
| # sequence is defined by the `action_horizon` field in the model config. This should be adjusted if your | |
| # LeRobot dataset is using different keys to represent the action. | |
| action_sequence_keys: Sequence[str] = ("actions",) | |
| # If true, will use the LeRobot dataset task to define the prompt. | |
| prompt_from_task: bool = False | |
| # Only used for RLDS data loader (ie currently only used for DROID). | |
| rlds_data_dir: str | None = None | |
| # Action space for DROID dataset. | |
| action_space: droid_rlds_dataset.DroidActionSpace | None = None | |
| # Action dimension for padding (used by Groot datasets) | |
| action_dim: int | None = None | |
| # Multi-dataset support for Groot datasets | |
| data_dirs: list[str] | None = None # List of data directories for multi-dataset | |
| dataset_weights: list[float] | None = None # Weights for each dataset in multi-dataset | |
| class GroupFactory(Protocol): | |
| def __call__(self, model_config: _model.BaseModelConfig) -> _transforms.Group: | |
| """Create a group.""" | |
| class ModelTransformFactory(GroupFactory): | |
| """Creates model transforms for standard pi0 models.""" | |
| # If provided, will determine the default prompt that be used by the model. | |
| default_prompt: str | None = None | |
| def __call__(self, model_config: _model.BaseModelConfig) -> _transforms.Group: | |
| match model_config.model_type: | |
| case _model.ModelType.PI0: | |
| return _transforms.Group( | |
| inputs=[ | |
| _transforms.InjectDefaultPrompt(self.default_prompt), | |
| _transforms.ResizeImages(224, 224), | |
| _transforms.TokenizePrompt( | |
| _tokenizer.PaligemmaTokenizer(model_config.max_token_len), | |
| ), | |
| ], | |
| ) | |
| case _model.ModelType.PI05: | |
| assert isinstance(model_config, pi0.Pi0Config) | |
| return _transforms.Group( | |
| inputs=[ | |
| _transforms.InjectDefaultPrompt(self.default_prompt), | |
| _transforms.ResizeImages(224, 224), | |
| _transforms.TokenizePrompt( | |
| _tokenizer.PaligemmaTokenizer(model_config.max_token_len), | |
| discrete_state_input=model_config.discrete_state_input, | |
| ), | |
| _transforms.PadStatesAndActions(model_config.action_dim), | |
| ], | |
| ) | |
| case _model.ModelType.PI0_FAST: | |
| return _transforms.Group( | |
| inputs=[ | |
| _transforms.InjectDefaultPrompt(self.default_prompt), | |
| _transforms.ResizeImages(224, 224), | |
| _transforms.TokenizeFASTInputs( | |
| _tokenizer.FASTTokenizer(model_config.max_token_len), | |
| ), | |
| ], | |
| outputs=[ | |
| _transforms.ExtractFASTActions( | |
| _tokenizer.FASTTokenizer(model_config.max_token_len), | |
| action_horizon=model_config.action_horizon, | |
| action_dim=model_config.action_dim, | |
| ) | |
| ], | |
| ) | |
| class DataConfigFactory(abc.ABC): | |
| # The LeRobot repo id. | |
| repo_id: str | None = None | |
| # Determines how the assets will be loaded. | |
| assets: AssetsConfig = dataclasses.field(default_factory=AssetsConfig) | |
| # Base config that will be updated by the factory. | |
| base_config: tyro.conf.Suppress[DataConfig | None] = None | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| """Create a data config.""" | |
| def create_base_config(self, assets_dirs: pathlib.Path) -> DataConfig: | |
| repo_id = self.repo_id if self.repo_id is not tyro.MISSING else None | |
| asset_id = self.assets.asset_id or repo_id | |
| base = self.base_config or DataConfig() | |
| # Preserve pre-supplied norm_stats; only load if not provided | |
| existing_stats = base.norm_stats | |
| loaded_stats = None if existing_stats is not None else self._load_norm_stats( | |
| epath.Path(self.assets.assets_dir or assets_dirs), asset_id | |
| ) | |
| return dataclasses.replace( | |
| base, | |
| repo_id=repo_id, | |
| asset_id=asset_id, | |
| norm_stats=existing_stats if existing_stats is not None else loaded_stats, | |
| ) | |
| def _load_norm_stats(self, assets_dir: epath.Path, asset_id: str | None) -> dict[str, _transforms.NormStats] | None: | |
| if asset_id is None: | |
| return None | |
| try: | |
| data_assets_dir = str(assets_dir / asset_id) | |
| norm_stats = _normalize.load(_download.maybe_download(data_assets_dir)) | |
| logging.info(f"Loaded norm stats from {data_assets_dir}") | |
| return norm_stats | |
| except FileNotFoundError: | |
| logging.info(f"Norm stats not found in {data_assets_dir}.") | |
| # Fallback: try to read and convert stats from repo meta | |
| # TODO: fix | |
| converted = _groot_openpi_dataset._convert_stats_from_repo_meta(asset_id) | |
| if converted is not None: | |
| logging.info(f"Converted norm stats from repo meta for {asset_id}") | |
| return converted | |
| return None | |
| class FakeDataConfig(DataConfigFactory): | |
| repo_id: str = "fake" | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| return DataConfig(repo_id=self.repo_id) | |
| class SimpleDataConfig(DataConfigFactory): | |
| # Factory for the data transforms. | |
| data_transforms: tyro.conf.Suppress[GroupFactory] = dataclasses.field(default_factory=GroupFactory) | |
| # Factory for the model transforms. | |
| model_transforms: tyro.conf.Suppress[GroupFactory] = dataclasses.field(default_factory=ModelTransformFactory) | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| return dataclasses.replace( | |
| self.create_base_config(assets_dirs), | |
| data_transforms=self.data_transforms(model_config), | |
| model_transforms=self.model_transforms(model_config), | |
| use_quantile_norm=model_config.model_type in (ModelType.PI0_FAST, ModelType.PI05), | |
| ) | |
| class LeRobotAlohaDataConfig(DataConfigFactory): | |
| # If true, will convert joint dimensions to deltas with respect to the current state before passing to the model. | |
| # Gripper dimensions will remain in absolute values. | |
| use_delta_joint_actions: bool = True | |
| # If provided, will be injected into the input data if the "prompt" key is not present. | |
| default_prompt: str | None = None | |
| # If true, this will convert the joint and gripper values from the standard Aloha space to | |
| # the space used by the pi internal runtime which was used to train the base model. People who | |
| # use standard Aloha data should set this to true. | |
| adapt_to_pi: bool = True | |
| # Repack transforms. | |
| repack_transforms: tyro.conf.Suppress[_transforms.Group] = dataclasses.field( | |
| default=_transforms.Group( | |
| inputs=[ | |
| _transforms.RepackTransform( | |
| { | |
| "images": {"cam_high": "observation.images.top"}, | |
| "state": "observation.state", | |
| "actions": "action", | |
| } | |
| ) | |
| ] | |
| ) | |
| ) | |
| # Action keys that will be used to read the action sequence from the dataset. | |
| action_sequence_keys: Sequence[str] = ("action",) | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| data_transforms = _transforms.Group( | |
| inputs=[aloha_policy.AlohaInputs(action_dim=model_config.action_dim, adapt_to_pi=self.adapt_to_pi)], | |
| outputs=[aloha_policy.AlohaOutputs(adapt_to_pi=self.adapt_to_pi)], | |
| ) | |
| if self.use_delta_joint_actions: | |
| delta_action_mask = _transforms.make_bool_mask(6, -1, 6, -1) | |
| data_transforms = data_transforms.push( | |
| inputs=[_transforms.DeltaActions(delta_action_mask)], | |
| outputs=[_transforms.AbsoluteActions(delta_action_mask)], | |
| ) | |
| model_transforms = ModelTransformFactory(default_prompt=self.default_prompt)(model_config) | |
| return dataclasses.replace( | |
| self.create_base_config(assets_dirs), | |
| repack_transforms=self.repack_transforms, | |
| data_transforms=data_transforms, | |
| model_transforms=model_transforms, | |
| action_sequence_keys=self.action_sequence_keys, | |
| ) | |
| class LeRobotLiberoDataConfig(DataConfigFactory): | |
| """ | |
| This config is used to configure transforms that are applied at various parts of the data pipeline. | |
| For your own dataset, you can copy this class and modify the transforms to match your dataset based on the | |
| comments below. | |
| """ | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| # The repack transform is *only* applied to the data coming from the dataset, | |
| # and *not* during inference. We can use it to make inputs from the dataset look | |
| # as close as possible to those coming from the inference environment (e.g. match the keys). | |
| # Below, we match the keys in the dataset (which we defined in the data conversion script) to | |
| # the keys we use in our inference pipeline (defined in the inference script for libero). | |
| # For your own dataset, first figure out what keys your environment passes to the policy server | |
| # and then modify the mappings below so your dataset's keys get matched to those target keys. | |
| # The repack transform simply remaps key names here. | |
| repack_transform = _transforms.Group( | |
| inputs=[ | |
| _transforms.RepackTransform( | |
| { | |
| "observation/image": "image", | |
| "observation/wrist_image": "wrist_image", | |
| "observation/state": "state", | |
| "actions": "actions", | |
| "prompt": "prompt", | |
| } | |
| ) | |
| ] | |
| ) | |
| # The data transforms are applied to the data coming from the dataset *and* during inference. | |
| # Below, we define the transforms for data going into the model (``inputs``) and the transforms | |
| # for data coming out of the model (``outputs``) (the latter is only used during inference). | |
| # We defined these transforms in `libero_policy.py`. You can check the detailed comments there for | |
| # how to modify the transforms to match your dataset. Once you created your own transforms, you can | |
| # replace the transforms below with your own. | |
| data_transforms = _transforms.Group( | |
| inputs=[libero_policy.LiberoInputs(action_dim=model_config.action_dim, model_type=model_config.model_type)], | |
| outputs=[libero_policy.LiberoOutputs()], | |
| ) | |
| # One additional data transform: pi0 models are trained on delta actions (relative to the first | |
| # state in each action chunk). IF your data has ``absolute`` actions (e.g. target joint angles) | |
| # you can uncomment the following line to convert the actions to delta actions. The only exception | |
| # is for the gripper actions which are always absolute. | |
| # In the example below, we would apply the delta conversion to the first 6 actions (joints) and | |
| # leave the 7th action (gripper) unchanged, i.e. absolute. | |
| # In Libero, the raw actions in the dataset are already delta actions, so we *do not* need to | |
| # apply a separate delta conversion (that's why it's commented out). Choose whether to apply this | |
| # transform based on whether your dataset uses ``absolute`` or ``delta`` actions out of the box. | |
| # TODO(karl): comment this out once we have updated the Libero checkpoints to not use | |
| # the delta action transform | |
| delta_action_mask = _transforms.make_bool_mask(6, -1) | |
| data_transforms = data_transforms.push( | |
| inputs=[_transforms.DeltaActions(delta_action_mask)], | |
| outputs=[_transforms.AbsoluteActions(delta_action_mask)], | |
| ) | |
| # Model transforms include things like tokenizing the prompt and action targets | |
| # You do not need to change anything here for your own dataset. | |
| model_transforms = ModelTransformFactory()(model_config) | |
| # We return all data transforms for training and inference. No need to change anything here. | |
| return dataclasses.replace( | |
| self.create_base_config(assets_dirs), | |
| repack_transforms=repack_transform, | |
| data_transforms=data_transforms, | |
| model_transforms=model_transforms, | |
| ) | |
| class RLDSDroidDataConfig(DataConfigFactory): | |
| """ | |
| Config for training on DROID, using RLDS data format (for efficient training on larger datasets). | |
| """ | |
| rlds_data_dir: str | None = None | |
| action_space: droid_rlds_dataset.DroidActionSpace | None = None | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| repack_transform = _transforms.Group( | |
| inputs=[ | |
| _transforms.RepackTransform( | |
| { | |
| "observation/exterior_image_1_left": "observation/image", | |
| "observation/wrist_image_left": "observation/wrist_image", | |
| "observation/joint_position": "observation/joint_position", | |
| "observation/gripper_position": "observation/gripper_position", | |
| "actions": "actions", | |
| "prompt": "prompt", | |
| } | |
| ) | |
| ] | |
| ) | |
| data_transforms = _transforms.Group( | |
| inputs=[droid_policy.DroidInputs(action_dim=model_config.action_dim, model_type=model_config.model_type)], | |
| outputs=[droid_policy.DroidOutputs()], | |
| ) | |
| if self.action_space == droid_rlds_dataset.DroidActionSpace.JOINT_POSITION: | |
| # Data loader returns absolute joint position actions -- convert to delta actions for training. | |
| delta_action_mask = _transforms.make_bool_mask(7, -1) | |
| data_transforms = data_transforms.push( | |
| inputs=[_transforms.DeltaActions(delta_action_mask)], | |
| outputs=[_transforms.AbsoluteActions(delta_action_mask)], | |
| ) | |
| model_transforms = ModelTransformFactory()(model_config) | |
| assert self.rlds_data_dir is not None, "Need to set rlds data dir for RLDS data loader." | |
| return dataclasses.replace( | |
| self.create_base_config(assets_dirs), | |
| repack_transforms=repack_transform, | |
| data_transforms=data_transforms, | |
| model_transforms=model_transforms, | |
| use_quantile_norm=model_config.model_type in (ModelType.PI0_FAST, ModelType.PI05), | |
| rlds_data_dir=self.rlds_data_dir, | |
| action_space=self.action_space, | |
| ) | |
| class LeRobotRobocasaDataConfig(DataConfigFactory): | |
| """Config for training on Groot datasets. | |
| Set `assets.asset_id` (inherited from DataConfigFactory) to a directory under | |
| `assets_dirs/` that holds the precomputed `norm_stats.json` (mean/std/q01/q99) for | |
| this RoboCasa mixture — produced by `scripts/compute_norm_stats_robocasa.py`. | |
| """ | |
| repo_id: str | None = None | |
| data_dirs: Any | None = None | |
| dataset_weights: list[float] | None = None | |
| action_dim: int | None = None | |
| # REPRO (2026-09-22): the ctc502 60k BASE checkpoint was trained upstream with plain | |
| # z-score (mean/std) normalisation -- upstream's config.py never assigns | |
| # use_quantile_norm anywhere, so it kept its `False` default, and the checkpoint's own | |
| # assets/norm_stats.json has q01/q99 = null accordingly. Serving those weights under | |
| # quantile norm applies a different transform than they were trained with. Set this on | |
| # an eval-only config for that lineage. Leave False for everything trained in THIS repo | |
| # (negmesh / mimicgen arms), which really is quantile-normalised under ctc502_qnorm. | |
| force_zscore_norm: bool = False | |
| def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: | |
| repack_transform = _transforms.Group() | |
| data_transforms = _transforms.Group( | |
| inputs=[robocasa_policy.RobocasaInputs(action_dim=model_config.action_dim, model_type=model_config.model_type)], | |
| outputs=[robocasa_policy.RobocasaOutputs()], | |
| ) | |
| model_transforms = ModelTransformFactory()(model_config) | |
| base = self.create_base_config(assets_dirs) | |
| # Fallback: if norm_stats are not available in assets, derive a quick approximation | |
| # from per-task stats.json files. We try this best-effort because eval-only nodes may | |
| # not have the dataset on disk; at inference time `create_trained_policy` will then | |
| # read norm_stats directly from the checkpoint's assets/. | |
| # NOTE: this fallback only fills mean/std (no q01/q99). For pi0-FAST training, prefer | |
| # running `scripts/compute_norm_stats_robocasa.py` so quantiles are accurate. | |
| fallback_norm_stats = None | |
| if base.norm_stats is None and self.data_dirs and len(self.data_dirs) > 0: | |
| try: | |
| if len(self.data_dirs) == 1: | |
| d = self.data_dirs[0] | |
| norm_stats = _groot_openpi_dataset._load_norm_stats_from_groot_dataset(d) | |
| if norm_stats is not None: | |
| fallback_norm_stats = norm_stats | |
| logging.info(f"Loaded norm stats from local data dir: {d}") | |
| else: | |
| norm_stats = _groot_openpi_dataset._load_norm_stats_from_groot_mixture_dataset(self.data_dirs) | |
| if norm_stats is not None: | |
| fallback_norm_stats = norm_stats | |
| logging.info(f"Loaded combined norm stats from {len(self.data_dirs)} data dirs") | |
| except FileNotFoundError as e: | |
| logging.warning( | |
| "Could not load norm stats from data_dirs (%s); will fall back to checkpoint assets at inference time.", | |
| e, | |
| ) | |
| return dataclasses.replace( | |
| base, | |
| repack_transforms=repack_transform, | |
| data_transforms=data_transforms, | |
| model_transforms=model_transforms, | |
| action_dim=model_config.action_dim, | |
| data_dirs=self.data_dirs, | |
| dataset_weights=self.dataset_weights, | |
| norm_stats=base.norm_stats or fallback_norm_stats, | |
| # pi0-FAST uses FAST tokenization which assumes inputs in [-1, 1]; switch to | |
| # quantile normalization (q01/q99 -> [-1, 1]) for that model. For pi0 (non-FAST) | |
| # we keep the default z-score normalization. | |
| use_quantile_norm=(not self.force_zscore_norm) | |
| and model_config.model_type in (_model.ModelType.PI0_FAST, _model.ModelType.PI05), | |
| ) | |
| class TrainConfig: | |
| # Name of the config. Must be unique. Will be used to reference this config. | |
| name: tyro.conf.Suppress[str] | |
| # Project name. | |
| project_name: str = "openpi" | |
| # Experiment name. Will be used to name the metadata and checkpoint directories. | |
| exp_name: str = tyro.MISSING | |
| # Defines the model config. Some attributes (action_dim, action_horizon, and max_token_len) are shared by all models | |
| # -- see BaseModelConfig. Specific model implementations (e.g., Pi0Config) inherit from BaseModelConfig and may | |
| # define additional attributes. | |
| model: _model.BaseModelConfig = dataclasses.field(default_factory=pi0.Pi0Config) | |
| # A weight loader can optionally load (possibly partial) weights from disk after the model is initialized. | |
| weight_loader: weight_loaders.WeightLoader = dataclasses.field(default_factory=weight_loaders.NoOpWeightLoader) | |
| lr_schedule: _optimizer.LRScheduleConfig = dataclasses.field(default_factory=_optimizer.CosineDecaySchedule) | |
| optimizer: _optimizer.OptimizerConfig = dataclasses.field(default_factory=_optimizer.AdamW) | |
| ema_decay: float | None = 0.99 | |
| # Specifies which weights should be frozen. | |
| freeze_filter: tyro.conf.Suppress[Filter] = dataclasses.field(default_factory=nnx.Nothing) | |
| # Determines the data to be trained on. | |
| data: DataConfigFactory = dataclasses.field(default_factory=FakeDataConfig) | |
| # Base directory for config assets (e.g., norm stats). | |
| assets_base_dir: str = "./assets" | |
| # Base directory for checkpoints. | |
| checkpoint_base_dir: str = "./checkpoints" | |
| # Random seed that will be used by random generators during training. | |
| seed: int = 42 | |
| # Global batch size. | |
| batch_size: int = 32 | |
| # Number of workers to use for the data loader. Increasing this number will speed up data loading but | |
| # will increase memory and CPU usage. | |
| num_workers: int = 2 | |
| # Number of train steps (batches) to run. | |
| num_train_steps: int = 30_000 | |
| # How often (in steps) to log training metrics. | |
| log_interval: int = 100 | |
| # How often (in steps) to save checkpoints. | |
| save_interval: int = 1000 | |
| # If set, checkpoints are saved at EXACTLY these steps and save_interval is ignored. | |
| # save_interval can only express a regular cadence; an irregular schedule needs a list. | |
| # NOTE: checkpoints.py hardcodes max_to_keep=1, so keep_period is what actually PRESERVES | |
| # a saved step. Any step listed here must be divisible by keep_period or it will be | |
| # pruned as soon as the next checkpoint lands. | |
| save_steps: tuple[int, ...] | None = None | |
| # If set, any existing checkpoints matching step % keep_period == 0 will not be deleted. | |
| keep_period: int | None = 5000 | |
| # If true, will overwrite the checkpoint directory if it already exists. | |
| overwrite: bool = False | |
| # If true, will resume training from the last checkpoint. | |
| resume: bool = False | |
| # If true, will enable wandb logging. | |
| wandb_enabled: bool = True | |
| # Used to pass metadata to the policy server. | |
| policy_metadata: dict[str, Any] | None = None | |
| # If the value is greater than 1, FSDP will be enabled and shard across number of specified devices; overall | |
| # device memory will be reduced but training could potentially be slower. | |
| # eg. if total device is 4 and fsdp devices is 2; then the model will shard to 2 devices and run | |
| # data parallel between 2 groups of devices. | |
| fsdp_devices: int = 1 | |
| def assets_dirs(self) -> pathlib.Path: | |
| """Get the assets directory for this config.""" | |
| return (pathlib.Path(self.assets_base_dir) / self.name).resolve() | |
| def checkpoint_dir(self) -> pathlib.Path: | |
| """Get the checkpoint directory for this config.""" | |
| if not self.exp_name: | |
| raise ValueError("--exp_name must be set") | |
| return (pathlib.Path(self.checkpoint_base_dir) / self.name / self.exp_name).resolve() | |
| def trainable_filter(self) -> nnx.filterlib.Filter: | |
| """Get the filter for the trainable parameters.""" | |
| return nnx.All(nnx.Param, nnx.Not(self.freeze_filter)) | |
| def __post_init__(self) -> None: | |
| if self.resume and self.overwrite: | |
| raise ValueError("Cannot resume and overwrite at the same time.") | |
| # Use `get_config` if you need to get a config by name in your code. | |
| _CONFIGS = [ | |
| # | |
| # Inference Aloha configs. | |
| # | |
| TrainConfig( | |
| name="pi0_aloha", | |
| model=pi0.Pi0Config(), | |
| data=LeRobotAlohaDataConfig( | |
| assets=AssetsConfig(asset_id="trossen"), | |
| ), | |
| policy_metadata={"reset_pose": [0, -1.5, 1.5, 0, 0, 0]}, | |
| ), | |
| TrainConfig( | |
| name="pi0_aloha_towel", | |
| model=pi0.Pi0Config(), | |
| data=LeRobotAlohaDataConfig( | |
| assets=AssetsConfig(asset_id="trossen"), | |
| default_prompt="fold the towel", | |
| ), | |
| policy_metadata={"reset_pose": [0, -1.5, 1.5, 0, 0, 0]}, | |
| ), | |
| TrainConfig( | |
| name="pi0_aloha_tupperware", | |
| model=pi0.Pi0Config(), | |
| data=LeRobotAlohaDataConfig( | |
| assets=AssetsConfig(asset_id="trossen"), | |
| default_prompt="open the tupperware and put the food on the plate", | |
| ), | |
| policy_metadata={"reset_pose": [0, -1.5, 1.5, 0, 0, 0]}, | |
| ), | |
| # | |
| # Inference DROID configs. | |
| # | |
| TrainConfig( | |
| name="pi0_droid", | |
| model=pi0.Pi0Config(action_horizon=10), | |
| data=SimpleDataConfig( | |
| assets=AssetsConfig(asset_id="droid"), | |
| data_transforms=lambda model: _transforms.Group( | |
| inputs=[droid_policy.DroidInputs(action_dim=model.action_dim)], | |
| outputs=[droid_policy.DroidOutputs()], | |
| ), | |
| base_config=DataConfig( | |
| prompt_from_task=True, | |
| ), | |
| ), | |
| ), | |
| TrainConfig( | |
| name="pi0_fast_droid", | |
| model=pi0_fast.Pi0FASTConfig(action_dim=8, action_horizon=10), | |
| data=SimpleDataConfig( | |
| assets=AssetsConfig(asset_id="droid"), | |
| data_transforms=lambda model: _transforms.Group( | |
| inputs=[droid_policy.DroidInputs(action_dim=model.action_dim, model_type=ModelType.PI0_FAST)], | |
| outputs=[droid_policy.DroidOutputs()], | |
| ), | |
| base_config=DataConfig( | |
| prompt_from_task=True, | |
| ), | |
| ), | |
| ), | |
| # | |
| # Fine-tuning Libero configs. | |
| # | |
| # These train configs define the hyperparameters for fine-tuning the base model on your own dataset. | |
| # They are used to define key elements like the dataset you are training on, the base checkpoint you | |
| # are using, and other hyperparameters like how many training steps to run or what learning rate to use. | |
| # For your own dataset, you can copy this class and modify the dataset name, and data transforms based on | |
| # the comments below. | |
| TrainConfig( | |
| # Change the name to reflect your model and dataset. | |
| name="pi0_libero", | |
| # Here you define the model config -- In this example we use pi0 as the model | |
| # architecture and perform *full* finetuning. in the examples below we show how to modify | |
| # this to perform *low-memory* (LORA) finetuning and use pi0-FAST as an alternative architecture. | |
| model=pi0.Pi0Config(), | |
| # Here you define the dataset you are training on. In this example we use the Libero | |
| # dataset. For your own dataset, you can change the repo_id to point to your dataset. | |
| # Also modify the DataConfig to use the new config you made for your dataset above. | |
| data=LeRobotLiberoDataConfig( | |
| repo_id="physical-intelligence/libero", | |
| base_config=DataConfig( | |
| # This flag determines whether we load the prompt (i.e. the task instruction) from the | |
| # ``task`` field in the LeRobot dataset. If set to True, the prompt will show up in | |
| # a field called ``prompt`` in the input dict. The recommended setting is True. | |
| prompt_from_task=True, | |
| ), | |
| ), | |
| # Here you define which pre-trained checkpoint you want to load to initialize the model. | |
| # This should match the model config you chose above -- i.e. in this case we use the pi0 base model. | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| # Below you can define other hyperparameters like the learning rate, number of training steps, etc. | |
| # Check the base TrainConfig class for a full list of available hyperparameters. | |
| num_train_steps=30_000, | |
| ), | |
| TrainConfig( | |
| name="pi0_libero_low_mem_finetune", | |
| # Here is an example of loading a pi0 model for LoRA fine-tuning. | |
| model=pi0.Pi0Config(paligemma_variant="gemma_2b_lora", action_expert_variant="gemma_300m_lora"), | |
| data=LeRobotLiberoDataConfig( | |
| repo_id="physical-intelligence/libero", | |
| base_config=DataConfig(prompt_from_task=True), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=30_000, | |
| # The freeze filter defines which parameters should be frozen during training. | |
| # We have a convenience function in the model config that returns the default freeze filter | |
| # for the given model config for LoRA finetuning. Just make sure it matches the model config | |
| # you chose above. | |
| freeze_filter=pi0.Pi0Config( | |
| paligemma_variant="gemma_2b_lora", action_expert_variant="gemma_300m_lora" | |
| ).get_freeze_filter(), | |
| # Turn off EMA for LoRA finetuning. | |
| ema_decay=None, | |
| ), | |
| TrainConfig( | |
| name="pi0_fast_libero", | |
| # Here is an example of loading a pi0-FAST model for full finetuning. | |
| # Modify action_dim and action_horizon to match your dataset (action horizon is equal to | |
| # the desired action chunk length). | |
| # The max_token_len is the maximum number of (non-image) tokens the model can handle. | |
| # This includes the tokenized prompt, proprioceptive state, and (FAST-tokenized) action tokens. | |
| # Choosing this value too small may chop off tokens at the end of your sequence (the code will throw | |
| # a warning), while choosing it too large will waste memory (since we pad each batch element to the | |
| # max_token_len). A good rule of thumb is to use approx 180 for single-arm robots, and approx 250 for | |
| # two-arm robots. Generally, err on the lower side here first, and potentially increase the value if | |
| # you see many warnings being thrown during training. | |
| model=pi0_fast.Pi0FASTConfig(action_dim=7, action_horizon=10, max_token_len=180), | |
| data=LeRobotLiberoDataConfig( | |
| repo_id="physical-intelligence/libero", | |
| base_config=DataConfig(prompt_from_task=True), | |
| ), | |
| # Note that we load the pi0-FAST base model checkpoint here. | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_fast_base/params"), | |
| num_train_steps=30_000, | |
| ), | |
| TrainConfig( | |
| name="pi0_fast_libero_low_mem_finetune", | |
| # Here is an example of loading a pi0-FAST model for LoRA finetuning. | |
| # For setting action_dim, action_horizon, and max_token_len, see the comments above. | |
| model=pi0_fast.Pi0FASTConfig( | |
| action_dim=7, action_horizon=10, max_token_len=180, paligemma_variant="gemma_2b_lora" | |
| ), | |
| data=LeRobotLiberoDataConfig( | |
| repo_id="physical-intelligence/libero", | |
| base_config=DataConfig(prompt_from_task=True), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_fast_base/params"), | |
| num_train_steps=30_000, | |
| # Again, make sure to match the model config above when extracting the freeze filter | |
| # that specifies which parameters should be frozen during LoRA finetuning. | |
| freeze_filter=pi0_fast.Pi0FASTConfig( | |
| action_dim=7, action_horizon=10, max_token_len=180, paligemma_variant="gemma_2b_lora" | |
| ).get_freeze_filter(), | |
| # Turn off EMA for LoRA finetuning. | |
| ema_decay=None, | |
| ), | |
| # | |
| # Fine-tuning Aloha configs. | |
| # | |
| # This is a test config that is used to illustate how train on a custom LeRobot dataset. | |
| # For instuctions on how to convert and train on your own Aloha dataset see examples/aloha_real/README.md | |
| TrainConfig( | |
| name="pi0_aloha_pen_uncap", | |
| model=pi0.Pi0Config(), | |
| data=LeRobotAlohaDataConfig( | |
| repo_id="physical-intelligence/aloha_pen_uncap_diverse", | |
| assets=AssetsConfig( | |
| assets_dir="gs://openpi-assets/checkpoints/pi0_base/assets", | |
| asset_id="trossen", | |
| ), | |
| default_prompt="uncap the pen", | |
| repack_transforms=_transforms.Group( | |
| inputs=[ | |
| _transforms.RepackTransform( | |
| { | |
| "images": { | |
| "cam_high": "observation.images.cam_high", | |
| "cam_left_wrist": "observation.images.cam_left_wrist", | |
| "cam_right_wrist": "observation.images.cam_right_wrist", | |
| }, | |
| "state": "observation.state", | |
| "actions": "action", | |
| } | |
| ) | |
| ] | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=20_000, | |
| ), | |
| # | |
| # Fine-tuning DROID configs. | |
| # | |
| TrainConfig( | |
| name="pi0_fast_droid_finetune", | |
| model=pi0_fast.Pi0FASTConfig( | |
| action_dim=8, | |
| action_horizon=16, | |
| max_token_len=180, | |
| ), | |
| data=RLDSDroidDataConfig( | |
| repo_id="droid", | |
| # Set this to the path to your DROID RLDS dataset (the parent directory of the `droid` directory). | |
| rlds_data_dir="<path_to_droid_rlds_dataset>", | |
| action_space=droid_rlds_dataset.DroidActionSpace.JOINT_POSITION, | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_fast_base/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=5e-5, | |
| decay_steps=1_000_000, | |
| decay_lr=5e-5, | |
| ), | |
| num_train_steps=100_000, # 100k steps should be sufficient, takes ~2 days on 8x H100s | |
| batch_size=256, | |
| log_interval=100, | |
| save_interval=5000, | |
| keep_period=20_000, | |
| num_workers=0, # Important: RLDS DataLoader requires num_workers=0, handles multi-processing internally | |
| ), | |
| # | |
| # ALOHA Sim configs. This config is used to demonstrate how to train on a simple simulated environment. | |
| # | |
| TrainConfig( | |
| name="pi0_aloha_sim", | |
| model=pi0.Pi0Config(), | |
| data=LeRobotAlohaDataConfig( | |
| repo_id="lerobot/aloha_sim_transfer_cube_human", | |
| default_prompt="Transfer cube", | |
| use_delta_joint_actions=False, | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=20_000, | |
| ), | |
| # | |
| # Debugging configs. | |
| # | |
| TrainConfig( | |
| name="debug", | |
| data=FakeDataConfig(), | |
| batch_size=2, | |
| model=pi0.Pi0Config(paligemma_variant="dummy", action_expert_variant="dummy"), | |
| save_interval=100, | |
| overwrite=True, | |
| exp_name="debug", | |
| num_train_steps=10, | |
| wandb_enabled=False, | |
| ), | |
| TrainConfig( | |
| name="debug_restore", | |
| data=FakeDataConfig(), | |
| batch_size=2, | |
| model=pi0.Pi0Config(paligemma_variant="dummy", action_expert_variant="dummy"), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("./checkpoints/debug/debug/9/params"), | |
| overwrite=True, | |
| exp_name="debug", | |
| num_train_steps=10, | |
| wandb_enabled=False, | |
| ), | |
| # | |
| # RoboCasa dataset configs. | |
| # | |
| TrainConfig( | |
| name="pi0_robocasa_target50", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target50"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=500000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_finetune_target_atomic_seen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_atomic_seen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("INSERT_CKPTPOINT_HERE"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=5000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_finetune_target_composite_seen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_composite_seen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("INSERT_CKPTPOINT_HERE"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=5000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_finetune_target_composite_unseen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_composite_unseen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("INSERT_CKPTPOINT_HERE"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=5000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_target_atomic_seen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_atomic_seen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_target_atomic_seen_random_weight_init", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_atomic_seen"], | |
| ), | |
| weight_loader=weight_loaders.NoOpWeightLoader(), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_target_atomic_seen_paligemma_init", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_atomic_seen"], | |
| ), | |
| weight_loader=weight_loaders.PaliGemmaWeightLoader(), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_target_composite_seen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_composite_seen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_target_composite_unseen", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["target_composite_unseen"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_pretrain_human300_mg60", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["pretrain_human300_mg60"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=100000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| TrainConfig( | |
| name="pi0_robocasa_pretrain_human300", | |
| model=pi0.Pi0Config( | |
| max_token_len=96, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["pretrain_human300"], | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_base/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=100000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=100000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| # | |
| # pi0-FAST counterpart of pi0_robocasa_pretrain_human300 (RoboCasa guideline: batch_size=64, num_train_steps=75_000). | |
| # RoboCasa single-arm: action_dim kept at 32 (pi0 default) so that state/action padding matches the | |
| # shipped (32,) norm stats; RobocasaOutputs trims back to the real 12-dim action at inference. | |
| # action_horizon=10, max_token_len=256 (smoke10 run showed prompts 180-210 tokens; 256 gives headroom). | |
| # `assets.asset_id` points at the precomputed mixture norm_stats (mean/std/q01/q99) generated | |
| # by `scripts/compute_norm_stats_robocasa.py`; pi0-FAST also uses quantile normalization. | |
| # | |
| TrainConfig( | |
| name="pi0_fast_robocasa_pretrain_human300", | |
| model=pi0_fast.Pi0FASTConfig( | |
| action_horizon=10, | |
| max_token_len=256, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=DATASET_SOUP_REGISTRY["pretrain_human300"], | |
| assets=AssetsConfig(asset_id="robocasa365_human300"), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi0_fast_base/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=75_000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=75_000, | |
| save_interval=5000, | |
| keep_period=10000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| # | |
| # Single-task fine-tune: PickPlaceCounterToCabinet, target split (500 human demos). | |
| # Starts from the pi0-FAST human300 checkpoint and reuses its quantile norm stats | |
| # (pi0-FAST needs q01/q99; the single-dataset fallback only yields mean/std). | |
| # | |
| TrainConfig( | |
| name="pi0_fast_robocasa_target_PickPlaceCounterToCabinet", | |
| model=pi0_fast.Pi0FASTConfig( | |
| action_horizon=10, | |
| max_token_len=256, | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[get_ds_meta("PickPlaceCounterToCabinet", "target", "human")], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/pi0_fast_robocasa_pretrain_human300", | |
| asset_id="robocasa365_human300", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("./checkpoints_v2_transfer/74999/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=500, | |
| peak_lr=1e-5, | |
| decay_steps=20_000, | |
| decay_lr=1e-6, | |
| ), | |
| num_train_steps=20_000, | |
| save_interval=2500, | |
| keep_period=5000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| # | |
| # Low-memory LoRA variant of the above (fits a single 80GB GPU): LoRA on the LLM, frozen base, no EMA. | |
| # | |
| TrainConfig( | |
| name="pi0_fast_robocasa_target_PickPlaceCounterToCabinet_lora", | |
| model=pi0_fast.Pi0FASTConfig( | |
| action_horizon=10, | |
| max_token_len=256, | |
| paligemma_variant="gemma_2b_lora", | |
| ), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[get_ds_meta("PickPlaceCounterToCabinet", "target", "human")], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/pi0_fast_robocasa_pretrain_human300", | |
| asset_id="robocasa365_human300", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("./checkpoints_v2_transfer/74999/params"), | |
| freeze_filter=pi0_fast.Pi0FASTConfig( | |
| action_horizon=10, | |
| max_token_len=256, | |
| paligemma_variant="gemma_2b_lora", | |
| ).get_freeze_filter(), | |
| ema_decay=None, | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=500, | |
| peak_lr=5e-5, | |
| decay_steps=20_000, | |
| decay_lr=5e-6, | |
| ), | |
| num_train_steps=20_000, | |
| save_interval=2500, | |
| keep_period=5000, | |
| batch_size=32, | |
| num_workers=4, | |
| ), | |
| # | |
| # pi0.5 single-task fine-tune on RoboCasa365 PickPlaceCounterToCabinet (target split, 500 human demos). | |
| # pi0.5: discretized state goes into the language prompt, flow timestep via adaRMSNorm; quantile norm | |
| # (needs q01/q99, so we reuse the human300 mixture stats). Starts from the released pi05_base weights. | |
| # | |
| TrainConfig( | |
| name="pi05_robocasa_target_PickPlaceCounterToCabinet", | |
| model=pi0.Pi0Config(pi05=True, action_horizon=10, max_token_len=200), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[get_ds_meta("PickPlaceCounterToCabinet", "target", "human")], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/pi0_fast_robocasa_pretrain_human300", | |
| asset_id="robocasa365_human300", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_base/params"), | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=20_000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=20_000, | |
| save_interval=2500, | |
| keep_period=5000, | |
| batch_size=64, | |
| num_workers=4, | |
| ), | |
| # | |
| # EVAL-ONLY config for the ctc502 60k BASE checkpoint (pi05_ctc502_60k / ctc502_59999_rawparams). | |
| # | |
| # Added 2026-09-22 after checking what that checkpoint was ACTUALLY trained with, rather | |
| # than trusting EVAL_SEEDS.md. Evidence (upstream repo Ronaldo-GOAT/pi05-neg-mesh, file | |
| # training/openpi/src/openpi/training/config.py, entry | |
| # "pi05_robocasa_pickplace_counter_to_cabinet_502_60k"): | |
| # * action_horizon=50, action_dim=32 -- NOT the 10 that | |
| # pi05_robocasa_target_PickPlaceCounterToCabinet declares. Since the client sizes its | |
| # flow-matching noise from the server's advertised action_horizon, serving at 10 both | |
| # changes the noise tensor and runs the action expert at a sequence length the model | |
| # never saw. | |
| # * no AssetsConfig override and repo_id=None -> asset_id=None -> the run fell through to | |
| # norm stats computed from the 502-demo dataset itself, and Orbax therefore wrote them | |
| # FLAT to <ckpt>/assets/norm_stats.json. asset_id=None here reproduces that: | |
| # create_trained_policy() loads <ckpt>/assets/norm_stats.json directly. | |
| # * upstream never assigns use_quantile_norm -> it stayed False -> z-score. Hence | |
| # force_zscore_norm=True. (That is also why the checkpoint's own norm_stats.json has | |
| # q01/q99 = null: nothing needed quantiles.) | |
| # Verified numerically: the checkpoint's flat norm_stats.json and this repo's | |
| # assets/robocasa_ctc502/ctc502_qnorm/norm_stats.json agree on every REAL dimension | |
| # (state 0-15, actions 0-11) to 4.4e-4 / 1.1e-7 -- they are the same ctc502 statistics. | |
| # They differ only on the PADDING dims (state 16-31, actions 12-31), where the flat file | |
| # uses std=1.0 (a no-op) and ctc502_qnorm uses std=0.0. By contrast robocasa365_human300 is | |
| # a different distribution entirely (state.mean differs by 0.86), so EVAL_SEEDS.md's | |
| # instruction to serve this checkpoint under robocasa365_human300 is wrong. | |
| # | |
| # NOT for training -- there is no weight_loader and no data here on purpose. | |
| TrainConfig( | |
| name="pi05_robocasa_ctc502_60k_eval", | |
| model=pi0.Pi0Config(pi05=True, action_dim=32, action_horizon=50, max_token_len=200), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=None, # eval-only: no dataset on the eval nodes, and the | |
| # norm-stat fallback must NOT fire (see force_zscore_norm) | |
| force_zscore_norm=True, # z-score, as trained upstream | |
| ), | |
| ema_decay=None, # the ctc502 lineage is raw non-EMA params | |
| num_train_steps=1, | |
| batch_size=1, | |
| num_workers=0, | |
| ), | |
| # VACE object-swap augmentation: 8 neg-mesh hard objects swapped into PnPCounterToCabinet | |
| # (vace-only, 256 eps = 8 objects x 32). Full fine-tune on the swap dataset for 30k steps. | |
| # | |
| # 2026-09-21: RETARGETED onto the ctc502 lineage, and now the SECOND ARM of a two-arm | |
| # experiment -- identical in every respect to the mimicgen_aug256 entry below except the | |
| # dataset, so the two are directly comparable. See that entry for why each of | |
| # action_horizon=50 / ema_decay=None / the ctc502_qnorm assets / the ctc502 60k | |
| # weight_loader is required; the same reasoning applies verbatim here. | |
| # The earlier note that "the robocasa365_human300 norm stats still apply" is no longer | |
| # true of this entry: the ctc502 weights were trained under ctc502_qnorm. | |
| TrainConfig( | |
| name="pi05_robocasa_target_PickPlaceCounterToCabinet_vace_negmesh", | |
| model=pi0.Pi0Config(pi05=True, action_horizon=50, max_token_len=200), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[ | |
| { | |
| "path": "/scratch/jonghoon/datasets/lerobot_swap_negmesh256_allintra_drop13", | |
| "horizon": 500, | |
| "filter_key": "500_demos", | |
| "task": "PickPlaceCounterToCabinet", | |
| "split": "target", | |
| "source": "human", | |
| } | |
| ], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/robocasa_ctc502", | |
| asset_id="ctc502_qnorm", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_ctc502_60k/params"), | |
| ema_decay=None, | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=30_000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=30_001, | |
| save_steps=(1_000, 1_500, 3_000, 5_000, 10_000, 15_000, 20_000, 25_000, 30_000), | |
| save_interval=2500, # ignored while save_steps is set | |
| keep_period=500, | |
| batch_size=64, | |
| fsdp_devices=1, | |
| num_workers=4, | |
| ), | |
| # MimicGen augmentation on the same 8 neg-mesh hard objects (mlnha/mimicgen-pi05-aug256, | |
| # 256 eps, natural distribution -- AluminumFoil006 yielded 0 of 640 attempts so 7 objects | |
| # are represented). Unlike the VACE set above these are NEW rollouts, so the published | |
| # release ships only rendered video + raw MimicGen HDF5; the LeRobot parquet is rebuilt by | |
| # jobs/mimicgen_build_lerobot.py, which drives robocasa's own reorder_hdf5_action / | |
| # reorder_hdf5_state so the simulator's arm-first layout is mapped onto this modality.json | |
| # rather than copied through (copying through is what scored 0/160 in the sibling GR00T run). | |
| # State/action/modality end up byte-compatible with the negmesh set. | |
| # | |
| # 2026-09-21: RETARGETED onto the ctc502 lineage. Four deliberate departures from the | |
| # negmesh config above -- all four are required together, none is cosmetic: | |
| # action_horizon=50 openpi's stock Pi0Config default (models/pi0_config.py), and what | |
| # the ctc502 lineage and the documented eval protocol both use. | |
| # action_dim stays at its 32 default. | |
| # ema_decay=None the ctc502 lineage disables EMA deliberately -- the EMA copy was | |
| # found broken for this task family, and the eval protocol reads raw | |
| # non-EMA params. (Without this it would inherit 0.99.) | |
| # assets ctc502_qnorm the ctc502 weights were trained under ctc502_qnorm normalisation, | |
| # whose quantiles differ materially from human300 (state q01[0] | |
| # +0.0296 vs -0.1537). Wrong stats here corrupt the warm start. | |
| # NB: the norm_stats.json shipped INSIDE pi05_ctc502_60k_rawparams/ | |
| # has q01/q99 = null and is unusable for pi0.5 quantile norm; the | |
| # file installed under ./assets/robocasa_ctc502/ctc502_qnorm comes | |
| # from Ronaldo-GOAT/pi05-neg-mesh and has all four fields, 32-D. | |
| # weight_loader warm start from the ctc502 60k export instead of pi05_base. It is | |
| # params-only (no train_state), so this is warm-start-only: never | |
| # --resume from it. jobs/ctc502_prep.sbatch stages it into the | |
| # openpi cache so nothing is fetched from GCS at train time. | |
| TrainConfig( | |
| name="pi05_robocasa_target_PickPlaceCounterToCabinet_mimicgen_aug256", | |
| model=pi0.Pi0Config(pi05=True, action_horizon=50, max_token_len=200), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[ | |
| { | |
| "path": "/scratch/jonghoon/datasets/lerobot_mimicgen_aug256_allintra", | |
| "horizon": 500, | |
| "filter_key": "500_demos", | |
| "task": "PickPlaceCounterToCabinet", | |
| "split": "target", | |
| "source": "human", | |
| } | |
| ], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/robocasa_ctc502", | |
| asset_id="ctc502_qnorm", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_ctc502_60k/params"), | |
| ema_decay=None, | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=30_000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=30_001, | |
| save_steps=(1_000, 1_500, 3_000, 5_000, 10_000, 15_000, 20_000, 25_000, 30_000), | |
| save_interval=2500, # ignored while save_steps is set | |
| keep_period=500, | |
| batch_size=64, | |
| fsdp_devices=1, | |
| num_workers=4, | |
| ), | |
| # Pose6DAug action-augmentation on the SAME 8 neg-mesh hard objects | |
| # (Ronaldo-GOAT/transfer :: dataset/lerobot_actaug256_neg_pi05, 256 eps = 8 x 32). | |
| # THIRD ARM of the comparison. The grasps here were re-earned by the pi0.5 ctc502 60k | |
| # policy -- the very checkpoint this config warm-starts from -- so the arm differs from | |
| # its two siblings ONLY in augmentation method: | |
| # vace_negmesh VACE video object-swap | |
| # mimicgen_aug256 MimicGen trajectory generation | |
| # pose6daug_neg256 Pose6DAug simulator action-augmentation <- this entry | |
| # Same 8 objects, same 256-episode budget, same task, same recipe, same warm start, | |
| # and the same 160-episode eval bank covers all three. | |
| # | |
| # Every field below is inherited verbatim from the two sibling arms -- see the | |
| # mimicgen entry above for why action_horizon=50 / ema_decay=None / ctc502_qnorm / | |
| # the ctc502 60k weight_loader are each required. THREE things differ, all deliberate: | |
| # | |
| # num_train_steps=5_001 The user's current instruction is to stop at 5k, NOT the | |
| # 30k the repo README suggests. The +1 is load-bearing: the | |
| # loop is `range(0, num_train_steps)`, so 5_000 has to be | |
| # inside it for a genuine 5000/ checkpoint to exist rather | |
| # than a 4999/ one. | |
| # save_steps every 500 Ten checkpoints, to study EARLY learning densely. Paired | |
| # with keep_period=500: checkpoints.py hardcodes | |
| # max_to_keep=1, so keep_period is the ONLY thing that stops | |
| # each checkpoint being deleted the moment the next lands. | |
| # Every listed step is divisible by 500, so none is pruned. | |
| # decay_steps=30_000 DELIBERATELY UNCHANGED, and this is the subtle one. The | |
| # cosine still decays over 30k even though the run stops at | |
| # 5k, so these 5k steps are bit-identical to the first 5k the | |
| # other two arms saw -- which is the entire point of a | |
| # controlled third arm. The cost is that the run ends about a | |
| # third of the way down the cosine at LR ~1.9e-5 rather than | |
| # at the 2.5e-6 floor. That is intended. Do NOT "fix" it to | |
| # 5_000; ours_config_assert.py asserts it stays 30_000. | |
| TrainConfig( | |
| name="pi05_robocasa_target_PickPlaceCounterToCabinet_pose6daug_neg256", | |
| model=pi0.Pi0Config(pi05=True, action_horizon=50, max_token_len=200), | |
| data=LeRobotRobocasaDataConfig( | |
| data_dirs=[ | |
| { | |
| "path": "/scratch/jonghoon/datasets/lerobot_actaug256_neg_pi05", | |
| "horizon": 500, | |
| "filter_key": "500_demos", | |
| "task": "PickPlaceCounterToCabinet", | |
| "split": "target", | |
| "source": "human", | |
| } | |
| ], | |
| assets=AssetsConfig( | |
| assets_dir="./assets/robocasa_ctc502", | |
| asset_id="ctc502_qnorm", | |
| ), | |
| ), | |
| weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_ctc502_60k/params"), | |
| ema_decay=None, | |
| lr_schedule=_optimizer.CosineDecaySchedule( | |
| warmup_steps=1_000, | |
| peak_lr=2.5e-5, | |
| decay_steps=30_000, | |
| decay_lr=2.5e-6, | |
| ), | |
| num_train_steps=5_001, | |
| save_steps=(500, 1_000, 1_500, 2_000, 2_500, 3_000, 3_500, 4_000, 4_500, 5_000), | |
| save_interval=2500, # ignored while save_steps is set | |
| keep_period=500, | |
| batch_size=64, | |
| fsdp_devices=1, | |
| num_workers=4, | |
| ), | |
| # --- DISABLED 2026-09-21 : LoRA variant, never used. See DISABLED_CONFIGS.md --- | |
| # # Low-memory LoRA variant (single 80GB GPU): LoRA on both PaliGemma and the action expert, no EMA. | |
| # TrainConfig( | |
| # name="pi05_robocasa_target_PickPlaceCounterToCabinet_lora", | |
| # model=pi0.Pi0Config( | |
| # pi05=True, | |
| # action_horizon=10, | |
| # max_token_len=200, | |
| # paligemma_variant="gemma_2b_lora", | |
| # action_expert_variant="gemma_300m_lora", | |
| # ), | |
| # data=LeRobotRobocasaDataConfig( | |
| # data_dirs=[get_ds_meta("PickPlaceCounterToCabinet", "target", "human")], | |
| # assets=AssetsConfig( | |
| # assets_dir="./assets/pi0_fast_robocasa_pretrain_human300", | |
| # asset_id="robocasa365_human300", | |
| # ), | |
| # ), | |
| # weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_base/params"), | |
| # freeze_filter=pi0.Pi0Config( | |
| # pi05=True, | |
| # action_horizon=10, | |
| # max_token_len=200, | |
| # paligemma_variant="gemma_2b_lora", | |
| # action_expert_variant="gemma_300m_lora", | |
| # ).get_freeze_filter(), | |
| # ema_decay=None, | |
| # lr_schedule=_optimizer.CosineDecaySchedule( | |
| # warmup_steps=500, | |
| # peak_lr=5e-5, | |
| # decay_steps=20_000, | |
| # decay_lr=5e-6, | |
| # ), | |
| # num_train_steps=20_000, | |
| # save_interval=2500, | |
| # keep_period=5000, | |
| # batch_size=32, | |
| # num_workers=4, | |
| # ), | |
| ] | |
| if len({config.name for config in _CONFIGS}) != len(_CONFIGS): | |
| raise ValueError("Config names must be unique.") | |
| _CONFIGS_DICT = {config.name: config for config in _CONFIGS} | |
| def cli() -> TrainConfig: | |
| return tyro.extras.overridable_config_cli({k: (k, v) for k, v in _CONFIGS_DICT.items()}) | |
| def get_config(config_name: str) -> TrainConfig: | |
| """Get a config by name.""" | |
| if config_name not in _CONFIGS_DICT: | |
| closest = difflib.get_close_matches(config_name, _CONFIGS_DICT.keys(), n=1, cutoff=0.0) | |
| closest_str = f" Did you mean '{closest[0]}'? " if closest else "" | |
| raise ValueError(f"Config '{config_name}' not found.{closest_str}") | |
| return _CONFIGS_DICT[config_name] | |