From 12740f6be0f6b2999bca2cab391879f69a3c806c Mon Sep 17 00:00:00 2001 From: Pepijn Date: Fri, 17 Jul 2026 17:41:20 +0200 Subject: [PATCH] refactor(rtc): share Pi05 training helpers --- src/lerobot/policies/pi05/modeling_pi05.py | 2 +- .../policies/pi052/configuration_pi052.py | 9 +---- src/lerobot/policies/pi052/modeling_pi052.py | 39 +------------------ tests/policies/common/test_vla_utils.py | 4 +- .../pi052/test_pi052_training_time_rtc.py | 2 +- 5 files changed, 7 insertions(+), 49 deletions(-) diff --git a/src/lerobot/policies/pi05/modeling_pi05.py b/src/lerobot/policies/pi05/modeling_pi05.py index 93f48e1f6..271732526 100644 --- a/src/lerobot/policies/pi05/modeling_pi05.py +++ b/src/lerobot/policies/pi05/modeling_pi05.py @@ -818,7 +818,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch` training_max_delay = int(getattr(self.config, "rtc_training_max_delay", 0)) if training_max_delay <= 0: raise ValueError( - "RTC mode='trained' requires a Pi052 checkpoint trained with " + "RTC mode='trained' requires a checkpoint trained with " "policy.rtc_training_max_delay > 0." ) trained_prefix, trained_prefix_mask = _prepare_trained_rtc_prefix( diff --git a/src/lerobot/policies/pi052/configuration_pi052.py b/src/lerobot/policies/pi052/configuration_pi052.py index 015c2d1fa..018204fda 100644 --- a/src/lerobot/policies/pi052/configuration_pi052.py +++ b/src/lerobot/policies/pi052/configuration_pi052.py @@ -123,9 +123,7 @@ class PI052Config(PI05Config): # Reuse each VLM prefix across independent denoising draws; 1 restores single-draw flow. flow_num_repeats: int = 5 - # Training-time RTC (arXiv:2512.05964). Zero preserves standard flow matching. - rtc_training_max_delay: int = 0 - """Maximum clean-prefix delay sampled for training-time RTC; zero disables it.""" + # Training-time RTC configuration is inherited from PI05Config. # PaLM-style z-loss stabilizes large-vocabulary CE; 0 disables it. text_ce_z_loss_weight: float = 1e-4 @@ -160,11 +158,6 @@ class PI052Config(PI05Config): raise ValueError("fast_tokenizer_validation_samples must be >= 1") if self.fast_tokenizer_max_reconstruction_rmse <= 0 or self.fast_tokenizer_max_dim_rmse <= 0: raise ValueError("FAST tokenizer reconstruction thresholds must be positive") - if not 0 <= self.rtc_training_max_delay < self.chunk_size: - raise ValueError( - "rtc_training_max_delay must satisfy " - f"0 <= delay < chunk_size ({self.chunk_size}), got {self.rtc_training_max_delay}" - ) if self.manual_attention_scope not in {"all", "action"}: raise ValueError( f"manual_attention_scope must be 'all' or 'action', got {self.manual_attention_scope!r}" diff --git a/src/lerobot/policies/pi052/modeling_pi052.py b/src/lerobot/policies/pi052/modeling_pi052.py index ed9347d42..ef0cbe77c 100644 --- a/src/lerobot/policies/pi052/modeling_pi052.py +++ b/src/lerobot/policies/pi052/modeling_pi052.py @@ -37,6 +37,8 @@ from ..pi05.modeling_pi05 import ( ActionSelectKwargs, PI05Policy, PI05Pytorch as PI05PytorchBase, + _build_flow_matching_inputs, + _sample_training_rtc_prefix_mask, make_att_2d_masks, ) from .configuration_pi052 import PI052Config @@ -211,43 +213,6 @@ def _enable_hf_kernels() -> None: logger.info("PI052: HF kernels (Liger) enabled — rope, geglu fused.") -def _sample_training_rtc_prefix_mask( - batch_size: int, - action_horizon: int, - max_delay: int, - device: torch.device, -) -> Tensor | None: - """Sample per-draw clean prefixes for training-time RTC.""" - if max_delay <= 0: - return None - delays = torch.randint(0, max_delay + 1, (batch_size,), device=device) - positions = torch.arange(action_horizon, device=device) - return positions.unsqueeze(0) < delays.unsqueeze(1) - - -def _build_flow_matching_inputs( - actions: Tensor, - noise: Tensor, - time: Tensor, - prefix_mask: Tensor | None, -) -> tuple[Tensor, Tensor]: - """Build noisy actions and scalar/per-token flow times. - - LeRobot's PI0.5 flow uses ``t=0`` for clean data and ``t=1`` for - noise, the reverse of the notation in arXiv:2512.05964. Consequently, - clean RTC prefix tokens receive ``t=0`` here. - """ - if prefix_mask is None: - model_time = time - expanded_time = time[:, None, None] - else: - model_time = time[:, None].expand_as(prefix_mask) - model_time = torch.where(prefix_mask, torch.zeros_like(model_time), model_time) - expanded_time = model_time.unsqueeze(-1) - x_t = expanded_time * noise + (1 - expanded_time) * actions - return x_t, model_time - - def _flow_loss_components(flow_per_dim: Tensor, prefix_mask: Tensor | None) -> tuple[Tensor, Tensor]: """Return each sample's postfix loss sum and valid-element count.""" if prefix_mask is None: diff --git a/tests/policies/common/test_vla_utils.py b/tests/policies/common/test_vla_utils.py index 4d57b15ee..61707eef6 100644 --- a/tests/policies/common/test_vla_utils.py +++ b/tests/policies/common/test_vla_utils.py @@ -56,8 +56,8 @@ def test_create_sinusoidal_pos_embedding_matches_openpi_formula(): def test_create_sinusoidal_pos_embedding_validation(): with pytest.raises(ValueError, match="divisible by 2"): create_sinusoidal_pos_embedding(torch.zeros(2), 7, 4e-3, 4.0, device=torch.device("cpu")) - with pytest.raises(ValueError, match="batch_size"): - create_sinusoidal_pos_embedding(torch.zeros(2, 2), 8, 4e-3, 4.0, device=torch.device("cpu")) + with pytest.raises(ValueError, match="must have shape"): + create_sinusoidal_pos_embedding(torch.zeros(2, 2, 2), 8, 4e-3, 4.0, device=torch.device("cpu")) def test_make_att_2d_masks_docstring_cases(): diff --git a/tests/policies/pi052/test_pi052_training_time_rtc.py b/tests/policies/pi052/test_pi052_training_time_rtc.py index 398203d64..6993f433c 100644 --- a/tests/policies/pi052/test_pi052_training_time_rtc.py +++ b/tests/policies/pi052/test_pi052_training_time_rtc.py @@ -16,13 +16,13 @@ from torch import nn pytest.importorskip("transformers") from lerobot.policies.pi05.modeling_pi05 import ( # noqa: E402 + _build_flow_matching_inputs, _prepare_trained_rtc_prefix, create_sinusoidal_pos_embedding, ) from lerobot.policies.pi052.configuration_pi052 import PI052Config # noqa: E402 from lerobot.policies.pi052.modeling_pi052 import ( # noqa: E402 PI05Pytorch as PI052Pytorch, - _build_flow_matching_inputs, _flow_loss_per_sample, _reduce_flow_loss, )