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,
    ),