File size: 3,786 Bytes
3acefc3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 | # 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,
),
|