Download openpi/code/rm65_config_excerpt.py from LGG100/max: direct link, hf CLI and curl.
- Browser
- Download file 3.79 kB
-
https://huggingface.co/LGG100/max/resolve/main/openpi/code/rm65_config_excerpt.py
- Command line
-
hf download hf://LGG100/max/openpi/code/rm65_config_excerpt.py
-
curl -L -o rm65_config_excerpt.py https://huggingface.co/LGG100/max/resolve/main/openpi/code/rm65_config_excerpt.py
3.79 kB
| # 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) ---- | |
| 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 | |
| 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, | |
| ), | |