# RM65 excerpt of openpi-fintune src/openpi/training/config.py. # Paste LeRobotRM65DataConfig next to the other DataConfigFactory classes and the TrainConfig into _CONFIGS. # Requires openpi.policies.b601_policy (included under code/src/openpi/policies/). # ---- DataConfigFactory (module level) ---- @dataclasses.dataclass(frozen=True) class LeRobotRM65DataConfig(DataConfigFactory): """Data config for fine-tuning on the RealMan RM65 right-arm LeRobot v2.1 dataset. Same layout as B601 -- 7-dim state/action (right_q1..q6 in rad, right_gripper 0=closed..100=open) and two cameras (top, wrist) -- so the B601 policy transforms are reused as-is. Unlike B601, observation.state and action share one coordinate convention, so the joint actions can be trained as deltas relative to the current state. """ # Injected as the prompt for every sample. Leave None to take the prompt from the dataset's tasks.jsonl. default_prompt: str | None = None # Train the 6 joints as deltas relative to the current state; the gripper stays absolute. use_delta_joint_actions: bool = False @override def create(self, assets_dirs: pathlib.Path, model_config: _model.BaseModelConfig) -> DataConfig: repack_transform = _transforms.Group( inputs=[ _transforms.RepackTransform( { "observation/top": "observation.images.top", "observation/wrist": "observation.images.wrist", "observation/state": "observation.state", "actions": "action", "prompt": "prompt", } ) ] ) data_transforms = _transforms.Group( inputs=[b601_policy.B601Inputs(model_type=model_config.model_type)], outputs=[b601_policy.B601Outputs()], ) if self.use_delta_joint_actions: 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 = ModelTransformFactory(default_prompt=self.default_prompt)(model_config) return dataclasses.replace( self.create_base_config(assets_dirs, model_config), repack_transforms=repack_transform, data_transforms=data_transforms, model_transforms=model_transforms, ) # ---- TrainConfig (entry of _CONFIGS) ---- # Same as pi05_rm65_4tasks, but on rm65/4tasks_v1_trim (scripts/trim_static_frames.py): the static head # of every episode is dropped and the static tail is cut to 10 frames, 66,121 -> 58,399 frames. # Short mid-episode pauses are kept. TrainConfig( name="pi05_rm65_4tasks_trim", model=pi0_config.Pi0Config( pi05=True, action_horizon=30, paligemma_variant="gemma_2b_lora", ), data=LeRobotRM65DataConfig( repo_id="rm65/4tasks_v1_trim", base_config=DataConfig( local_root=pathlib.Path("/shared/user64/workspace/yuhao/pi/data"), prompt_from_task=True, action_sequence_keys=("action",), ), ), weight_loader=weight_loaders.CheckpointWeightLoader("gs://openpi-assets/checkpoints/pi05_base/params"), freeze_filter=pi0_config.Pi0Config( pi05=True, action_horizon=30, paligemma_variant="gemma_2b_lora", ).get_freeze_filter(), num_train_steps=30_000, batch_size=32, num_workers=8, ema_decay=None, ),