robopro_jax_30000 / jax_30000 /train_config.py
mzxuan's picture
Move JAX pi05 30000 checkpoint from repo root into jax_30000/
c5224a4 verified
Raw
History Blame Contribute Delete
2.47 kB
# Train / eval config for this checkpoint: `pi05_robopro_top_cam_jax`
#
# This is the exact openpi TrainConfig used to fine-tune and to load this
# checkpoint. To use it, add the TrainConfig(...) entry below to the `_CONFIGS`
# list in `src/openpi/training/config.py` of your openpi checkout, then load it
# with `openpi.training.config.get_config("pi05_robopro_top_cam_jax")`.
#
# Symbols referenced (already imported at the top of openpi's config.py):
# TrainConfig, DataConfig, LeRobotAlohaDataConfig
# pi0_config = openpi.models.pi0_config
# _transforms = openpi.transforms
# weight_loaders = openpi.training.weight_loaders
# _optimizer = openpi.training.optimizer
#
# Dataset: robopro top-cam LeRobot v2.1 (`roboreal_lerobot`), robot_type roboreal,
# 25 fps, 14-DoF dual-arm, cams countertop/left/right. Set
# HF_LEROBOT_HOME=<parent> so repo_id `roboreal_lerobot` resolves locally.
# Base weights: JAX pi05_base from gs://openpi-assets (auto-download).
TrainConfig(
name="pi05_robopro_top_cam_jax",
model=pi0_config.Pi0Config(pi05=True),
data=LeRobotAlohaDataConfig(
repo_id="roboreal_lerobot",
# Map the dataset's raw feature keys -> the model's expected keys.
# NOTE: cam_high is fed from the COUNTERTOP (overhead) camera.
repack_transforms=_transforms.Group(inputs=[
_transforms.RepackTransform({
"images": {
"cam_high": "observation.images.countertop",
"cam_left_wrist": "observation.images.left",
"cam_right_wrist": "observation.images.right",
},
"state": "observation.state",
"actions": "action",
"prompt": "prompt",
})
]),
base_config=DataConfig(
prompt_from_task=True,
),
# (defaults inherited from LeRobotAlohaDataConfig:)
# adapt_to_pi=True, use_delta_joint_actions=True
# -> arm joints trained as delta, grippers absolute;
# AbsoluteActions on output => returned actions are ABSOLUTE.
),
weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_base/params"),
lr_schedule=_optimizer.CosineDecaySchedule(decay_steps=30_000),
num_train_steps=30_000,
batch_size=192, # 3 GPUs = 64/GPU (must be divisible by device count)
num_workers=16,
fsdp_devices=1,
),