max / openpi /code /rm65_config_excerpt.py
LGG100's picture
Add files using upload-large-folder tool
3acefc3 verified
Raw History Blame Contribute Delete
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) ----
@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,
),