diff --git a/README.md b/README.md
index 63081acae..5842c307b 100644
--- a/README.md
+++ b/README.md
@@ -101,13 +101,13 @@ lerobot-train \
--dataset.repo_id=lerobot/aloha_mobile_cabinet
```
-| Category | Models |
-| -------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ |
-| **Imitation Learning** | [ACT](./docs/source/policy_act_README.md), [Diffusion](./docs/source/policy_diffusion_README.md), [VQ-BeT](./docs/source/policy_vqbet_README.md), [Multitask DiT Policy](./docs/source/policy_multi_task_dit_README.md) |
-| **Reinforcement Learning** | [HIL-SERL](./docs/source/hilserl.mdx), [TDMPC](./docs/source/policy_tdmpc_README.md) & QC-FQL (coming soon) |
-| **VLAs Models** | [Pi0](./docs/source/pi0.mdx), [Pi0Fast](./docs/source/pi0fast.mdx), [Pi0.5](./docs/source/pi05.mdx), [GR00T N1.7](./docs/source/policy_groot_README.md), [SmolVLA](./docs/source/policy_smolvla_README.md), [XVLA](./docs/source/xvla.mdx), [EO-1](./docs/source/eo1.mdx), [MolmoAct2](./docs/source/molmoact2.mdx), [WALL-OSS](./docs/source/walloss.mdx), [EVO1](./docs/source/evo1.mdx) |
-| **World Models** | [VLA-JEPA](./docs/source/vla_jepa.mdx), [LingBot-VA](./docs/source/lingbot_va.mdx), [FastWAM](./docs/source/fastwam.mdx) |
-| **Reward Models** | [SARM](./docs/source/sarm.mdx), [TOPReward](./docs/source/topreward.mdx), [Robometer](./docs/source/robometer.mdx) |
+| Category | Models |
+| -------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
+| **Imitation Learning** | [ACT](./docs/source/policy_act_README.md), [Diffusion](./docs/source/policy_diffusion_README.md), [VQ-BeT](./docs/source/policy_vqbet_README.md), [Multitask DiT Policy](./docs/source/policy_multi_task_dit_README.md) |
+| **Reinforcement Learning** | [HIL-SERL](./docs/source/hilserl.mdx), [TDMPC](./docs/source/policy_tdmpc_README.md) & QC-FQL (coming soon) |
+| **VLAs Models** | [Pi0](./docs/source/pi0.mdx), [Pi0Fast](./docs/source/pi0fast.mdx), [Pi0.5](./docs/source/pi05.mdx), [Pi052](./docs/source/pi052.mdx), [GR00T N1.7](./docs/source/policy_groot_README.md), [SmolVLA](./docs/source/policy_smolvla_README.md), [XVLA](./docs/source/xvla.mdx), [EO-1](./docs/source/eo1.mdx), [MolmoAct2](./docs/source/molmoact2.mdx), [WALL-OSS](./docs/source/walloss.mdx), [EVO1](./docs/source/evo1.mdx) |
+| **World Models** | [VLA-JEPA](./docs/source/vla_jepa.mdx), [LingBot-VA](./docs/source/lingbot_va.mdx), [FastWAM](./docs/source/fastwam.mdx) |
+| **Reward Models** | [SARM](./docs/source/sarm.mdx), [TOPReward](./docs/source/topreward.mdx), [Robometer](./docs/source/robometer.mdx) |
Similarly to the hardware, you can easily implement your own policy & leverage LeRobot's data collection, training, and visualization tools, and share your model to the HF Hub.
diff --git a/docs/source/_toctree.yml b/docs/source/_toctree.yml
index 7f7a34e6a..cc657bf3a 100644
--- a/docs/source/_toctree.yml
+++ b/docs/source/_toctree.yml
@@ -63,6 +63,8 @@
title: π₀-FAST (Pi0Fast)
- local: pi05
title: π₀.₅ (Pi05)
+ - local: pi052
+ title: π₀.₅ with language supervision (Pi052)
- local: molmoact2
title: MolmoAct2
- local: vla_jepa
diff --git a/docs/source/pi052.mdx b/docs/source/pi052.mdx
new file mode 100644
index 000000000..dbea9bd69
--- /dev/null
+++ b/docs/source/pi052.mdx
@@ -0,0 +1,274 @@
+# π₀.₅ with language supervision (Pi052)
+
+Pi052 extends [Pi05](./pi05) with a trainable PaliGemma language head and a
+runtime that alternates language generation with action generation. A single
+checkpoint can predict a low-level subtask, optionally update memory or answer
+visual questions, and condition its flow-matching action expert on that text.
+
+Use Pi05 when you only need task-conditioned actions. Use Pi052 when the policy
+must generate or consume intermediate language during a rollout.
+
+## How Pi052 differs from Pi05
+
+| Capability | Pi05 | Pi052 |
+| ------------------- | ------------------------------------------------------ | --------------------------------------------------------------------------------- |
+| Action model | PaliGemma vision-language prefix + Gemma action expert | Same base architecture |
+| Language head | Not trained for runtime generation | Re-enabled and trained with text cross-entropy |
+| Action conditioning | Episode task | Active low-level subtask plus normalized robot state |
+| Training targets | Flow-matching actions | Flow actions, recipe-selected text, and optional FAST action tokens |
+| Dataset requirement | Standard images, state, actions, and task | The same fields plus language annotations for every language capability you train |
+| Rollout | Direct task-to-action policy | Hierarchical task → subtask → action loop, with optional memory and VQA |
+
+Pi052 can initialize from a Pi05 checkpoint. The policy architecture remains
+compatible, while Pi052 builds its own processors so recipe labels and FAST
+labels are not silently replaced by the Pi05 processor stack.
+
+## Install
+
+Install LeRobot with the PI dependencies:
+
+```bash
+git clone https://github.com/huggingface/lerobot.git
+cd lerobot
+python -m venv .venv
+source .venv/bin/activate
+pip install -e ".[pi]"
+```
+
+The `pi` extra includes the PaliGemma/FAST dependencies. Install
+`liger-kernel` for the supported fused training kernels; optional FlashRT
+backends also require the Hugging Face `kernels` package and a supported CUDA
+GPU.
+
+## Prepare language-annotated data
+
+Pi052 does not infer supervised subtasks from a normal LeRobot dataset during
+training. The dataset must contain the language targets used by the selected
+recipe in the optional `language_persistent` and `language_events` columns.
+
+At minimum, annotate a continuous `subtask` timeline so each training frame has
+an active low-level instruction. Add `memory`, VQA, interjections, and speech
+annotations only if the recipe trains those capabilities.
+
+The provided recipes are:
+
+| Recipe | Required annotations | Trains |
+| ------------------------------------- | ----------------------------------------------------------------------- | ------------------------------------------------------------------ |
+| `recipes/subtask.yaml` | `subtask` | Subtask prediction and subtask-conditioned actions |
+| `recipes/subtask_joint.yaml` | `subtask` | Paper-style joint sequence: subtask text and actions in one sample |
+| `recipes/subtask_mem.yaml` | `subtask`, `memory` | Subtasks, actions, and memory updates |
+| `recipes/subtask_mem_vqa_speech.yaml` | `subtask`, `memory`, `vqa`; interjection/speech rows for those branches | Subtasks, actions, memory, VQA, and spoken replies |
+
+The blend recipes factorize training into separate high-level (task → subtask)
+and low-level (subtask → actions) samples, matching how inference decomposes
+π(a|o, subtask)·π(subtask|o, task). `recipes/subtask_joint.yaml` instead uses
+the π0.5 paper's single-sequence layout — the supervised subtask span is
+attended causally and conditions the FAST and flow losses in the same forward.
+Checkpoints trained with the joint recipe must set
+`--policy.joint_subtask_conditioning=true` at inference so the flow prefix
+rebuilds the same layout (task turn with state, then the generated subtask as a
+causal assistant turn); leave it `false` for the blend recipes.
+
+Use `lerobot-annotate` to generate these columns. The repository includes a
+Hugging Face Jobs launcher that you can edit for your source and destination
+datasets. For a local annotation run, first install
+`pip install -e ".[annotations]"`:
+
+```bash
+HF_TOKEN=hf_... uv run python examples/annotations/run_hf_job.py
+```
+
+Before a long training run, inspect several episodes and verify that subtasks
+are temporally correct and cover the full demonstration. See
+[Annotation Pipeline](./annotation_pipeline) for generation and validation, and
+[Language Columns and Recipes](./language_and_recipes) for the schema and
+recipe resolver.
+
+
+ If a dataset has no language columns, recipe rendering becomes a no-op and
+ Pi052 falls back to the plain Pi05 prompt path. This is useful for
+ compatibility but does not train the language planner.
+
+
+## Train Pi052
+
+This example initializes Pi052 from the public Pi05 base checkpoint and trains
+the default subtask-and-memory recipe:
+
+```bash
+lerobot-train \
+ --dataset.repo_id=${HF_USER}/my_language_annotated_dataset \
+ --policy.type=pi052 \
+ --policy.pretrained_path=lerobot/pi05_base \
+ --policy.recipe_path=recipes/subtask_mem.yaml \
+ --policy.dtype=bfloat16 \
+ --policy.device=cuda \
+ --policy.freeze_vision_encoder=false \
+ --policy.gradient_checkpointing=true \
+ --batch_size=8 \
+ --steps=30000 \
+ --output_dir=outputs/pi052 \
+ --job_name=pi052 \
+ --wandb.enable=true
+```
+
+For subtask-only data, change the recipe to `recipes/subtask.yaml` and disable
+memory during rollout. Start with a small run and confirm that W&B examples show
+the expected prompt, text target, and action endpoints before scaling up.
+
+### Main training controls
+
+| Option | Default | Purpose |
+| ----------------------------------- | -------------------------: | ------------------------------------------------------------------- |
+| `policy.recipe_path` | `recipes/subtask_mem.yaml` | Selects the language/action objective mixture |
+| `policy.text_loss_weight` | `1.0` | Language-head cross-entropy weight; `0` disables text training |
+| `policy.flow_loss_weight` | `10.0` | Continuous action flow-loss weight |
+| `policy.enable_fast_action_loss` | `true` | Adds discrete FAST action-token supervision |
+| `policy.fast_action_loss_weight` | `1.0` | FAST cross-entropy weight |
+| `policy.knowledge_insulation` | `true` | Blocks action-loss gradients through the VLM K/V path |
+| `policy.flow_num_repeats` | `5` | Reuses one VLM prefix for independent denoising targets |
+| `policy.lm_head_lr_scale` | `1.0` | Scales language-head learning rate; `1.0` uses the base rate |
+| `policy.fast_skip_tokens` | `1152` | FAST id offset; skips ``+`` so VQA and FAST never collide |
+| `policy.joint_subtask_conditioning` | `false` | Rebuilds the joint-sequence prefix at inference (see recipes) |
+
+`fast_skip_tokens=1152` places FAST codes below PaliGemma's `` range.
+openpi's pi0-FAST convention is `128` (FAST occupies the `` ids); use that
+value only to stay weight-compatible with checkpoints trained that way, and
+avoid combining it with the VQA recipe, whose `` targets would share
+embedding rows with FAST codes.
+
+The loss weights are starting points, not dataset-independent constants. Track
+flow loss and text/FAST losses separately, and inspect generated subtasks rather
+than selecting a checkpoint from total loss alone.
+
+### Dataset-specific FAST tokenizer
+
+The universal FAST tokenizer works out of the box. For a large or
+embodiment-specific dataset, Pi052 can fit and cache a tokenizer on normalized
+actions before training:
+
+```bash
+lerobot-train \
+ ... \
+ --policy.auto_fit_fast_tokenizer=true \
+ --policy.fast_tokenizer_fit_samples=4096
+```
+
+The fit runs once per dataset/tokenizer configuration. Keep
+`auto_fit_fast_tokenizer=false` when you do not want the extra preprocessing
+pass.
+
+## Training performance
+
+Pi052 uses optimized training paths by default:
+
+- batches repeated flow targets and suffix projections instead of replaying
+ small operations in Python;
+- caches constant action masks and computes RoPE positions once per forward;
+- selects the text/FAST cross-entropy implementation from target shape and
+ sparsity;
+- skips the mathematically dead VLM/vision backward on knowledge-insulated,
+ flow-only batches;
+- uses native non-reentrant SigLIP layer checkpointing when gradient
+ checkpointing is enabled; and
+- retains the Liger RoPE/GeGLU kernels while avoiding the slower LayerNorm
+ patch at SigLIP shapes.
+
+Optional training backends are disabled by default:
+
+| Option | When to try it |
+| -------------------------------------- | ------------------------------------------------------------------------------------------------- |
+| `policy.use_flashrt_adarms=true` | Fused adaptive RMSNorm and gated residuals on supported CUDA GPUs |
+| `policy.use_compiled_text_ce=true` | Compiled materialized-logit CE buckets |
+| `policy.use_compiled_vision=true` | Compiled vision only when the vision pass has no gradients |
+| `policy.use_flex_attention=true` | Profiled CUDA setups with knowledge insulation and `flow_num_repeats > 1`; otherwise SDPA is used |
+| `policy.use_manual_attention=true` | Explicitly profiled shapes where materialized attention is faster |
+| `policy.manual_attention_scope=action` | Restricts manual attention to action queries |
+
+Do not enable every backend blindly. Flex and manual attention are mutually
+exclusive, and attention/AdaRMS alternatives require knowledge insulation.
+The benchmark-best configuration used compiled text CE and FlashRT AdaRMS,
+with Flex/manual attention and compiled vision disabled.
+
+### Reported training benchmarks
+
+These benchmarks measure complete optimizer steps with three real camera
+inputs, BF16 transformer/action execution, FP32 vision, fused AdamW, and no
+video decoding or network I/O. Results vary with GPU, batch shape, annotation
+mixture, and checkpointing:
+
+| Workload | RTX PRO 6000 Blackwell | A100 80 GB |
+| -------------------------- | -------------------------: | -------------------------: |
+| Full flow + text, batch 1 | 4.75× vs checkpointing off | 3.33× vs checkpointing off |
+| Full flow + text, batch 8 | 2.16× vs checkpointing off | 1.66× vs checkpointing off |
+| Full flow + text, batch 64 | 1.24× vs checkpointing on | 1.15× vs checkpointing on |
+| Flow-only, batch 1 | 3.70× vs checkpointing off | 3.58× vs checkpointing off |
+| Flow-only, batch 64 | 3.76× vs checkpointing on | 3.61× vs checkpointing on |
+
+On those 80 GB GPUs, full training was fastest without gradient checkpointing
+through batch 8, then required checkpointing at batch 16 and above. Treat that
+as a tuning rule to test on your hardware, not a universal threshold. Flow-only
+means both text and FAST supervision are disabled; it is useful for action-only
+ablation or post-training but does not learn the language runtime.
+
+## Inference performance
+
+Pi052 has two inference loops, and both avoid repeatedly encoding the expensive
+multimodal prefix:
+
+1. **Action denoising** encodes the image/language prefix once, reuses its KV
+ cache across flow steps, precomputes the timestep schedule on-device, and
+ crops temporary suffix K/V instead of cloning the prefix cache.
+2. **Language decoding** uses autoregressive KV caching, so each new token only
+ processes the sampled token against cached image/language keys instead of
+ rerunning the full prefix.
+
+The runtime also runs language and actions at different rates. Increase
+`--subtask_chunks_per_gen` when a subtask remains valid across several action
+chunks, lower `--high_level_hz`, or use `--direct_subtask` to bypass language
+generation entirely. These settings reduce compute but also slow replanning.
+
+`--fp8` enables the optional FlashRT inference MLP swap on supported CUDA GPUs.
+It calibrates on the first observation and falls back to BF16 when unavailable;
+because FP8 can change outputs slightly, validate task success before using it
+for production rollouts.
+
+## Run a checkpoint
+
+RoboCasa:
+
+```bash
+MUJOCO_GL=egl lerobot-rollout \
+ --policy.path=lerobot/pi052_robocasa \
+ --sim --sim.task=CloseFridge --sim.split=pretrain \
+ --task="close the fridge" \
+ --disable_memory \
+ --sim.render_size=384 \
+ --sim.views=robot0_agentview_left,robot0_eye_in_hand,robot0_agentview_right \
+ --mode=action --ctrl_hz=20
+```
+
+Open `http://localhost:8010` for the live view. Without
+`--sim.direct_subtask`, Pi052 generates the low-level subtask; with it, each
+prompt becomes the action policy's subtask directly.
+
+The same runtime supports real robots. See [Interactive language
+control](./inference#interactive-language-control) for the real-arm command,
+safety behavior, and runtime controls.
+
+## Troubleshooting
+
+- **No text loss or generated subtasks:** confirm the selected recipe can bind
+ the annotations on sampled frames and that `policy.text_loss_weight > 0`.
+- **Subtasks look plausible but actions fail:** verify subtask boundaries,
+ normalized state/action statistics, and that low-level recipe samples are
+ present.
+- **Text collapses to repeated or location tokens:** inspect text-target
+ coverage, language-head learning rate, and the balance between flow, FAST,
+ and text losses.
+- **Out of memory:** reduce batch size first, then enable gradient
+ checkpointing. Do not enable compiled or alternative attention backends
+ without profiling their memory on your camera count.
+- **Slow rollout:** separate action latency from language latency, then tune
+ `--subtask_chunks_per_gen`, `--high_level_hz`, and the number of flow
+ inference steps.
diff --git a/docs/source/pi0fast.mdx b/docs/source/pi0fast.mdx
index 15dff8071..3c4fc86fd 100644
--- a/docs/source/pi0fast.mdx
+++ b/docs/source/pi0fast.mdx
@@ -109,15 +109,21 @@ lerobot-train \
### Key Training Parameters
-| Parameter | Description | Default |
-| -------------------------------------- | -------------------------------------------------- | ------------------------------- |
-| `--policy.gradient_checkpointing=true` | Reduces memory usage significantly during training | `false` |
-| `--policy.dtype=bfloat16` | Use mixed precision training for efficiency | `float32` |
-| `--policy.chunk_size` | Number of action steps to predict (action horizon) | `50` |
-| `--policy.n_action_steps` | Number of action steps to execute | `50` |
-| `--policy.max_action_tokens` | Maximum number of FAST tokens per action chunk | `256` |
-| `--policy.action_tokenizer_name` | FAST tokenizer to use | `lerobot/fast-action-tokenizer` |
-| `--policy.compile_model=true` | Enable torch.compile for faster training | `false` |
+| Parameter | Description | Default |
+| --------------------------------------- | -------------------------------------------------- | ------------------------------- |
+| `--policy.gradient_checkpointing=true` | Reduces memory usage significantly during training | `false` |
+| `--policy.dtype=bfloat16` | Use mixed precision training for efficiency | `float32` |
+| `--policy.chunk_size` | Number of action steps to predict (action horizon) | `50` |
+| `--policy.n_action_steps` | Number of action steps to execute | `50` |
+| `--policy.max_action_tokens` | Maximum number of FAST tokens per action chunk | `256` |
+| `--policy.action_tokenizer_name` | FAST tokenizer to use | `lerobot/fast-action-tokenizer` |
+| `--policy.auto_fit_fast_tokenizer=true` | Fit and cache a tokenizer for the training dataset | `false` |
+| `--policy.compile_model=true` | Enable torch.compile for faster training | `false` |
+
+Set `--policy.auto_fit_fast_tokenizer=true` to sample action chunks from the
+training dataset and cache a fitted tokenizer under
+`~/.cache/lerobot/fast_tokenizers`. This also works when fine-tuning with
+`--policy.path`; leave it disabled to retain the checkpoint's tokenizer.
## Inference
diff --git a/pyproject.toml b/pyproject.toml
index 9e88e8eca..868865f53 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -150,6 +150,7 @@ pygame-dep = ["pygame>=2.5.1,<2.7.0"]
# There is no cmeel-urdfdom 5.x; <5 selects the 4.x ABI the placo/pin wheels are built against.
placo-dep = ["placo>=0.9.6,<0.9.16", "cmeel-urdfdom>=4,<5", "cmeel-tinyxml2<11"]
transformers-dep = ["transformers>=5.4.0,<5.6.0"]
+sentencepiece-dep = ["sentencepiece>=0.2.0,<0.3.0"] # FAST action tokenizer backend (pi052, pi0_fast)
grpcio-dep = ["grpcio>=1.73.1,<2.0.0", "protobuf>=6.31.1,<8.0.0"]
accelerate-dep = ["accelerate>=1.14.0,<2.0.0"]
can-dep = ["python-can>=4.2.0,<5.0.0"]
@@ -212,7 +213,7 @@ wallx = [
"torchdiffeq>=0.2.4,<0.3.0",
"lerobot[qwen-vl-utils-dep]",
]
-pi = ["lerobot[transformers-dep]", "lerobot[scipy-dep]"]
+pi = ["lerobot[transformers-dep]", "lerobot[scipy-dep]", "lerobot[sentencepiece-dep]"]
molmoact2 = ["lerobot[transformers-dep]", "lerobot[peft-dep]", "lerobot[scipy-dep]"]
smolvla = ["lerobot[transformers-dep]", "num2words>=0.5.14,<0.6.0", "lerobot[accelerate-dep]"]
multi_task_dit = ["lerobot[transformers-dep]", "lerobot[diffusers-dep]"]
diff --git a/scripts/backfill_pi052_action_tokenizer.py b/scripts/backfill_pi052_action_tokenizer.py
new file mode 100644
index 000000000..d97f62d37
--- /dev/null
+++ b/scripts/backfill_pi052_action_tokenizer.py
@@ -0,0 +1,227 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Audit or backfill checkpoint-local FAST artifacts for PI052 model repositories."""
+
+from __future__ import annotations
+
+import argparse
+import hashlib
+import io
+import json
+from pathlib import Path, PurePosixPath
+from typing import Any
+
+from huggingface_hub import CommitOperationAdd, HfApi, hf_hub_download
+
+DEFAULT_REPOSITORIES = (
+ "pepijn223/pi052_atomic4_01_baseline",
+ "pepijn223/pi052_atomic4_02_lr_1e5",
+ "pepijn223/pi052_atomic4_03_recipe_50_50",
+ "pepijn223/pi052_atomic4_04_flow_weight_10",
+ "pepijn223/pi052_atomic4_05_flow_repeat_1",
+ "pepijn223/pi052_atomic4_06_ki_off",
+)
+CHECKPOINT_DIRECTORIES = (
+ "",
+ "checkpoints/003000/pretrained_model",
+ "checkpoints/006000/pretrained_model",
+ "checkpoints/009000/pretrained_model",
+ "checkpoints/012000/pretrained_model",
+)
+TOKENIZER_DIRECTORY = "action_tokenizer"
+
+
+def artifact_fingerprint(files: list[tuple[str, bytes]]) -> str:
+ digest = hashlib.sha256()
+ for relative_path, content in sorted(files):
+ encoded_path = relative_path.encode()
+ digest.update(len(encoded_path).to_bytes(8, "big"))
+ digest.update(encoded_path)
+ digest.update(len(content).to_bytes(8, "big"))
+ digest.update(content)
+ return digest.hexdigest()
+
+
+def tokenizer_files(tokenizer_path: Path) -> list[tuple[str, Path]]:
+ return [
+ (path.relative_to(tokenizer_path).as_posix(), path)
+ for path in sorted(tokenizer_path.rglob("*"))
+ if path.is_file()
+ ]
+
+
+def _repo_path(directory: str, filename: str) -> str:
+ return (PurePosixPath(directory) / filename).as_posix() if directory else filename
+
+
+def _download_json(repo_id: str, path_in_repo: str, revision: str | None = None) -> dict[str, Any]:
+ path = hf_hub_download(repo_id, path_in_repo, repo_type="model", revision=revision)
+ return json.loads(Path(path).read_text())
+
+
+def make_portable_preprocessor(config: dict[str, Any]) -> dict[str, Any]:
+ config = json.loads(json.dumps(config))
+ action_steps = [
+ step for step in config["steps"] if step.get("registry_name") == "action_tokenizer_processor"
+ ]
+ if len(action_steps) != 1:
+ raise ValueError(f"Expected one action tokenizer step, found {len(action_steps)}")
+ action_step = action_steps[0]
+ action_step["config"]["action_tokenizer_name"] = TOKENIZER_DIRECTORY
+ action_step["artifacts"] = {"action_tokenizer_name": TOKENIZER_DIRECTORY}
+
+ recipe_steps = [
+ step for step in config["steps"] if step.get("registry_name") == "render_messages_processor"
+ ]
+ if len(recipe_steps) != 1 or not recipe_steps[0].get("config", {}).get("recipe"):
+ raise ValueError("PI052 preprocessor does not contain an embedded training recipe")
+ return config
+
+
+def _json_operation(path_in_repo: str, content: dict[str, Any]) -> CommitOperationAdd:
+ serialized = (json.dumps(content, indent=2) + "\n").encode()
+ return CommitOperationAdd(path_in_repo=path_in_repo, path_or_fileobj=io.BytesIO(serialized))
+
+
+def prepare_operations(
+ repo_id: str,
+ tokenizer_path: Path,
+ revision: str | None = None,
+) -> list[CommitOperationAdd]:
+ operations: list[CommitOperationAdd] = []
+ files = tokenizer_files(tokenizer_path)
+ for directory in CHECKPOINT_DIRECTORIES:
+ preprocessor_path = _repo_path(directory, "policy_preprocessor.json")
+ operations.append(
+ _json_operation(
+ preprocessor_path,
+ make_portable_preprocessor(_download_json(repo_id, preprocessor_path, revision)),
+ )
+ )
+ for relative_path, local_path in files:
+ operations.append(
+ CommitOperationAdd(
+ path_in_repo=_repo_path(
+ directory,
+ f"{TOKENIZER_DIRECTORY}/{relative_path}",
+ ),
+ path_or_fileobj=str(local_path),
+ )
+ )
+ return operations
+
+
+def audit_repository(
+ api: HfApi,
+ repo_id: str,
+ expected_tokenizer_fingerprint: str,
+ revision: str | None = None,
+) -> None:
+ info = api.model_info(repo_id, revision=revision)
+ repository_files = {sibling.rfilename for sibling in info.siblings or []}
+
+ for directory in CHECKPOINT_DIRECTORIES:
+ preprocessor_path = _repo_path(directory, "policy_preprocessor.json")
+ policy_config_path = _repo_path(directory, "config.json")
+ postprocessor_path = _repo_path(directory, "policy_postprocessor.json")
+ for required_path in (preprocessor_path, policy_config_path, postprocessor_path):
+ if required_path not in repository_files:
+ raise FileNotFoundError(f"{repo_id}@{revision or 'main'} is missing {required_path}")
+
+ preprocessor = _download_json(repo_id, preprocessor_path, revision)
+ portable_preprocessor = make_portable_preprocessor(preprocessor)
+ if preprocessor != portable_preprocessor:
+ raise ValueError(f"{repo_id}:{preprocessor_path} is not portable")
+
+ normalizer_steps = [
+ step for step in preprocessor["steps"] if step.get("registry_name") == "normalizer_processor"
+ ]
+ if len(normalizer_steps) != 1 or "state_file" not in normalizer_steps[0]:
+ raise ValueError(f"{repo_id}:{preprocessor_path} is missing normalizer state metadata")
+ normalizer_path = _repo_path(directory, normalizer_steps[0]["state_file"])
+ if normalizer_path not in repository_files:
+ raise FileNotFoundError(f"{repo_id} is missing {normalizer_path}")
+
+ remote_tokenizer_files: list[tuple[str, bytes]] = []
+ for relative_path in _tokenizer_relative_paths(repository_files, directory):
+ path_in_repo = _repo_path(directory, f"{TOKENIZER_DIRECTORY}/{relative_path}")
+ downloaded = hf_hub_download(repo_id, path_in_repo, repo_type="model", revision=revision)
+ remote_tokenizer_files.append((relative_path, Path(downloaded).read_bytes()))
+ fingerprint = artifact_fingerprint(remote_tokenizer_files)
+ if fingerprint != expected_tokenizer_fingerprint:
+ raise ValueError(
+ f"{repo_id}:{_repo_path(directory, TOKENIZER_DIRECTORY)} fingerprint "
+ f"{fingerprint} != {expected_tokenizer_fingerprint}"
+ )
+
+
+def _tokenizer_relative_paths(repository_files: set[str], directory: str) -> list[str]:
+ prefix = _repo_path(directory, TOKENIZER_DIRECTORY).rstrip("/") + "/"
+ paths = sorted(path.removeprefix(prefix) for path in repository_files if path.startswith(prefix))
+ if not paths:
+ raise FileNotFoundError(f"Missing tokenizer artifact directory {prefix.rstrip('/')}")
+ return paths
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument("--tokenizer-path", type=Path, required=True)
+ parser.add_argument("--repo-id", action="append", dest="repo_ids")
+ parser.add_argument("--revision")
+ parser.add_argument("--dry-run", action="store_true")
+ parser.add_argument("--audit-only", action="store_true")
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ tokenizer_path = args.tokenizer_path.resolve()
+ if not tokenizer_path.is_dir():
+ raise FileNotFoundError(f"Tokenizer directory does not exist: {tokenizer_path}")
+
+ files = tokenizer_files(tokenizer_path)
+ fingerprint = artifact_fingerprint([(relative_path, path.read_bytes()) for relative_path, path in files])
+ api = HfApi()
+ repositories = tuple(args.repo_ids or DEFAULT_REPOSITORIES)
+ print(f"Tokenizer fingerprint: {fingerprint}")
+
+ for repo_id in repositories:
+ if args.audit_only:
+ audit_repository(api, repo_id, fingerprint, args.revision)
+ print(f"AUDIT OK {repo_id}@{args.revision or 'main'}")
+ continue
+
+ operations = prepare_operations(repo_id, tokenizer_path, args.revision)
+ if args.dry_run:
+ print(f"DRY RUN {repo_id}: {len(operations)} files")
+ for operation in operations:
+ print(f" {operation.path_in_repo}")
+ continue
+
+ commit = api.create_commit(
+ repo_id=repo_id,
+ repo_type="model",
+ operations=operations,
+ commit_message="Embed fitted FAST tokenizer for portable PI052 checkpoints",
+ revision=args.revision,
+ )
+ audit_repository(api, repo_id, fingerprint, commit.oid)
+ print(f"BACKFILLED {repo_id}@{commit.oid}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/src/lerobot/optim/optimizers.py b/src/lerobot/optim/optimizers.py
index 0a462e1aa..9d3959525 100644
--- a/src/lerobot/optim/optimizers.py
+++ b/src/lerobot/optim/optimizers.py
@@ -104,6 +104,8 @@ class AdamWConfig(OptimizerConfig):
eps: float = 1e-8
weight_decay: float = 1e-2
grad_clip_norm: float = 10.0
+ foreach: bool | None = None
+ fused: bool | None = None
def build(self, params: OptimizerParams) -> torch.optim.Optimizer:
kwargs = asdict(self)
diff --git a/src/lerobot/policies/__init__.py b/src/lerobot/policies/__init__.py
index a95d23b91..04ebaea75 100644
--- a/src/lerobot/policies/__init__.py
+++ b/src/lerobot/policies/__init__.py
@@ -28,6 +28,7 @@ from .multi_task_dit.configuration_multi_task_dit import MultiTaskDiTConfig as M
from .pi0.configuration_pi0 import PI0Config as PI0Config
from .pi0_fast.configuration_pi0_fast import PI0FastConfig as PI0FastConfig
from .pi05.configuration_pi05 import PI05Config as PI05Config
+from .pi052.configuration_pi052 import PI052Config as PI052Config
from .pretrained import PreTrainedPolicy as PreTrainedPolicy
from .smolvla.configuration_smolvla import SmolVLAConfig as SmolVLAConfig
from .tdmpc.configuration_tdmpc import TDMPCConfig as TDMPCConfig
@@ -56,6 +57,7 @@ __all__ = [
"PI0Config",
"PI0FastConfig",
"PI05Config",
+ "PI052Config",
"SmolVLAConfig",
"TDMPCConfig",
"VLAJEPAConfig",
diff --git a/src/lerobot/policies/factory.py b/src/lerobot/policies/factory.py
index 1848a6ffd..577de892d 100644
--- a/src/lerobot/policies/factory.py
+++ b/src/lerobot/policies/factory.py
@@ -137,6 +137,12 @@ class ProcessorConfigKwargs(TypedDict, total=False):
preprocessor_overrides: dict[str, Any] | None
postprocessor_overrides: dict[str, Any] | None
dataset_stats: dict[str, dict[str, torch.Tensor]] | None
+ # Dataset source used by policies that optionally fit processor artifacts.
+ dataset_repo_id: str | None
+ dataset_root: str | None
+ dataset_revision: str | None
+ dataset_episodes: list[int] | None
+ dataset_exclude_episodes: list[int] | None
dataset_meta: Any | None
@@ -171,6 +177,10 @@ def make_pre_post_processors(
ValueError: If no processor factory exists for the given policy configuration type.
"""
if pretrained_path:
+ # Register the PI052-only stateful tokenizer step before deserializing its pipeline.
+ if policy_cfg.type == "pi052":
+ from .pi052 import processor_pi052 as _processor_pi052 # noqa: F401
+
if isinstance(policy_cfg, GrootConfig):
from .groot.processor_groot import make_groot_pre_post_processors_from_pretrained
@@ -190,12 +200,29 @@ def make_pre_post_processors(
),
)
+ preprocessor_overrides = dict(kwargs.get("preprocessor_overrides") or {})
+ if policy_cfg.type == "pi0_fast" and getattr(policy_cfg, "auto_fit_fast_tokenizer", False):
+ from .pi052.fit_fast_tokenizer import resolve_fast_tokenizer
+
+ fitted_tokenizer = resolve_fast_tokenizer(
+ policy_cfg,
+ kwargs.get("dataset_repo_id"),
+ kwargs.get("dataset_root"),
+ kwargs.get("dataset_stats"),
+ kwargs.get("dataset_revision"),
+ kwargs.get("dataset_episodes"),
+ kwargs.get("dataset_exclude_episodes"),
+ )
+ tokenizer_overrides = dict(preprocessor_overrides.get("action_tokenizer_processor") or {})
+ tokenizer_overrides["action_tokenizer_name"] = fitted_tokenizer
+ preprocessor_overrides["action_tokenizer_processor"] = tokenizer_overrides
+
preprocessor = PolicyProcessorPipeline.from_pretrained(
pretrained_model_name_or_path=pretrained_path,
config_filename=kwargs.get(
"preprocessor_config_filename", f"{POLICY_PREPROCESSOR_DEFAULT_NAME}.json"
),
- overrides=kwargs.get("preprocessor_overrides", {}),
+ overrides=preprocessor_overrides,
to_transition=batch_to_transition,
to_output=transition_to_batch,
revision=pretrained_revision,
@@ -227,6 +254,11 @@ def make_pre_post_processors(
config=policy_cfg,
dataset_stats=kwargs.get("dataset_stats"),
dataset_meta=kwargs.get("dataset_meta"),
+ dataset_repo_id=kwargs.get("dataset_repo_id"),
+ dataset_root=kwargs.get("dataset_root"),
+ dataset_revision=kwargs.get("dataset_revision"),
+ episodes=kwargs.get("dataset_episodes"),
+ exclude_episodes=kwargs.get("dataset_exclude_episodes"),
)
@@ -424,6 +456,7 @@ def _make_processors_from_policy_config(
config: PreTrainedConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
dataset_meta: Any | None = None,
+ **optional_kwargs: Any,
) -> tuple[Any, Any]:
"""Create pre- and post-processors from a policy configuration using dynamic imports.
@@ -459,7 +492,9 @@ def _make_processors_from_policy_config(
function = getattr(module, function_name, None)
if function is None:
raise ValueError(f"Processor for policy type '{policy_type}' is not implemented.")
+ parameters = inspect.signature(function).parameters
call_kwargs: dict[str, Any] = {"dataset_stats": dataset_stats}
- if "dataset_meta" in inspect.signature(function).parameters:
+ if "dataset_meta" in parameters:
call_kwargs["dataset_meta"] = dataset_meta
+ call_kwargs.update({name: value for name, value in optional_kwargs.items() if name in parameters})
return function(config, **call_kwargs)
diff --git a/src/lerobot/policies/pi05/modeling_pi05.py b/src/lerobot/policies/pi05/modeling_pi05.py
index d45f5a5c2..b232c3906 100644
--- a/src/lerobot/policies/pi05/modeling_pi05.py
+++ b/src/lerobot/policies/pi05/modeling_pi05.py
@@ -15,21 +15,26 @@
# limitations under the License.
import builtins
+import json
import logging
+import math
from collections import deque
from pathlib import Path
from typing import TYPE_CHECKING, Literal, TypedDict, Unpack
import torch
import torch.nn.functional as F # noqa: N812
+from safetensors.torch import load_file
from torch import Tensor, nn
from lerobot.utils.import_utils import _transformers_available, require_package
# Conditional import for type checking and lazy loading
if TYPE_CHECKING or _transformers_available:
+ from transformers.cache_utils import DynamicCache
from transformers.models.auto import CONFIG_MAPPING
from transformers.models.gemma import modeling_gemma
+ from transformers.utils import cached_file
from ..pi_gemma import (
PaliGemmaForConditionalGenerationWithPiGemma,
@@ -39,27 +44,22 @@ if TYPE_CHECKING or _transformers_available:
)
else:
CONFIG_MAPPING = None
+ DynamicCache = None
modeling_gemma = None
PiGemmaForCausalLM = None
_gated_residual = None
layernorm_forward = None
PaliGemmaForConditionalGenerationWithPiGemma = None
+ cached_file = None
from lerobot.configs import PreTrainedConfig
from lerobot.utils.constants import (
ACTION,
OBS_LANGUAGE_ATTENTION_MASK,
OBS_LANGUAGE_TOKENS,
+ OPENPI_ATTENTION_MASK_VALUE,
)
-from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
-from ..common.vla_utils import (
- clone_past_key_values,
- create_sinusoidal_pos_embedding,
- make_att_2d_masks,
- pad_vector,
- prepare_attention_masks_4d,
- resize_with_pad_torch,
-)
+from ..common.flow_matching import sample_noise, sample_time_beta
from ..pretrained import PreTrainedPolicy, T
from ..rtc.modeling_rtc import RTCProcessor
from .configuration_pi05 import DEFAULT_IMAGE_SIZE, PI05Config
@@ -71,6 +71,251 @@ class ActionSelectKwargs(TypedDict, total=False):
execution_horizon: int | None
+_SAFETENSORS_FILE = "model.safetensors"
+_SAFETENSORS_INDEX = "model.safetensors.index.json"
+
+
+def _resolve_weight_files(
+ pretrained_name_or_path: str | Path,
+ *,
+ force_download: bool,
+ resume_download: bool | None,
+ proxies: dict | None,
+ token: str | bool | None,
+ cache_dir: str | Path | None,
+ local_files_only: bool,
+ revision: str | None,
+) -> list[Path]:
+ model_id = str(pretrained_name_or_path)
+ local_dir = Path(model_id)
+ load_kwargs = {
+ "revision": revision,
+ "cache_dir": cache_dir,
+ "force_download": force_download,
+ "resume_download": resume_download,
+ "proxies": proxies,
+ "token": token,
+ "local_files_only": local_files_only,
+ }
+
+ if local_dir.is_dir():
+ index_path = local_dir / _SAFETENSORS_INDEX
+ single_path = local_dir / _SAFETENSORS_FILE
+ else:
+ resolved_index = cached_file(
+ model_id,
+ _SAFETENSORS_INDEX,
+ _raise_exceptions_for_missing_entries=False,
+ **load_kwargs,
+ )
+ index_path = Path(resolved_index) if resolved_index is not None else None
+ single_path = None
+ if index_path is None:
+ resolved_file = cached_file(model_id, _SAFETENSORS_FILE, **load_kwargs)
+ single_path = Path(resolved_file) if resolved_file is not None else None
+
+ if index_path is None or not index_path.is_file():
+ if single_path is None or not single_path.is_file():
+ raise FileNotFoundError(f"No {_SAFETENSORS_FILE} found in {model_id!r}.")
+ return [single_path]
+
+ index = json.loads(index_path.read_text())
+ shard_names = sorted(set(index.get("weight_map", {}).values()))
+ if not shard_names:
+ raise ValueError(f"Invalid safetensors index without a weight_map: {index_path}")
+ if local_dir.is_dir():
+ files = [local_dir / name for name in shard_names]
+ else:
+ files = []
+ for name in shard_names:
+ resolved_file = cached_file(model_id, name, **load_kwargs)
+ if resolved_file is None:
+ raise FileNotFoundError(f"Checkpoint shard {name!r} not found in {model_id!r}.")
+ files.append(Path(resolved_file))
+ missing = [str(path) for path in files if not path.is_file()]
+ if missing:
+ raise FileNotFoundError(f"Missing checkpoint shards: {missing}")
+ return files
+
+
+def _load_weight_files(files: list[Path]) -> dict[str, Tensor]:
+ state_dict: dict[str, Tensor] = {}
+ for path in files:
+ shard = load_file(path)
+ overlap = state_dict.keys() & shard.keys()
+ if overlap:
+ raise ValueError(f"Duplicate checkpoint keys in {path}: {sorted(overlap)[:5]}")
+ state_dict.update(shard)
+ return state_dict
+
+
+def get_safe_dtype(target_dtype, device_type):
+ """Get a safe dtype for the given device type."""
+ if device_type == "mps" and target_dtype == torch.float64:
+ return torch.float32
+ if device_type == "cpu":
+ # CPU doesn't support bfloat16, use float32 instead
+ if target_dtype == torch.bfloat16:
+ return torch.float32
+ if target_dtype == torch.float64:
+ return torch.float64
+ return target_dtype
+
+
+def create_sinusoidal_pos_embedding( # see openpi `create_sinusoidal_pos_embedding` (exact copy)
+ time: torch.Tensor, dimension: int, min_period: float, max_period: float, device="cpu"
+) -> Tensor:
+ """Computes sine-cosine positional embedding vectors for scalar positions."""
+ if dimension % 2 != 0:
+ raise ValueError(f"dimension ({dimension}) must be divisible by 2")
+
+ if time.ndim != 1:
+ raise ValueError("The time tensor is expected to be of shape `(batch_size, )`.")
+
+ dtype = get_safe_dtype(torch.float64, device.type)
+ fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
+ period = min_period * (max_period / min_period) ** fraction
+
+ # Compute the outer product
+ scaling_factor = 1.0 / period * 2 * math.pi
+ sin_input = scaling_factor[None, :] * time[:, None]
+ return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
+
+
+def sample_beta(alpha, beta, bsize, device): # see openpi `sample_beta` (exact copy)
+ # Beta sampling uses _sample_dirichlet which isn't implemented for MPS, so sample on CPU
+ alpha_t = torch.tensor(alpha, dtype=torch.float32)
+ beta_t = torch.tensor(beta, dtype=torch.float32)
+ dist = torch.distributions.Beta(alpha_t, beta_t)
+ return dist.sample((bsize,)).to(device)
+
+
+def make_att_2d_masks(pad_masks, att_masks): # see openpi `make_att_2d_masks` (exact copy)
+ """Copied from big_vision.
+
+ Tokens can attend to valid inputs tokens which have a cumulative mask_ar
+ smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
+ setup several types of attention, for example:
+
+ [[1 1 1 1 1 1]]: pure causal attention.
+
+ [[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
+ themselves and the last 3 tokens have a causal attention. The first
+ entry could also be a 1 without changing behaviour.
+
+ [[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
+ block can attend all previous blocks and all tokens on the same block.
+
+ Args:
+ input_mask: bool[B, N] true if its part of the input, false if padding.
+ mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
+ it and 0 where it shares the same attention mask as the previous token.
+ """
+ if att_masks.ndim != 2:
+ raise ValueError(att_masks.ndim)
+ if pad_masks.ndim != 2:
+ raise ValueError(pad_masks.ndim)
+
+ cumsum = torch.cumsum(att_masks, dim=1)
+ att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
+ pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
+ return att_2d_masks & pad_2d_masks
+
+
+def clone_past_key_values(past_key_values):
+ """Clone the DynamicCache returned by prefix prefill for compiled denoising."""
+ return DynamicCache(
+ tuple(
+ (keys.clone(), values.clone(), sliding_window) for keys, values, sliding_window in past_key_values
+ )
+ )
+
+
+def pad_vector(vector, new_dim):
+ """Pad the last dimension of a vector to new_dim with zeros.
+
+ Can be (batch_size x sequence_length x features_dimension)
+ or (batch_size x features_dimension)
+ """
+ if vector.shape[-1] >= new_dim:
+ return vector
+ return F.pad(vector, (0, new_dim - vector.shape[-1]))
+
+
+def resize_with_pad_torch( # see openpi `resize_with_pad_torch` (exact copy)
+ images: torch.Tensor,
+ height: int,
+ width: int,
+ mode: str = "bilinear",
+) -> torch.Tensor:
+ """PyTorch version of resize_with_pad. Resizes an image to a target height and width without distortion
+ by padding with black. If the image is float32, it must be in the range [-1, 1].
+
+ Args:
+ images: Tensor of shape [*b, h, w, c] or [*b, c, h, w]
+ height: Target height
+ width: Target width
+ mode: Interpolation mode ('bilinear', 'nearest', etc.)
+
+ Returns:
+ Resized and padded tensor with same shape format as input
+ """
+ # Check if input is in channels-last format [*b, h, w, c] or channels-first [*b, c, h, w]
+ if images.shape[-1] <= 4: # Assume channels-last format
+ channels_last = True
+ if images.dim() == 3:
+ images = images.unsqueeze(0) # Add batch dimension
+ images = images.permute(0, 3, 1, 2) # [b, h, w, c] -> [b, c, h, w]
+ else:
+ channels_last = False
+ if images.dim() == 3:
+ images = images.unsqueeze(0) # Add batch dimension
+
+ batch_size, channels, cur_height, cur_width = images.shape
+
+ # Calculate resize ratio
+ ratio = max(cur_width / width, cur_height / height)
+ resized_height = int(cur_height / ratio)
+ resized_width = int(cur_width / ratio)
+
+ # Resize
+ resized_images = F.interpolate(
+ images,
+ size=(resized_height, resized_width),
+ mode=mode,
+ align_corners=False if mode == "bilinear" else None,
+ )
+
+ # Handle dtype-specific clipping
+ if images.dtype == torch.uint8:
+ resized_images = torch.round(resized_images).clamp(0, 255).to(torch.uint8)
+ elif images.dtype == torch.float32:
+ resized_images = resized_images.clamp(0.0, 1.0)
+ else:
+ raise ValueError(f"Unsupported image dtype: {images.dtype}")
+
+ # Calculate padding
+ pad_h0, remainder_h = divmod(height - resized_height, 2)
+ pad_h1 = pad_h0 + remainder_h
+ pad_w0, remainder_w = divmod(width - resized_width, 2)
+ pad_w1 = pad_w0 + remainder_w
+
+ # Pad
+ constant_value = 0 if images.dtype == torch.uint8 else 0.0
+ padded_images = F.pad(
+ resized_images,
+ (pad_w0, pad_w1, pad_h0, pad_h1), # left, right, top, bottom
+ mode="constant",
+ value=constant_value,
+ )
+
+ # Convert back to original format if needed
+ if channels_last:
+ padded_images = padded_images.permute(0, 2, 3, 1) # [b, c, h, w] -> [b, h, w, c]
+
+ return padded_images
+
+
# Define the complete layer computation function for gradient checkpointing
def compute_layer_complete(inputs_embeds, attention_mask, position_ids, adarms_cond, layers, rotary_emb):
query_states = []
@@ -401,6 +646,12 @@ class PaliGemmaWithExpertModel(
class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
"""Core PI05 PyTorch model."""
+ use_hf_vision_checkpointing_api = False
+ checkpoint_vision_embeddings = True
+ use_typed_attention_masks = False
+ use_on_device_suffix_mask = False
+ precompute_denoise_times = False
+
def __init__(self, config: PI05Config, rtc_processor: RTCProcessor | None = None):
super().__init__()
self.config = config
@@ -444,7 +695,11 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
"""Enable gradient checkpointing for memory optimization."""
self.gradient_checkpointing_enabled = True
self.paligemma_with_expert.paligemma.model.language_model.gradient_checkpointing = True
- self.paligemma_with_expert.paligemma.model.vision_tower.gradient_checkpointing = True
+ vision_tower = self.paligemma_with_expert.paligemma.model.vision_tower
+ if self.use_hf_vision_checkpointing_api:
+ vision_tower.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
+ else:
+ vision_tower.gradient_checkpointing = True
self.paligemma_with_expert.gemma_expert.model.gradient_checkpointing = True
logging.info("Enabled gradient checkpointing for PI05Pytorch model")
@@ -452,7 +707,11 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
"""Disable gradient checkpointing."""
self.gradient_checkpointing_enabled = False
self.paligemma_with_expert.paligemma.model.language_model.gradient_checkpointing = False
- self.paligemma_with_expert.paligemma.model.vision_tower.gradient_checkpointing = False
+ vision_tower = self.paligemma_with_expert.paligemma.model.vision_tower
+ if self.use_hf_vision_checkpointing_api:
+ vision_tower.gradient_checkpointing_disable()
+ else:
+ vision_tower.gradient_checkpointing = False
self.paligemma_with_expert.gemma_expert.model.gradient_checkpointing = False
logging.info("Disabled gradient checkpointing for PI05Pytorch model")
@@ -467,6 +726,14 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
)
return func(*args, **kwargs)
+ def _prepare_attention_masks_4d(self, att_2d_masks, dtype=None):
+ """Helper method to prepare 4D attention masks for transformer."""
+ att_2d_masks_4d = att_2d_masks[:, None, :, :]
+ result = torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
+ if dtype is not None:
+ result = result.to(dtype=dtype)
+ return result
+
def sample_noise(self, shape, device):
return sample_noise(shape, device)
@@ -488,13 +755,16 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
pad_masks = []
att_masks = []
- # Process images
- for img, img_mask in zip(images, img_masks, strict=True):
+ if self.checkpoint_vision_embeddings:
- def image_embed_func(img):
- return self.paligemma_with_expert.embed_image(img)
+ def embed_image(img):
+ return self._apply_checkpoint(self.paligemma_with_expert.embed_image, img)
- img_emb = self._apply_checkpoint(image_embed_func, img)
+ img_embs = [embed_image(img) for img in images]
+ else:
+ img_embs = [self.paligemma_with_expert.embed_image(img) for img in images]
+
+ for img_emb, img_mask in zip(img_embs, img_masks, strict=True):
bsize, num_img_embs = img_emb.shape[:2]
embs.append(img_emb)
@@ -556,8 +826,15 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
# Set attention masks so that image, language and state inputs do not attend to action tokens
att_masks += [1] + ([0] * (self.config.chunk_size - 1))
- att_masks = torch.tensor(att_masks, dtype=action_emb.dtype, device=action_emb.device)
- att_masks = att_masks[None, :].expand(bsize, len(att_masks))
+
+ if self.use_on_device_suffix_mask:
+ n = len(att_masks)
+ att_masks = torch.zeros(n, dtype=action_emb.dtype, device=action_emb.device)
+ att_masks[0] = 1
+ att_masks = att_masks[None, :].expand(bsize, n)
+ else:
+ att_masks = torch.tensor(att_masks, dtype=action_emb.dtype, device=action_emb.device)
+ att_masks = att_masks[None, :].expand(bsize, len(att_masks))
return action_emb, pad_masks, att_masks, adarms_cond
@@ -583,7 +860,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
position_ids = torch.cumsum(pad_masks, dim=1) - 1
- att_2d_masks_4d = prepare_attention_masks_4d(att_2d_masks)
+ att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks)
def forward_func(prefix_embs, suffix_embs, att_2d_masks_4d, position_ids, adarms_cond):
(_, suffix_out), _ = self.paligemma_with_expert.forward(
@@ -641,7 +918,8 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
prefix_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks)
prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1
- prefix_att_2d_masks_4d = prepare_attention_masks_4d(prefix_att_2d_masks)
+ mask_dtype = prefix_embs.dtype if self.use_typed_attention_masks else None
+ prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(prefix_att_2d_masks, dtype=mask_dtype)
self.paligemma_with_expert.paligemma.model.language_model.config._attn_implementation = "eager" # noqa: SLF001
_, past_key_values = self.paligemma_with_expert.forward(
@@ -652,21 +930,52 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
use_cache=True,
)
- return euler_integrate(
- lambda input_x_t, current_timestep: self.denoise_step(
- prefix_pad_masks=prefix_pad_masks,
- past_key_values=past_key_values,
- x_t=input_x_t,
- timestep=current_timestep,
- ),
- noise,
- num_steps,
- rtc_processor=self.rtc_processor,
- rtc_enabled=self._rtc_enabled(),
- inference_delay=kwargs.get("inference_delay"),
- prev_chunk_left_over=kwargs.get("prev_chunk_left_over"),
- execution_horizon=kwargs.get("execution_horizon"),
- )
+ dt = -1.0 / num_steps
+
+ times = None
+ if self.precompute_denoise_times:
+ times = torch.tensor(
+ [1.0 + step * dt for step in range(num_steps)], dtype=torch.float32, device=device
+ )
+
+ x_t = noise
+ for step in range(num_steps):
+ time = 1.0 + step * dt
+ if times is None:
+ time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
+ else:
+ time_tensor = times[step].expand(bsize)
+
+ def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
+ return self.denoise_step(
+ prefix_pad_masks=prefix_pad_masks,
+ past_key_values=past_key_values,
+ x_t=input_x_t,
+ timestep=current_timestep,
+ )
+
+ if self._rtc_enabled():
+ inference_delay = kwargs.get("inference_delay")
+ prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
+ execution_horizon = kwargs.get("execution_horizon")
+
+ v_t = self.rtc_processor.denoise_step(
+ x_t=x_t,
+ prev_chunk_left_over=prev_chunk_left_over,
+ inference_delay=inference_delay,
+ time=time,
+ original_denoise_step_partial=denoise_step_partial_call,
+ execution_horizon=execution_horizon,
+ )
+ else:
+ v_t = denoise_step_partial_call(x_t)
+
+ x_t = x_t + dt * v_t
+
+ if self.rtc_processor is not None and self.rtc_processor.is_debug_enabled():
+ self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
+
+ return x_t
def denoise_step(
self,
@@ -689,7 +998,7 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
- full_att_2d_masks_4d = prepare_attention_masks_4d(full_att_2d_masks)
+ full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks)
self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001
past_key_values = clone_past_key_values(past_key_values)
@@ -713,6 +1022,9 @@ class PI05Policy(PreTrainedPolicy):
config_class = PI05Config
name = "pi05"
+ model_class = PI05Pytorch
+ eval_after_pretrained_load = False
+ show_openpi_disclaimer = True
def __init__(
self,
@@ -730,7 +1042,7 @@ class PI05Policy(PreTrainedPolicy):
# Initialize the core PI05 model
self.init_rtc_processor()
- self.model = PI05Pytorch(config, rtc_processor=self.rtc_processor)
+ self.model = self.model_class(config, rtc_processor=self.rtc_processor)
# Enable gradient checkpointing if requested
if config.gradient_checkpointing:
@@ -756,16 +1068,16 @@ class PI05Policy(PreTrainedPolicy):
strict: bool = True,
**kwargs,
) -> T:
- """Override the from_pretrained method to handle key remapping and display important disclaimer."""
- print(
- "The PI05 model is a direct port of the OpenPI implementation. \n"
- "This implementation follows the original OpenPI structure for compatibility. \n"
- "Original implementation: https://github.com/Physical-Intelligence/openpi"
- )
+ """Load PI05-compatible single-file or sharded safetensors checkpoints."""
+ if cls.show_openpi_disclaimer:
+ print(
+ "The PI05 model is a direct port of the OpenPI implementation. \n"
+ "This implementation follows the original OpenPI structure for compatibility. \n"
+ "Original implementation: https://github.com/Physical-Intelligence/openpi"
+ )
if pretrained_name_or_path is None:
raise ValueError("pretrained_name_or_path is required")
- # Use provided config if available, otherwise create default config
if config is None:
config = PreTrainedConfig.from_pretrained(
pretrained_name_or_path=pretrained_name_or_path,
@@ -779,85 +1091,35 @@ class PI05Policy(PreTrainedPolicy):
**kwargs,
)
- # Initialize model without loading weights
- # Check if dataset_stats were provided in kwargs
model = cls(config, **kwargs)
-
- # Load state dict (expects keys with "model." prefix)
- try:
- print(f"Loading model from: {pretrained_name_or_path}")
- try:
- from transformers.utils import cached_file
-
- resolved_file = cached_file(
- pretrained_name_or_path,
- "model.safetensors",
- cache_dir=kwargs.get("cache_dir"),
- force_download=kwargs.get("force_download", False),
- resume_download=kwargs.get("resume_download"),
- proxies=kwargs.get("proxies"),
- token=kwargs.get("token"),
- revision=kwargs.get("revision"),
- local_files_only=kwargs.get("local_files_only", False),
- )
- from safetensors.torch import load_file
-
- original_state_dict = load_file(resolved_file)
- print("✓ Loaded state dict from model.safetensors")
- except Exception as e:
- print(f"Could not load state dict from remote files: {e}")
- print("Returning model without loading pretrained weights")
- return model
-
- # First, fix any key differences (see openpi model.py, _fix_pytorch_state_dict_keys)
- fixed_state_dict = model._fix_pytorch_state_dict_keys(original_state_dict, model.config)
-
- # Then add "model." prefix for all keys that don't already have it
- remapped_state_dict = {}
- remap_count = 0
-
- for key, value in fixed_state_dict.items():
- if not key.startswith("model."):
- new_key = f"model.{key}"
- remapped_state_dict[new_key] = value
- remap_count += 1
- else:
- remapped_state_dict[key] = value
-
- if remap_count > 0:
- print(f"Remapped {remap_count} state dict keys")
-
- # Load the remapped state dict into the model
- missing_keys, unexpected_keys = model.load_state_dict(remapped_state_dict, strict=strict)
-
- if missing_keys:
- print(f"Missing keys when loading state dict: {len(missing_keys)} keys")
- if len(missing_keys) <= 5:
- for key in missing_keys:
- print(f" - {key}")
- else:
- for key in missing_keys[:5]:
- print(f" - {key}")
- print(f" ... and {len(missing_keys) - 5} more")
-
- if unexpected_keys:
- print(f"Unexpected keys when loading state dict: {len(unexpected_keys)} keys")
- if len(unexpected_keys) <= 5:
- for key in unexpected_keys:
- print(f" - {key}")
- else:
- for key in unexpected_keys[:5]:
- print(f" - {key}")
- print(f" ... and {len(unexpected_keys) - 5} more")
-
- if not missing_keys and not unexpected_keys:
- print("All keys loaded successfully!")
-
- except Exception as e:
- print(f"Warning: Could not load state dict: {e}")
-
+ files = _resolve_weight_files(
+ pretrained_name_or_path,
+ force_download=force_download,
+ resume_download=resume_download,
+ proxies=proxies,
+ token=token,
+ cache_dir=cache_dir,
+ local_files_only=local_files_only,
+ revision=revision,
+ )
+ fixed_state_dict = model._fix_pytorch_state_dict_keys(_load_weight_files(files), model.config)
+ remapped_state_dict = {
+ key if key.startswith("model.") else f"model.{key}": value
+ for key, value in fixed_state_dict.items()
+ }
+ remapped_state_dict = model._prepare_pretrained_state_dict(remapped_state_dict)
+ missing_keys, unexpected_keys = model.load_state_dict(remapped_state_dict, strict=strict)
+ if missing_keys:
+ logging.warning("Missing %s checkpoint keys: %s", cls.name, missing_keys)
+ if unexpected_keys:
+ logging.warning("Unexpected %s checkpoint keys: %s", cls.name, unexpected_keys)
+ if model.eval_after_pretrained_load:
+ model.eval()
return model
+ def _prepare_pretrained_state_dict(self, state_dict: dict[str, Tensor]) -> dict[str, Tensor]:
+ return state_dict
+
def _fix_pytorch_state_dict_keys(
self, state_dict, model_config
): # see openpi `BaseModelConfig, _fix_pytorch_state_dict_keys`
@@ -1028,12 +1290,16 @@ class PI05Policy(PreTrainedPolicy):
# Action queue logic for n_action_steps > 1
if len(self._action_queue) == 0:
- actions = self.predict_action_chunk(batch)[:, : self.config.n_action_steps]
+ action_batch = self._prepare_action_batch(batch)
+ actions = self.predict_action_chunk(action_batch)[:, : self.config.n_action_steps]
# Transpose to get shape (n_action_steps, batch_size, action_dim)
self._action_queue.extend(actions.transpose(0, 1))
return self._action_queue.popleft()
+ def _prepare_action_batch(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
+ return batch
+
@torch.no_grad()
def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
"""Predict a chunk of actions given environment observations."""
diff --git a/src/lerobot/policies/pi052/__init__.py b/src/lerobot/policies/pi052/__init__.py
new file mode 100644
index 000000000..d9857c5c2
--- /dev/null
+++ b/src/lerobot/policies/pi052/__init__.py
@@ -0,0 +1,19 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI052 configuration; model and processors are imported lazily by their factories."""
+
+from .configuration_pi052 import PI052Config
+
+__all__ = ["PI052Config"]
diff --git a/src/lerobot/policies/pi052/configuration_pi052.py b/src/lerobot/policies/pi052/configuration_pi052.py
new file mode 100644
index 000000000..01acd3b72
--- /dev/null
+++ b/src/lerobot/policies/pi052/configuration_pi052.py
@@ -0,0 +1,195 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI0.5 with hierarchical text generation and flow-matched actions."""
+
+from dataclasses import dataclass
+
+from lerobot.configs import PreTrainedConfig
+from lerobot.optim.optimizers import AdamWConfig
+
+from ..pi05.configuration_pi05 import PI05Config
+
+
+@PreTrainedConfig.register_subclass("pi052")
+@dataclass
+class PI052Config(PI05Config):
+ """PI0.5 configuration for recipe-driven text and action supervision."""
+
+ # Recipe / language stack ---------------------------------------------
+ recipe_path: str | None = "recipes/subtask_mem.yaml"
+ """Recipe path relative to ``src/lerobot/configs/``, or ``None`` for the plain PI0.5 prompt."""
+
+ apply_chat_template: bool = False
+ """Whether to apply a tokenizer chat template.
+
+ PaliGemma defaults to plain recipe-rendered prefixes because it is not chat-pretrained.
+ """
+
+ # Balance frequent recipe text supervision against the paper's α=10 flow weight.
+ text_loss_weight: float = 1.0
+ """LM-head cross-entropy weight; ``0`` disables text training."""
+
+ flow_loss_weight: float = 10.0
+ """Weight on action-expert flow matching relative to text supervision."""
+
+ # Backbone training ---------------------------------------------------
+ unfreeze_lm_head: bool = True
+ """Keep PaliGemma's language head trainable for hierarchical inference."""
+
+ # Optional context dropout improves tolerance to missing or stale language state.
+ plan_dropout_prob: float = 0.0
+ memory_dropout_prob: float = 0.0
+ subtask_dropout_prob: float = 0.0
+
+ # FAST adds discrete-action CE to the text and flow objectives from paper §III.B-C.
+ enable_fast_action_loss: bool = True
+ """Add FAST-tokenized action cross-entropy to text CE and flow matching."""
+
+ action_tokenizer_name: str = "physical-intelligence/fast"
+ """HF identifier for the FAST action tokenizer."""
+
+ max_action_tokens: int = 256
+ """Maximum number of FAST tokens per action chunk."""
+
+ fast_skip_tokens: int = 1152
+ """Number of top-of-vocab tokens the FAST id mapping skips.
+
+ 1152 skips PaliGemma's 128 ```` and 1024 ```` special tokens so
+ FAST codes land in plain-text ids below 256000 and never collide with the
+ ```` targets used for VQA. openpi's pi0-FAST convention is 128 (FAST
+ occupies the ```` range); use 128 only to stay weight-compatible with
+ checkpoints trained that way."""
+
+ fast_action_loss_weight: float = 1.0
+ """Weight on FAST action-token CE relative to continuous-flow supervision."""
+
+ subtask_replan_steps: int = 0
+ """Environment steps between subtask generations during evaluation.
+
+ Non-positive values regenerate each action chunk while still refreshing the action prompt every chunk.
+ """
+
+ joint_subtask_conditioning: bool = False
+ """Condition low-level action inference on the task plus the generated subtask.
+
+ Matches paper-style joint-sequence recipes (``recipes/subtask_joint.yaml``)
+ where one sample supervises the subtask text and conditions the action
+ losses on it: the inference prefix becomes
+ ``User: {task}, State: ...;\\nAssistant: {subtask}`` with the subtask
+ span attended causally, exactly as trained. Leave ``False`` for the blend
+ recipes, whose low-level samples use ``User: {subtask}, State: ...;``."""
+
+ auto_fit_fast_tokenizer: bool = False
+ """Fit and cache a dataset-specific FAST tokenizer before training.
+
+ Disabled by default to avoid the extra dataset pass and use the universal tokenizer.
+ """
+
+ fast_tokenizer_cache_dir: str = "~/.cache/lerobot/fast_tokenizers"
+ """Where fitted FAST tokenizers are stored. ``~`` expands."""
+
+ fast_tokenizer_fit_samples: int = 1024
+ """Number of action chunks sampled when fitting FAST."""
+
+ fast_tokenizer_validation_samples: int = 256
+ """Held-out action chunks used to validate tokenizer reconstruction."""
+
+ fast_tokenizer_max_reconstruction_rmse: float = 0.10
+ """Maximum normalized RMSE allowed across held-out action chunks."""
+
+ fast_tokenizer_max_dim_rmse: float = 0.20
+ """Maximum normalized RMSE allowed for any nonconstant action dimension."""
+
+ # Knowledge insulation detaches VLM K/V from action-loss gradients (paper §III.B).
+ knowledge_insulation: bool = True
+ """Block action-loss gradients through VLM keys and values."""
+
+ # Optional training backends. Defaults preserve the eager/SDPA path.
+ use_flashrt_adarms: bool = False
+ """Use FlashRT adaptive RMSNorm kernels when available."""
+
+ use_compiled_text_ce: bool = False
+ """Compile the materialized-logits text and FAST CE path."""
+
+ use_compiled_vision: bool = False
+ """Compile the SigLIP tower for no-grad flow and inference passes."""
+
+ use_flex_attention: bool = False
+ """Use FlexAttention for amortized KI, with SDPA fallback where unsupported."""
+
+ use_manual_attention: bool = False
+ """Use materialized-logits attention for explicitly profiled KI shapes."""
+
+ manual_attention_scope: str = "all"
+ """Apply manual attention to all KI queries or only action queries."""
+
+ # Scale language-head updates relative to the base optimizer schedule.
+ lm_head_lr_scale: float = 1.0
+
+ # Scale backbone and action-expert optimizer groups independently.
+ backbone_lr_scale: float = 1.0
+ action_expert_lr_scale: float = 1.0
+
+ # Reuse each VLM prefix across independent denoising draws; 1 restores single-draw flow.
+ flow_num_repeats: int = 5
+
+ # PaLM-style z-loss stabilizes large-vocabulary CE; 0 disables it.
+ text_ce_z_loss_weight: float = 1e-4
+
+ use_flashrt_fp8_mlp: bool = False
+ """Enable calibrated FlashRT FP8 kernels for Gemma and SigLIP MLPs.
+
+ Apply after loading with ``PI052Policy.apply_flashrt_fp8_mlp``; unavailable kernels keep BF16.
+ """
+
+ # Keep serialized PI052 AdamW options local because PI05Config lacks them.
+ optimizer_foreach: bool | None = False
+ optimizer_fused: bool | None = True
+
+ def get_optimizer_preset(self) -> AdamWConfig:
+ return AdamWConfig(
+ lr=self.optimizer_lr,
+ betas=self.optimizer_betas,
+ eps=self.optimizer_eps,
+ weight_decay=self.optimizer_weight_decay,
+ grad_clip_norm=self.optimizer_grad_clip_norm,
+ foreach=self.optimizer_foreach,
+ fused=self.optimizer_fused,
+ )
+
+ def __post_init__(self) -> None:
+ super().__post_init__()
+ if self.enable_fast_action_loss and not self.recipe_path:
+ raise ValueError("PI052 FAST action loss requires recipe_path to build action supervision.")
+ if self.text_loss_weight > 0 and self.unfreeze_lm_head:
+ self.train_expert_only = False
+ if self.flow_num_repeats < 1:
+ raise ValueError(f"flow_num_repeats must be >= 1, got {self.flow_num_repeats}")
+ if self.fast_tokenizer_validation_samples < 1:
+ 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 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}"
+ )
+ if self.use_flex_attention and self.use_manual_attention:
+ raise ValueError("use_flex_attention and use_manual_attention are mutually exclusive")
+ if self.use_flex_attention and self.flow_num_repeats == 1:
+ raise ValueError("use_flex_attention requires flow_num_repeats > 1")
+ if not self.knowledge_insulation and (
+ self.use_flex_attention or self.use_manual_attention or self.use_flashrt_adarms
+ ):
+ raise ValueError("KI attention and AdaRMS optimizations require knowledge_insulation=True")
diff --git a/src/lerobot/policies/pi052/fit_fast_tokenizer.py b/src/lerobot/policies/pi052/fit_fast_tokenizer.py
new file mode 100644
index 000000000..e0dd36e12
--- /dev/null
+++ b/src/lerobot/policies/pi052/fit_fast_tokenizer.py
@@ -0,0 +1,522 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Fit and cache a FAST tokenizer for a dataset's action distribution.
+
+Training invokes this automatically when FAST loss and automatic fitting are enabled.
+"""
+
+from __future__ import annotations
+
+import hashlib
+import json
+import logging
+import os
+import shutil
+import time
+from pathlib import Path
+from typing import Any
+
+import numpy as np
+
+logger = logging.getLogger(__name__)
+
+# ``ProcessorMixin.save_pretrained`` writes this shared cache sentinel.
+_CACHE_SENTINEL = "processor_config.json"
+
+
+def _is_global_leader() -> bool:
+ return int(os.environ.get("RANK", "0")) == 0
+
+
+def _jsonable(value: Any) -> Any:
+ if hasattr(value, "detach"):
+ value = value.detach().cpu().numpy()
+ if isinstance(value, np.ndarray):
+ return value.tolist()
+ if isinstance(value, dict):
+ return {key: _jsonable(item) for key, item in sorted(value.items())}
+ if isinstance(value, (list, tuple)):
+ return [_jsonable(item) for item in value]
+ return value
+
+
+def _dataset_signature(
+ dataset_repo_id: str,
+ base_tokenizer_name: str,
+ n_samples: int,
+ chunk_size: int,
+ normalization_mode: str,
+ dataset_revision: str | None = None,
+ episodes: list[int] | None = None,
+ exclude_episodes: list[int] | None = None,
+ action_stats: dict | None = None,
+ use_relative_actions: bool = False,
+ relative_action_mask: list[bool] | None = None,
+ validation_samples: int = 256,
+ max_reconstruction_rmse: float = 0.10,
+ max_dim_rmse: float = 0.20,
+) -> str:
+ """Hash every input that changes the fitted action distribution."""
+ payload = {
+ "dataset_repo_id": dataset_repo_id,
+ "dataset_revision": dataset_revision,
+ "base_tokenizer_name": base_tokenizer_name,
+ "n_samples": n_samples,
+ "chunk_size": chunk_size,
+ "normalization_mode": normalization_mode,
+ "episodes": episodes,
+ "exclude_episodes": exclude_episodes,
+ "action_stats": action_stats,
+ "use_relative_actions": use_relative_actions,
+ "relative_action_mask": relative_action_mask,
+ "validation_samples": validation_samples,
+ "max_reconstruction_rmse": max_reconstruction_rmse,
+ "max_dim_rmse": max_dim_rmse,
+ }
+ encoded = json.dumps(_jsonable(payload), sort_keys=True, separators=(",", ":")).encode()
+ return hashlib.sha256(encoded).hexdigest()[:16]
+
+
+def _select_episode_indices(
+ available_episodes: list[int],
+ episodes: list[int] | None,
+ exclude_episodes: list[int] | None,
+) -> list[int]:
+ allowed = set(episodes) if episodes is not None else set(available_episodes)
+ excluded = set(exclude_episodes or [])
+ return [episode for episode in available_episodes if episode in allowed and episode not in excluded]
+
+
+def _apply_relative_actions(
+ actions: np.ndarray,
+ states: np.ndarray,
+ relative_action_mask: list[bool] | None,
+) -> np.ndarray:
+ """Match RelativeActionsProcessorStep before tokenizer fitting."""
+ action_dim = actions.shape[-1]
+ mask = list(relative_action_mask) if relative_action_mask is not None else [True] * action_dim
+ if len(mask) < action_dim:
+ mask.extend([True] * (action_dim - len(mask)))
+ mask_array = np.asarray(mask[:action_dim], dtype=np.float32)
+ relative = actions.copy()
+ relative -= states[:, None, :action_dim] * mask_array
+ return relative
+
+
+def _normalize_actions(
+ actions: np.ndarray,
+ normalization_mode: str,
+ action_stats: dict | None = None,
+) -> np.ndarray:
+ """Match the action normalization applied by the training preprocessor."""
+ mode = getattr(normalization_mode, "value", normalization_mode).upper()
+ flat = actions.reshape(-1, actions.shape[-1])
+ stats = action_stats or {}
+
+ def stat(name: str, fallback) -> np.ndarray:
+ value = stats.get(name)
+ if value is None:
+ value = fallback()
+ if hasattr(value, "detach"):
+ value = value.detach().cpu().numpy()
+ return np.asarray(value, dtype=np.float32)
+
+ if mode == "IDENTITY":
+ return actions
+ if mode == "MEAN_STD":
+ mean = stat("mean", lambda: flat.mean(axis=0))
+ std = stat("std", lambda: flat.std(axis=0))
+ return ((actions - mean) / np.where(std == 0, 1e-8, std)).astype(np.float32)
+ if mode in {"QUANTILES", "QUANTILE10"}:
+ low_name, high_name, low_q, high_q = (
+ ("q01", "q99", 0.01, 0.99) if mode == "QUANTILES" else ("q10", "q90", 0.10, 0.90)
+ )
+ low = stat(low_name, lambda: np.quantile(flat, low_q, axis=0))
+ high = stat(high_name, lambda: np.quantile(flat, high_q, axis=0))
+ elif mode == "MIN_MAX":
+ low = stat("min", lambda: flat.min(axis=0))
+ high = stat("max", lambda: flat.max(axis=0))
+ else:
+ raise ValueError(f"Unsupported FAST tokenizer normalization mode: {mode}")
+
+ return (2.0 * (actions - low) / np.where(high == low, 1e-8, high - low) - 1.0).astype(np.float32)
+
+
+def _validate_fast_reconstruction(
+ tokenizer: Any,
+ actions: np.ndarray,
+ max_reconstruction_rmse: float,
+ max_dim_rmse: float,
+) -> tuple[dict[str, Any], np.ndarray]:
+ """Decode held-out chunks and reject tokenizers with excessive quantization error."""
+ decoded = np.asarray(tokenizer.decode(tokenizer(actions)), dtype=np.float32)
+ if decoded.shape != actions.shape:
+ raise RuntimeError(
+ f"FAST tokenizer reconstruction shape mismatch: expected {actions.shape}, got {decoded.shape}."
+ )
+ if not np.isfinite(decoded).all():
+ raise RuntimeError("FAST tokenizer reconstruction contains non-finite values.")
+
+ squared_error = np.square(decoded - actions)
+ rmse = float(np.sqrt(squared_error.mean()))
+ dim_rmse = np.sqrt(squared_error.mean(axis=(0, 1)))
+ nonconstant_dims = np.ptp(actions, axis=(0, 1)) > 1e-8
+ max_observed_dim_rmse = float(dim_rmse[nonconstant_dims].max(initial=0.0))
+ report = {
+ "num_validation_chunks": int(actions.shape[0]),
+ "reconstruction_rmse": rmse,
+ "max_dim_rmse": max_observed_dim_rmse,
+ "dim_rmse": dim_rmse.tolist(),
+ "max_reconstruction_rmse": max_reconstruction_rmse,
+ "max_allowed_dim_rmse": max_dim_rmse,
+ }
+ if rmse > max_reconstruction_rmse or max_observed_dim_rmse > max_dim_rmse:
+ raise RuntimeError(
+ "FAST tokenizer reconstruction error exceeds the configured limit: "
+ f"rmse={rmse:.4f} (max {max_reconstruction_rmse:.4f}), "
+ f"max_dim_rmse={max_observed_dim_rmse:.4f} (max {max_dim_rmse:.4f})."
+ )
+ return report, decoded
+
+
+def _load_fast_fitter(base_tokenizer_name: str) -> Any:
+ """Load FAST's fitting implementation without requiring its universal BPE weights."""
+ from transformers import AutoProcessor # noqa: PLC0415
+
+ try:
+ return AutoProcessor.from_pretrained(base_tokenizer_name, trust_remote_code=True)
+ except ValueError as error:
+ if base_tokenizer_name != "physical-intelligence/fast":
+ raise
+ logger.warning(
+ "Could not load the universal FAST tokenizer backend; loading its fitting class directly: %s",
+ error,
+ )
+ from transformers.dynamic_module_utils import get_class_from_dynamic_module # noqa: PLC0415
+
+ return get_class_from_dynamic_module(
+ "processing_action_tokenizer.UniversalActionProcessor",
+ base_tokenizer_name,
+ )
+
+
+def fit_fast_tokenizer(
+ *,
+ dataset_repo_id: str,
+ cache_dir: str | Path,
+ base_tokenizer_name: str = "physical-intelligence/fast",
+ n_samples: int = 1024,
+ chunk_size: int = 50,
+ seed: int = 42,
+ dataset_root: str | Path | None = None,
+ dataset_revision: str | None = None,
+ episodes: list[int] | None = None,
+ exclude_episodes: list[int] | None = None,
+ normalization_mode: str = "QUANTILES",
+ action_stats: dict | None = None,
+ use_relative_actions: bool = False,
+ relative_action_mask: list[bool] | None = None,
+ validation_samples: int = 256,
+ max_reconstruction_rmse: float = 0.10,
+ max_dim_rmse: float = 0.20,
+) -> str:
+ """Fit a FAST tokenizer on a LeRobot dataset's action distribution.
+
+ Args:
+ dataset_repo_id: HF Hub repo id of the LeRobotDataset to fit on.
+ cache_dir: Directory under which to save (and look up) fitted
+ tokenizers. The actual save path is
+ ``{cache_dir}/{signature}``.
+ base_tokenizer_name: HF identifier for the base FAST tokenizer
+ to finetune from. ``physical-intelligence/fast`` is the
+ universal one.
+ n_samples: Number of action chunks to sample for the fit. The
+ FAST paper uses a few thousand; ``1024`` is a good default
+ for medium datasets.
+ chunk_size: Length of each action chunk (matches
+ ``policy.chunk_size``). The FAST tokenizer is fit on
+ sequences of this length.
+ seed: RNG seed for sample selection.
+
+ Returns:
+ The local path to the fitted tokenizer. Passed directly to
+ ``--policy.action_tokenizer_name`` for the training run.
+
+ Raises:
+ ImportError: If the ``transformers`` library doesn't expose
+ ``AutoProcessor`` or the FAST tokenizer doesn't have a
+ ``.fit()`` method (then you're on an older FAST snapshot —
+ update to the current published model).
+ FileNotFoundError: If the dataset can't be loaded.
+ """
+ cache_dir = Path(cache_dir)
+ normalization_mode = getattr(normalization_mode, "value", normalization_mode).upper()
+ sig = _dataset_signature(
+ dataset_repo_id,
+ base_tokenizer_name,
+ n_samples,
+ chunk_size,
+ normalization_mode,
+ dataset_revision,
+ episodes,
+ exclude_episodes,
+ action_stats,
+ use_relative_actions,
+ relative_action_mask,
+ validation_samples,
+ max_reconstruction_rmse,
+ max_dim_rmse,
+ )
+ out_dir = cache_dir / sig
+
+ if out_dir.exists() and (out_dir / _CACHE_SENTINEL).exists():
+ logger.info(
+ "FAST tokenizer cache hit: %s — re-using fitted tokenizer for dataset=%s base=%s n_samples=%d",
+ out_dir,
+ dataset_repo_id,
+ base_tokenizer_name,
+ n_samples,
+ )
+ return str(out_dir)
+
+ # One global rank populates the shared cache; every other rank waits for the atomic publish.
+ is_leader = _is_global_leader()
+ if not is_leader:
+ timeout_s = 1800.0 # 30 min — covers ~1024-sample fits on cold caches
+ start = time.monotonic()
+ while not (out_dir / _CACHE_SENTINEL).exists():
+ if time.monotonic() - start > timeout_s:
+ raise RuntimeError(
+ f"FAST tokenizer fit: non-leader rank timed out after "
+ f"{timeout_s:.0f}s waiting for {out_dir / _CACHE_SENTINEL}. "
+ "Leader rank likely crashed during the fit."
+ )
+ time.sleep(2.0)
+ logger.info("FAST tokenizer ready (leader populated cache): %s", out_dir)
+ return str(out_dir)
+
+ logger.info(
+ "FAST tokenizer cache miss — fitting on dataset=%s base=%s n_samples=%d chunk_size=%d → %s",
+ dataset_repo_id,
+ base_tokenizer_name,
+ n_samples,
+ chunk_size,
+ out_dir,
+ )
+
+ # Read action columns directly to avoid video decoding and bound memory to sampled episodes.
+ rng = np.random.default_rng(seed)
+ actions_buf: list[np.ndarray] = []
+
+ # Read v3 parquet shards directly to avoid split lookup failures and repeated metadata parsing.
+ import pyarrow as _pa # noqa: PLC0415
+ import pyarrow.parquet as _pq # noqa: PLC0415
+
+ if dataset_root is not None:
+ snap = Path(dataset_root)
+ else:
+ from huggingface_hub import snapshot_download # noqa: PLC0415
+
+ snap = Path(
+ snapshot_download(repo_id=dataset_repo_id, repo_type="dataset", revision=dataset_revision)
+ )
+ data_files = sorted((snap / "data").glob("chunk-*/file-*.parquet"))
+ if not data_files:
+ raise RuntimeError(f"FAST fit: no ``data/chunk-*/file-*.parquet`` shards found under {snap!s}.")
+
+ columns = ["episode_index", "action"]
+ if use_relative_actions:
+ columns.append("observation.state")
+ tables = [_pq.read_table(f, columns=columns) for f in data_files]
+ table = _pa.concat_tables(tables)
+ eps = table["episode_index"].to_numpy()
+ acts_col = table["action"]
+ # Normalize Arrow action representations into an (N, D) array.
+ try:
+ acts = np.stack(acts_col.to_numpy(zero_copy_only=False)).astype(np.float32)
+ except Exception: # noqa: BLE001
+ # Fallback path for nested-list types: flatten via to_pylist().
+ acts = np.asarray(acts_col.to_pylist(), dtype=np.float32)
+ if acts.ndim != 2:
+ raise RuntimeError(f"FAST fit: expected ``action`` rows to be 1-D vectors; got shape {acts.shape}.")
+ states = None
+ if use_relative_actions:
+ try:
+ states = np.stack(table["observation.state"].to_numpy(zero_copy_only=False)).astype(np.float32)
+ except Exception: # noqa: BLE001
+ states = np.asarray(table["observation.state"].to_pylist(), dtype=np.float32)
+ if states.ndim != 2:
+ raise RuntimeError(
+ f"FAST fit: expected ``observation.state`` rows to be 1-D vectors; got {states.shape}."
+ )
+
+ # Sort once because episode order is only guaranteed within each shard.
+ order = np.argsort(eps, kind="stable")
+ eps_sorted = eps[order]
+ boundaries = np.searchsorted(eps_sorted, np.arange(int(eps_sorted.max()) + 2))
+ ep_to_slice: dict[int, tuple[int, int]] = {
+ int(ep): (int(boundaries[ep]), int(boundaries[ep + 1]))
+ for ep in range(len(boundaries) - 1)
+ if boundaries[ep] < boundaries[ep + 1]
+ }
+ num_episodes = len(ep_to_slice)
+ # ``acts`` is in original (un-sorted-by-episode) row order; reorder
+ # so per-episode slices are contiguous.
+ acts = acts[order]
+ if states is not None:
+ states = states[order]
+
+ ep_indices = _select_episode_indices(list(ep_to_slice), episodes, exclude_episodes)
+ if not ep_indices:
+ raise RuntimeError("FAST fit: episode selection is empty after applying exclusions.")
+ total_samples = n_samples + validation_samples
+ samples_per_episode = max(1, (total_samples + len(ep_indices) - 1) // len(ep_indices))
+ collected = 0
+ eps_visited = 0
+ short_episodes = 0
+ states_buf: list[np.ndarray] = []
+ for ep_idx in rng.permutation(ep_indices):
+ if collected >= total_samples:
+ break
+ start, stop = ep_to_slice[int(ep_idx)]
+ ep_actions = acts[start:stop]
+ if ep_actions.shape[0] < chunk_size:
+ short_episodes += 1
+ continue
+ starts = rng.integers(0, ep_actions.shape[0] - chunk_size + 1, size=samples_per_episode)
+ for s in starts:
+ actions_buf.append(ep_actions[int(s) : int(s) + chunk_size])
+ if states is not None:
+ states_buf.append(states[start + int(s)])
+ collected += 1
+ if collected >= total_samples:
+ break
+ eps_visited += 1
+
+ if not actions_buf:
+ raise RuntimeError(
+ f"FAST fit collected zero action chunks from {dataset_repo_id!r}: "
+ f"all {num_episodes} episodes were shorter than chunk_size="
+ f"{chunk_size} ({short_episodes} too short) or had an unreadable "
+ "``action`` column. Lower ``chunk_size`` to match your episode "
+ "lengths."
+ )
+
+ actions = np.stack(actions_buf, axis=0).astype(np.float32) # (N, H, D)
+ if states is not None:
+ actions = _apply_relative_actions(actions, np.stack(states_buf), relative_action_mask)
+ logger.info(
+ "FAST fit: collected %d chunks of shape %s from %d episodes",
+ actions.shape[0],
+ actions.shape[1:],
+ eps_visited,
+ )
+
+ actions = _normalize_actions(actions, normalization_mode, action_stats)
+
+ base = _load_fast_fitter(base_tokenizer_name)
+ if not hasattr(base, "fit"):
+ raise ImportError(
+ f"Base FAST tokenizer {base_tokenizer_name!r} has no ``.fit()`` "
+ "method — your transformers / model snapshot is too old. Update "
+ "to the current ``physical-intelligence/fast`` revision."
+ )
+
+ if actions.shape[0] < total_samples:
+ raise RuntimeError(
+ f"FAST fit collected {actions.shape[0]} chunks, but {total_samples} are required "
+ f"for {n_samples} fit and {validation_samples} validation chunks."
+ )
+ fit_actions = actions[:n_samples]
+ validation_actions = actions[n_samples:total_samples]
+ fitted = base.fit(fit_actions)
+ validation_report, decoded_actions = _validate_fast_reconstruction(
+ fitted,
+ validation_actions,
+ max_reconstruction_rmse,
+ max_dim_rmse,
+ )
+ cache_dir.mkdir(parents=True, exist_ok=True)
+ staging_dir = cache_dir / f".{sig}.tmp-{os.getpid()}"
+ shutil.rmtree(staging_dir, ignore_errors=True)
+ fitted.save_pretrained(str(staging_dir))
+ (staging_dir / "reconstruction_validation.json").write_text(
+ json.dumps(validation_report, indent=2) + "\n"
+ )
+ np.savez_compressed(
+ staging_dir / "reconstruction_examples.npz",
+ original=validation_actions[:8],
+ decoded=decoded_actions[:8],
+ )
+ if out_dir.exists():
+ shutil.rmtree(out_dir)
+ staging_dir.replace(out_dir)
+ logger.info("FAST fit: saved fitted tokenizer to %s", out_dir)
+ return str(out_dir)
+
+
+def resolve_fast_tokenizer(
+ config: Any,
+ dataset_repo_id: str | None,
+ dataset_root: str | Path | None = None,
+ dataset_stats: dict | None = None,
+ dataset_revision: str | None = None,
+ episodes: list[int] | None = None,
+ exclude_episodes: list[int] | None = None,
+) -> str:
+ """Return the configured tokenizer, fitting a cached dataset-specific one when requested."""
+ if not getattr(config, "auto_fit_fast_tokenizer", False) or dataset_repo_id is None:
+ return config.action_tokenizer_name
+
+ relative_action_mask = None
+ if getattr(config, "use_relative_actions", False):
+ action_names = getattr(config, "action_feature_names", None)
+ exclude_tokens = [
+ str(name).lower() for name in getattr(config, "relative_exclude_joints", []) if name
+ ]
+ if action_names is not None and exclude_tokens:
+ relative_action_mask = [
+ not any(token == str(name).lower() or token in str(name).lower() for token in exclude_tokens)
+ for name in action_names
+ ]
+
+ fit_kwargs = {
+ "dataset_repo_id": dataset_repo_id,
+ "cache_dir": Path(config.fast_tokenizer_cache_dir).expanduser(),
+ "base_tokenizer_name": config.action_tokenizer_name,
+ "n_samples": config.fast_tokenizer_fit_samples,
+ "chunk_size": config.chunk_size,
+ "dataset_root": dataset_root,
+ "dataset_revision": dataset_revision,
+ "episodes": episodes,
+ "exclude_episodes": exclude_episodes,
+ "normalization_mode": config.normalization_mapping.get("ACTION", "QUANTILES"),
+ "action_stats": (dataset_stats or {}).get("action"),
+ "use_relative_actions": getattr(config, "use_relative_actions", False),
+ "relative_action_mask": relative_action_mask,
+ }
+ validation_fields = {
+ "validation_samples": "fast_tokenizer_validation_samples",
+ "max_reconstruction_rmse": "fast_tokenizer_max_reconstruction_rmse",
+ "max_dim_rmse": "fast_tokenizer_max_dim_rmse",
+ }
+ fit_kwargs.update(
+ {
+ argument: getattr(config, attribute)
+ for argument, attribute in validation_fields.items()
+ if hasattr(config, attribute)
+ }
+ )
+ return fit_fast_tokenizer(**fit_kwargs)
diff --git a/src/lerobot/policies/pi052/flashrt_fp8.py b/src/lerobot/policies/pi052/flashrt_fp8.py
new file mode 100644
index 000000000..dc0ac3515
--- /dev/null
+++ b/src/lerobot/policies/pi052/flashrt_fp8.py
@@ -0,0 +1,263 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Optional FlashRT FP8 MLP kernels with one-pass calibration and BF16 fallback."""
+
+from __future__ import annotations
+
+import logging
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F # noqa: N812
+
+logger = logging.getLogger(__name__)
+
+_FP8_MAX = 448.0
+
+
+def _roundtrip_fp8(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
+ """Quantize->dequantize an activation through FP8 E4M3 at ``scale`` (f32)."""
+ q = torch.clamp(x.float() / scale.float(), -_FP8_MAX, _FP8_MAX).to(torch.float8_e4m3fn)
+ return q.float() * scale.float()
+
+
+_SWIGLU_REPO = "flashrt/flashrt-fp8-swiglu-ffn"
+_GELU_REPO = "flashrt/flashrt-fp8-ffn"
+_GEMM_REPO = "flashrt/flashrt-gemm-epilogues"
+
+
+def _get_kernel(repo: str):
+ """Load a cached FlashRT Hub package."""
+ from kernels import get_kernel
+
+ return get_kernel(repo, version=1)
+
+
+def _quantize_fp8(weight: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ scale = max(weight.detach().float().abs().max().item(), 1e-12) / _FP8_MAX
+ fp8 = torch.clamp(weight.float() / scale, -_FP8_MAX, _FP8_MAX).to(torch.float8_e4m3fn)
+ return fp8.contiguous(), torch.tensor([scale], dtype=torch.float32)
+
+
+def _static_scale(amax: float, safety: float) -> torch.Tensor:
+ return torch.tensor([max(amax, 1e-12) / _FP8_MAX * safety], dtype=torch.float32)
+
+
+class _FlashRTGeGLU(nn.Module):
+ """FP8 Gemma GeGLU MLP."""
+
+ def __init__(self, mlp, in_amax, hid_amax, ffn_ops, quant_ops, safety, fuse_weight=None):
+ super().__init__()
+ self.ffn_ops = ffn_ops
+ self.quant_ops = quant_ops
+ self.in_features = mlp.gate_proj.weight.shape[1]
+ device = mlp.gate_proj.weight.device
+ gate_up = torch.cat([mlp.gate_proj.weight, mlp.up_proj.weight], dim=0).float()
+ # Fold fixed RMSNorm weights into GEMM; adaptive norms use identity scaling.
+ if fuse_weight is not None:
+ f = 1.0 + fuse_weight.detach().float()
+ gate_up = gate_up * f[None, :]
+ channel_scale = (1.0 / f).to(torch.bfloat16)
+ else:
+ channel_scale = torch.ones(self.in_features, dtype=torch.bfloat16)
+ gate_up_fp8, gate_up_scale = _quantize_fp8(gate_up)
+ down_fp8, down_scale = _quantize_fp8(mlp.down_proj.weight)
+ self.register_buffer("gate_up_fp8", gate_up_fp8.to(device))
+ self.register_buffer("down_fp8", down_fp8.to(device))
+ self.register_buffer("gate_up_scale", gate_up_scale.to(device))
+ self.register_buffer("down_scale", down_scale.to(device))
+ self.register_buffer("input_scale", _static_scale(in_amax, safety).to(device))
+ self.register_buffer("hidden_scale", _static_scale(hid_amax, safety).to(device))
+ self.register_buffer("channel_scale", channel_scale.to(device))
+ self.safety = safety
+ self.calibrating = False
+ self._ia = 0.0
+ self._ha = 0.0
+
+ def _calibrate_step(self, x):
+ # Track input and hidden maxima on live FP8-propagated activations.
+ flat = x.reshape(-1, self.in_features).to(torch.bfloat16)
+ xq = flat.float() * self.channel_scale.float()
+ self._ia = max(self._ia, xq.abs().max().item())
+ self.input_scale.copy_(_static_scale(self._ia, self.safety).to(self.input_scale.device))
+ xdq = _roundtrip_fp8(xq, self.input_scale)
+ wdq = self.gate_up_fp8.float() * self.gate_up_scale.float()
+ gate, up = (xdq @ wdq.t()).chunk(2, dim=-1)
+ hidden = F.gelu(gate, approximate="tanh") * up
+ self._ha = max(self._ha, hidden.abs().max().item())
+ self.hidden_scale.copy_(_static_scale(self._ha, self.safety).to(self.hidden_scale.device))
+
+ def forward(self, x):
+ if self.calibrating:
+ self._calibrate_step(x)
+ shape = x.shape
+ flat = x.reshape(-1, self.in_features).to(torch.bfloat16)
+ x_fp8 = self.quant_ops.channel_scale_quantize_fp8_static_bf16(
+ flat, self.channel_scale, self.input_scale
+ )
+ out = self.ffn_ops.fp8_geglu_mlp_bf16(
+ x_fp8,
+ self.gate_up_fp8,
+ self.down_fp8,
+ self.input_scale,
+ self.gate_up_scale,
+ self.hidden_scale,
+ self.down_scale,
+ )
+ return out.reshape(shape)
+
+
+class _FlashRTGeluMLP(nn.Module):
+ """FP8 SigLIP GELU MLP."""
+
+ def __init__(self, mlp, in_amax, hid_amax, ffn_ops, quant_ops, safety):
+ super().__init__()
+ self.ffn_ops = ffn_ops
+ self.quant_ops = quant_ops
+ self.in_features = mlp.fc1.weight.shape[1]
+ self.out_features = mlp.fc2.weight.shape[0]
+ device = mlp.fc1.weight.device
+ up_fp8, up_scale = _quantize_fp8(mlp.fc1.weight)
+ down_fp8, down_scale = _quantize_fp8(mlp.fc2.weight)
+ self.register_buffer("up_fp8", up_fp8.to(device))
+ self.register_buffer("down_fp8", down_fp8.to(device))
+ self.register_buffer("up_scale", up_scale.to(device))
+ self.register_buffer("down_scale", down_scale.to(device))
+ self.register_buffer("up_bias", mlp.fc1.bias.detach().to(torch.bfloat16))
+ self.register_buffer("down_bias", mlp.fc2.bias.detach().to(torch.bfloat16))
+ self.register_buffer("input_scale", _static_scale(in_amax, safety).to(device))
+ self.register_buffer("hidden_scale", _static_scale(hid_amax, safety).to(device))
+ self.register_buffer(
+ "channel_scale", torch.ones(self.in_features, device=device, dtype=torch.bfloat16)
+ )
+ self.safety = safety
+ self.calibrating = False
+ self._ia = 0.0
+ self._ha = 0.0
+
+ def _calibrate_step(self, x):
+ flat = x.reshape(-1, self.in_features).to(torch.bfloat16)
+ self._ia = max(self._ia, flat.float().abs().max().item())
+ self.input_scale.copy_(_static_scale(self._ia, self.safety).to(self.input_scale.device))
+ xdq = _roundtrip_fp8(flat.float(), self.input_scale)
+ hid = (xdq @ (self.up_fp8.float() * self.up_scale.float()).t()) + self.up_bias.float()
+ hid = F.gelu(hid, approximate="tanh")
+ self._ha = max(self._ha, hid.abs().max().item())
+ self.hidden_scale.copy_(_static_scale(self._ha, self.safety).to(self.hidden_scale.device))
+
+ def forward(self, x):
+ if self.calibrating:
+ self._calibrate_step(x)
+ shape = x.shape
+ dtype = x.dtype
+ flat = x.reshape(-1, self.in_features).to(torch.bfloat16)
+ x_fp8 = self.quant_ops.channel_scale_quantize_fp8_static_bf16(
+ flat, self.channel_scale, self.input_scale
+ )
+ out = self.ffn_ops.fp8_gelu_mlp_bf16(
+ x_fp8,
+ self.up_fp8,
+ self.up_bias,
+ self.down_fp8,
+ self.down_bias,
+ self.input_scale,
+ self.up_scale,
+ self.hidden_scale,
+ self.down_scale,
+ )
+ return out.reshape(*shape[:-1], self.out_features).to(dtype)
+
+
+def _siglip_mlps(model) -> list:
+ tower = model.paligemma_with_expert.paligemma.model.vision_tower
+ return [m for _, m in tower.named_modules() if type(m).__name__ == "SiglipMLP"]
+
+
+def _run_forward(policy, batches) -> None:
+ """Run eager action prediction so calibration reaches Python module forwards."""
+ model = policy.model
+ saved = {name: vars(model).pop(name) for name in ("sample_actions", "forward") if name in vars(model)}
+ with torch.inference_mode():
+ for batch in batches:
+ policy.predict_action_chunk(
+ {k: (v.clone() if torch.is_tensor(v) else v) for k, v in batch.items()}
+ )
+ torch.cuda.synchronize()
+ vars(model).update(saved)
+
+
+def _fixed_norm_weight(norm):
+ """Return a fixed RMSNorm fold weight, or ``None`` for adaptive norms."""
+ return norm.weight if getattr(norm, "dense", None) is None else None
+
+
+def _fp8_supported(device) -> bool:
+ """Return whether the device supports FP8 E4M3 tensor cores (CUDA SM >= 8.9)."""
+ if device.type != "cuda" or not torch.cuda.is_available():
+ return False
+ major, minor = torch.cuda.get_device_capability(device)
+ return (major, minor) >= (8, 9)
+
+
+def apply_fp8_mlp(policy, batch, *, safety: float = 1.05) -> bool:
+ """Replace Gemma and SigLIP MLPs with FlashRT FP8 kernels calibrated on the supplied batch.
+
+ Returns ``False`` without modifying BF16 execution when FP8 or its kernels are unavailable.
+ """
+ device = next(policy.parameters()).device
+ if not _fp8_supported(device):
+ logger.warning(
+ "PI052: device %s has no FP8 (E4M3) support (needs CUDA SM>=8.9); keeping BF16.",
+ device,
+ )
+ return False
+ batches = batch if isinstance(batch, (list, tuple)) else [batch]
+ try:
+ ffn_ops = _get_kernel(_SWIGLU_REPO)
+ gelu_ops = _get_kernel(_GELU_REPO)
+ quant_ops = _get_kernel(_GEMM_REPO)
+ except Exception as exc: # noqa: BLE001
+ logger.warning("PI052: FlashRT FP8 kernels unavailable (%s); keeping BF16.", exc)
+ return False
+
+ model = policy.model
+ calibrating = []
+
+ gemma_layers = list(model.paligemma_with_expert.gemma_expert.model.layers) + list(
+ model.paligemma_with_expert.paligemma.model.language_model.layers
+ )
+ for layer in gemma_layers:
+ fw = _fixed_norm_weight(layer.post_attention_layernorm)
+ layer.mlp = _FlashRTGeGLU(layer.mlp, 1.0, 1.0, ffn_ops, quant_ops, safety, fuse_weight=fw).to(device)
+ calibrating.append(layer.mlp)
+
+ siglip = _siglip_mlps(model)
+ for mlp_parent in model.paligemma_with_expert.paligemma.model.vision_tower.vision_model.encoder.layers:
+ mlp_parent.mlp = _FlashRTGeluMLP(mlp_parent.mlp, 1.0, 1.0, gelu_ops, quant_ops, safety).to(device)
+ calibrating.append(mlp_parent.mlp)
+
+ # Calibrate every swapped module in one FP8-propagated forward.
+ for m in calibrating:
+ m.calibrating = True
+ _run_forward(policy, batches)
+ for m in calibrating:
+ m.calibrating = False
+
+ logger.info(
+ "PI052: FlashRT FP8 enabled (%d Gemma + %d SigLIP MLPs).",
+ len(gemma_layers),
+ len(siglip),
+ )
+ return True
diff --git a/src/lerobot/policies/pi052/inference/__init__.py b/src/lerobot/policies/pi052/inference/__init__.py
new file mode 100644
index 000000000..8e8f1b3c7
--- /dev/null
+++ b/src/lerobot/policies/pi052/inference/__init__.py
@@ -0,0 +1,19 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI052 adapter for the policy-agnostic language runtime."""
+
+from .pi052_adapter import PI052PolicyAdapter
+
+__all__ = ["PI052PolicyAdapter"]
diff --git a/src/lerobot/policies/pi052/inference/pi052_adapter.py b/src/lerobot/policies/pi052/inference/pi052_adapter.py
new file mode 100644
index 000000000..a570ef1af
--- /dev/null
+++ b/src/lerobot/policies/pi052/inference/pi052_adapter.py
@@ -0,0 +1,254 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI052 actions and text generation for the generic language runtime."""
+
+from __future__ import annotations
+
+import logging
+from typing import Any
+
+from lerobot.runtime import RuntimeState
+from lerobot.runtime.adapter import BaseLanguageAdapter
+
+logger = logging.getLogger(__name__)
+
+_LOC_TOKENIZER_CACHE: dict[str, Any] = {}
+
+
+class PI052PolicyAdapter(BaseLanguageAdapter):
+ """Runtime bridge for PI052 policies."""
+
+ def select_action(self, observation: dict[str, Any], state: RuntimeState) -> Any:
+ import torch # noqa: PLC0415
+
+ from lerobot.utils.constants import ( # noqa: PLC0415
+ OBS_LANGUAGE_ATTENTION_MASK,
+ OBS_LANGUAGE_TOKENS,
+ OBS_STATE,
+ )
+
+ subtask = state.language_context.get("subtask") or state.task or ""
+ # Match the training prompt by conditioning on both subtask and discretized state.
+ state_str = None
+ obs_state = observation.get(OBS_STATE)
+ if isinstance(obs_state, torch.Tensor) and obs_state.numel() > 0:
+ from lerobot.policies.pi052.text_processor_pi052 import discretize_state_str # noqa: PLC0415
+
+ state_row = obs_state[0] if obs_state.ndim > 1 else obs_state
+ state_str = discretize_state_str(state_row)
+
+ batch = dict(observation)
+ if getattr(self.policy.config, "joint_subtask_conditioning", False):
+ # Joint sequences keep the task turn (with state) and render the
+ # subtask as a causal assistant turn, exactly as trained.
+ from transformers import AutoTokenizer # noqa: PLC0415
+
+ from lerobot.policies.pi052.text_processor_pi052 import ( # noqa: PLC0415
+ encode_prompt_with_targets,
+ register_paligemma_loc_tokens,
+ )
+ from lerobot.utils.constants import OBS_LANGUAGE_CAUSAL_MARKS # noqa: PLC0415
+
+ task = state.task or ""
+ task_content = task if state_str is None else f"{task}, State: {state_str};"
+ tok_name = getattr(self.policy.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
+ tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
+ ids, attn, marks = encode_prompt_with_targets(
+ tokenizer,
+ [
+ {"role": "user", "content": task_content},
+ {"role": "assistant", "content": subtask},
+ ],
+ target_indices=[1],
+ )
+ device = getattr(self.policy.config, "device", None)
+ if device is not None:
+ ids, attn, marks = ids.to(device), attn.to(device), marks.to(device)
+ batch[OBS_LANGUAGE_TOKENS] = ids
+ batch[OBS_LANGUAGE_ATTENTION_MASK] = attn
+ batch[OBS_LANGUAGE_CAUSAL_MARKS] = marks
+ else:
+ content = subtask if state_str is None else f"{subtask}, State: {state_str};"
+ text_batch = _build_text_batch(
+ self.policy,
+ [{"role": "user", "content": content}],
+ add_generation_prompt=False,
+ )
+ batch[OBS_LANGUAGE_TOKENS] = text_batch["lang_tokens"]
+ batch[OBS_LANGUAGE_ATTENTION_MASK] = text_batch["lang_masks"]
+ return self.policy.predict_action_chunk(batch)
+
+ def generate_text(
+ self,
+ kind: str,
+ observation: dict[str, Any] | None,
+ state: RuntimeState,
+ user_text: str | None = None,
+ ) -> str:
+ messages = self.build_messages(kind, state, user_text=user_text)
+ if kind == "subtask" and getattr(self.policy.config, "joint_subtask_conditioning", False):
+ # Joint samples carry state on the task turn, so the subtask must be
+ # generated from the same state-bearing prompt.
+ import torch # noqa: PLC0415
+
+ from lerobot.policies.pi052.text_processor_pi052 import discretize_state_str # noqa: PLC0415
+ from lerobot.utils.constants import OBS_STATE # noqa: PLC0415
+
+ obs_state = (observation or {}).get(OBS_STATE)
+ if isinstance(obs_state, torch.Tensor) and obs_state.numel() > 0:
+ state_row = obs_state[0] if obs_state.ndim > 1 else obs_state
+ for m in reversed(messages):
+ if m.get("role") == "user":
+ m["content"] = f"{m.get('content', '')}, State: {discretize_state_str(state_row)};"
+ break
+ return _generate_with_policy(
+ self.policy,
+ messages,
+ observation=observation,
+ state=state,
+ label=f"{kind} gen",
+ min_new_tokens=self.gen.min_new_tokens,
+ temperature=self.gen.temperature,
+ top_p=self.gen.top_p,
+ suppress_loc_tokens=True, # all runtime text is prose; never emit
+ )
+
+ def build_messages(
+ self,
+ kind: str,
+ state: RuntimeState,
+ *,
+ user_text: str | None = None,
+ ) -> list[dict[str, Any]]:
+ if kind in ("subtask", "plan"):
+ return [{"role": "user", "content": state.task or ""}]
+ if kind == "memory":
+ messages = [{"role": "user", "content": state.task or ""}]
+ if state.language_context.get("memory"):
+ messages.append(
+ {"role": "assistant", "content": f"Previous memory: {state.language_context['memory']}"}
+ )
+ if state.extra.get("prior_subtask"):
+ messages.append(
+ {"role": "user", "content": f"Completed subtask: {state.extra['prior_subtask']}"}
+ )
+ return messages
+ if kind == "interjection":
+ messages = [{"role": "user", "content": state.task or ""}]
+ if state.language_context.get("plan"):
+ messages.append(
+ {"role": "assistant", "content": f"Previous plan:\n{state.language_context['plan']}"}
+ )
+ if user_text:
+ messages.append({"role": "user", "content": user_text})
+ return messages
+ raise ValueError(f"Unknown PI052 text kind: {kind}")
+
+
+def _get_loc_tokenizer(tok_name: str, auto_tokenizer_cls: Any, register_loc_fn: Any) -> Any:
+ tokenizer = _LOC_TOKENIZER_CACHE.get(tok_name)
+ if tokenizer is None:
+ tokenizer = register_loc_fn(auto_tokenizer_cls.from_pretrained(tok_name))
+ _LOC_TOKENIZER_CACHE[tok_name] = tokenizer
+ return tokenizer
+
+
+def _build_text_batch(
+ policy: Any,
+ prompt_messages: list[dict[str, Any]],
+ *,
+ add_generation_prompt: bool = True,
+) -> dict[str, Any]:
+ import torch # noqa: PLC0415
+ from transformers import AutoTokenizer # noqa: PLC0415
+
+ from lerobot.policies.pi052.text_processor_pi052 import ( # noqa: PLC0415
+ _flatten_say_tool_calls,
+ _format_messages,
+ _strip_blocks,
+ register_paligemma_loc_tokens,
+ )
+
+ tok_name = getattr(policy.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
+ tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
+
+ messages = [_strip_blocks(_flatten_say_tool_calls(m)) for m in prompt_messages]
+ prompt, _spans = _format_messages(messages)
+ if add_generation_prompt:
+ # No trailing space: SentencePiece folds it into the first target token
+ # ("▁move"), so a space-suffixed prefill ends in a lone "▁" the model
+ # never saw at this position during training.
+ prompt = prompt + "Assistant:"
+
+ encoded = tokenizer(prompt, return_tensors="pt")
+ ids = encoded["input_ids"]
+ attn = encoded.get("attention_mask")
+ if attn is None and tokenizer.pad_token_id is not None:
+ attn = ids != tokenizer.pad_token_id
+ if attn is not None and hasattr(attn, "dtype") and attn.dtype != torch.bool:
+ attn = attn.bool()
+
+ device = getattr(getattr(policy, "config", None), "device", None)
+ if device is not None:
+ try:
+ ids = ids.to(device)
+ if attn is not None and hasattr(attn, "to"):
+ attn = attn.to(device)
+ except Exception as exc: # noqa: BLE001
+ logger.debug("could not move pi052 lang tokens to %s: %s", device, exc)
+ return {"lang_tokens": ids, "lang_masks": attn, "tokenizer": tokenizer}
+
+
+def _generate_with_policy(
+ policy: Any,
+ messages: list[dict[str, Any]],
+ *,
+ observation: dict[str, Any] | None = None,
+ state: RuntimeState | None = None,
+ label: str = "select_message",
+ min_new_tokens: int = 0,
+ temperature: float = 0.0,
+ top_p: float = 1.0,
+ suppress_loc_tokens: bool = False,
+) -> str:
+ if not hasattr(policy, "select_message"):
+ if state is not None:
+ state.log(f" [warn] policy has no select_message — skipping {label}")
+ return ""
+ text_batch = _build_text_batch(policy, messages)
+ try:
+ from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS # noqa: PLC0415
+
+ batch: dict[str, Any] = {
+ OBS_LANGUAGE_TOKENS: text_batch["lang_tokens"],
+ OBS_LANGUAGE_ATTENTION_MASK: text_batch["lang_masks"],
+ }
+ if observation:
+ for k, v in observation.items():
+ if isinstance(k, str) and k.startswith("observation.") and k not in batch:
+ batch[k] = v
+ return policy.select_message(
+ batch,
+ tokenizer=text_batch["tokenizer"],
+ min_new_tokens=min_new_tokens,
+ temperature=temperature,
+ top_p=top_p,
+ suppress_loc_tokens=suppress_loc_tokens,
+ )
+ except Exception as exc: # noqa: BLE001
+ logger.warning("%s failed: %s", label, exc, exc_info=logger.isEnabledFor(logging.DEBUG))
+ if state is not None:
+ state.log(f" [warn] {label} failed: {type(exc).__name__}: {exc}")
+ return ""
diff --git a/src/lerobot/policies/pi052/modeling_pi052.py b/src/lerobot/policies/pi052/modeling_pi052.py
new file mode 100644
index 000000000..b22423f46
--- /dev/null
+++ b/src/lerobot/policies/pi052/modeling_pi052.py
@@ -0,0 +1,2037 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI0.5 with joint flow/text training and hierarchical language inference."""
+
+from __future__ import annotations
+
+import logging
+import types
+from contextlib import nullcontext
+from typing import Any, Unpack
+
+import torch
+from torch import Tensor
+from torch.nn import functional
+
+from lerobot.utils.constants import (
+ ACTION,
+ OBS_LANGUAGE_ATTENTION_MASK,
+ OBS_LANGUAGE_CAUSAL_MARKS,
+ OBS_LANGUAGE_TOKENS,
+ OBS_STATE,
+)
+
+from ..pi05.modeling_pi05 import (
+ ActionSelectKwargs,
+ PI05Policy,
+ PI05Pytorch as PI05PytorchBase,
+ make_att_2d_masks,
+)
+from .configuration_pi052 import PI052Config
+
+logger = logging.getLogger(__name__)
+
+
+class PI05Pytorch(PI05PytorchBase): # see openpi `PI0Pytorch`
+ """Core PI05 PyTorch model."""
+
+ use_hf_vision_checkpointing_api = True
+ checkpoint_vision_embeddings = False
+ use_typed_attention_masks = True
+ use_on_device_suffix_mask = True
+ precompute_denoise_times = True
+
+ def forward(self, images, img_masks, tokens, masks, actions, noise, time) -> Tensor:
+ """Do a full training forward pass and compute the loss."""
+ time_expanded = time[:, None, None]
+ x_t = time_expanded * noise + (1 - time_expanded) * actions
+ u_t = noise - actions
+
+ # The prefix has no gradient path to a flow-only loss under KI.
+ suppress_prefix = bool(getattr(self.config, "knowledge_insulation", False))
+ with torch.no_grad() if suppress_prefix else nullcontext():
+ prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix(
+ images, img_masks, tokens, masks
+ )
+ suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, time)
+
+ if (
+ self.paligemma_with_expert.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype
+ == torch.bfloat16
+ ):
+ suffix_embs = suffix_embs.to(dtype=torch.bfloat16)
+ prefix_embs = prefix_embs.to(dtype=torch.bfloat16)
+
+ pad_masks = torch.cat([prefix_pad_masks, suffix_pad_masks], dim=1)
+ att_masks = torch.cat([prefix_att_masks, suffix_att_masks], dim=1)
+
+ att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
+ position_ids = torch.cumsum(pad_masks, dim=1) - 1
+
+ att_2d_masks_4d = self._prepare_attention_masks_4d(att_2d_masks, dtype=prefix_embs.dtype)
+
+ # Transformer and vision layers own their checkpoint boundaries.
+ ki_kwargs = {"suppress_prefix_grads": True} if suppress_prefix else {}
+ (_, suffix_out), _ = self.paligemma_with_expert.forward(
+ attention_mask=att_2d_masks_4d,
+ position_ids=position_ids,
+ past_key_values=None,
+ inputs_embeds=[prefix_embs, suffix_embs],
+ use_cache=False,
+ adarms_cond=[None, adarms_cond],
+ **ki_kwargs,
+ )
+
+ suffix_out = suffix_out[:, -self.config.chunk_size :]
+ suffix_out = suffix_out.to(dtype=torch.float32)
+
+ def action_out_proj_func(suffix_out):
+ return self.action_out_proj(suffix_out)
+
+ v_t = self._apply_checkpoint(action_out_proj_func, suffix_out)
+
+ return functional.mse_loss(u_t, v_t, reduction="none")
+
+ def denoise_step(
+ self,
+ prefix_pad_masks,
+ past_key_values,
+ x_t,
+ timestep,
+ ):
+ """Apply one denoising step of the noise `x_t` at a given timestep."""
+ suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = self.embed_suffix(x_t, timestep)
+
+ suffix_len = suffix_pad_masks.shape[1]
+ batch_size = prefix_pad_masks.shape[0]
+ prefix_len = prefix_pad_masks.shape[1]
+
+ prefix_pad_2d_masks = prefix_pad_masks[:, None, :].expand(batch_size, suffix_len, prefix_len)
+ suffix_att_2d_masks = make_att_2d_masks(suffix_pad_masks, suffix_att_masks)
+ full_att_2d_masks = torch.cat([prefix_pad_2d_masks, suffix_att_2d_masks], dim=2)
+
+ prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None]
+ position_ids = prefix_offsets + torch.cumsum(suffix_pad_masks, dim=1) - 1
+
+ full_att_2d_masks_4d = self._prepare_attention_masks_4d(full_att_2d_masks, dtype=suffix_embs.dtype)
+ self.paligemma_with_expert.gemma_expert.model.config._attn_implementation = "eager" # noqa: SLF001
+
+ # Crop appended suffix K/V after each step instead of copying the read-only prefix cache.
+ outputs_embeds, _ = self.paligemma_with_expert.forward(
+ attention_mask=full_att_2d_masks_4d,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ inputs_embeds=[None, suffix_embs],
+ use_cache=False,
+ adarms_cond=[None, adarms_cond],
+ )
+ past_key_values.crop(prefix_len)
+
+ suffix_out = outputs_embeds[1]
+ suffix_out = suffix_out[:, -self.config.chunk_size :]
+ suffix_out = suffix_out.to(dtype=torch.float32)
+ return self.action_out_proj(suffix_out)
+
+ def sample_actions(
+ self,
+ images,
+ img_masks,
+ tokens,
+ masks,
+ noise=None,
+ num_steps=None,
+ lang_causal_marks=None,
+ **kwargs,
+ ) -> Tensor:
+ """Sample actions, optionally marking trailing language positions causal.
+
+ ``lang_causal_marks`` (B, L_lang bool) flags generated-subtask tokens so
+ joint-sequence checkpoints see the same causal prefix layout at
+ inference as during training (``_mark_target_span_causal``).
+ """
+ self._lang_causal_marks = lang_causal_marks
+ try:
+ return super().sample_actions(
+ images, img_masks, tokens, masks, noise=noise, num_steps=num_steps, **kwargs
+ )
+ finally:
+ self._lang_causal_marks = None
+
+ def embed_prefix(self, images, img_masks, tokens, masks):
+ prefix_embs, prefix_pad, prefix_att = super().embed_prefix(images, img_masks, tokens, masks)
+ marks = getattr(self, "_lang_causal_marks", None)
+ if marks is not None:
+ prefix_att = _apply_causal_language_marks(prefix_att, marks.to(prefix_att.device))
+ return prefix_embs, prefix_pad, prefix_att
+
+
+# The universal `physical-intelligence/fast` tokenizer (and dataset refits of it)
+# uses 1024 BPE codes; text generation must mask any that map below the range.
+_FAST_ACTION_VOCAB_SIZE = 1024
+
+
+_HF_KERNELS_ENABLED = False
+
+
+def _enable_hf_kernels() -> None:
+ """Patch supported PaliGemma operations before constructing the model."""
+ global _HF_KERNELS_ENABLED
+ if _HF_KERNELS_ENABLED:
+ return
+ try:
+ from liger_kernel.transformers import apply_liger_kernel_to_paligemma # noqa: PLC0415
+ except ImportError:
+ logger.warning(
+ "PI052: liger-kernel is not installed; skipping fused Triton "
+ "kernels. Install with ``pip install liger-kernel``."
+ )
+ return
+ apply_liger_kernel_to_paligemma(
+ rope=True,
+ geglu=True,
+ # Liger LayerNorm regresses at SigLIP shapes; RoPE and GeGLU remain enabled.
+ layer_norm=False,
+ rms_norm=False,
+ cross_entropy=False,
+ fused_linear_cross_entropy=False,
+ )
+ _HF_KERNELS_ENABLED = True
+ logger.info("PI052: HF kernels (Liger) enabled — rope, geglu fused.")
+
+
+def _reduce_action_loss(per_sample: Tensor, predict_actions_t: Tensor | None, reduction: str) -> Tensor:
+ """Mask non-action samples and apply the requested batch reduction."""
+ if predict_actions_t is None:
+ return per_sample if reduction == "none" else per_sample.mean()
+ mask = predict_actions_t.to(per_sample.dtype)
+ if reduction == "none":
+ return per_sample * mask
+ return (per_sample * mask).sum() / mask.sum().clamp(min=1.0)
+
+
+# Materialized logits win at VLA token counts; larger dense targets use Liger.
+_LOGITS_CE_MAX_POSITIONS = 2048
+
+
+def _lin_ce_small(
+ flat_hidden: Tensor,
+ lm_head_weight: Tensor,
+ flat_labels: Tensor,
+ z_loss_weight: float = 0.0,
+) -> Tensor:
+ """Small-N linear CE on materialized logits (see ``_lin_ce_flat``)."""
+ logits = (flat_hidden @ lm_head_weight.t()).float()
+ n_valid = (flat_labels != -100).sum().clamp(min=1)
+ loss = functional.cross_entropy(logits, flat_labels, ignore_index=-100, reduction="sum") / n_valid
+ if z_loss_weight > 0:
+ lse = torch.logsumexp(logits, dim=-1)
+ valid = (flat_labels != -100).to(lse.dtype)
+ loss = loss + float(z_loss_weight) * (lse.square() * valid).sum() / n_valid
+ return loss
+
+
+# Built lazily so importing this module does not invoke Dynamo.
+_compiled_lin_ce_small = None
+
+
+def _get_compiled_lin_ce_small():
+ global _compiled_lin_ce_small
+ if _compiled_lin_ce_small is None:
+ _compiled_lin_ce_small = torch.compile(_lin_ce_small, dynamic=False)
+ return _compiled_lin_ce_small
+
+
+def _lin_ce_flat(
+ flat_hidden: Tensor,
+ lm_head_weight: Tensor,
+ flat_labels: Tensor,
+ z_loss_weight: float = 0.0,
+ compiled: bool = False,
+) -> Tensor:
+ """Dispatch sparse targets to fixed logits buckets and dense targets to Liger."""
+ if flat_hidden.shape[0] > _LOGITS_CE_MAX_POSITIONS:
+ valid = flat_labels != -100
+ compact_hidden = flat_hidden[valid]
+ compact_labels = flat_labels[valid]
+ compact_rows = compact_hidden.shape[0]
+
+ if compact_rows == 0:
+ return _lin_ce_flat(
+ functional.pad(compact_hidden, (0, 0, 0, 1)),
+ lm_head_weight,
+ functional.pad(compact_labels, (0, 1), value=-100),
+ z_loss_weight,
+ compiled=compiled,
+ )
+
+ # Fixed power-of-two buckets avoid shape churn while keeping sparse
+ # supervision on the materialized-logits path.
+ bucket_rows = 1 << (compact_rows - 1).bit_length()
+ if bucket_rows < flat_hidden.shape[0]:
+ weighted_losses = []
+ for start in range(0, compact_rows, _LOGITS_CE_MAX_POSITIONS):
+ end = min(start + _LOGITS_CE_MAX_POSITIONS, compact_rows)
+ rows = end - start
+ chunk_rows = 1 << (rows - 1).bit_length()
+ hidden_chunk = compact_hidden[start:end]
+ labels_chunk = compact_labels[start:end]
+ pad_rows = chunk_rows - rows
+ if pad_rows:
+ hidden_chunk = functional.pad(hidden_chunk, (0, 0, 0, pad_rows))
+ labels_chunk = functional.pad(labels_chunk, (0, pad_rows), value=-100)
+ chunk_loss = _lin_ce_flat(
+ hidden_chunk,
+ lm_head_weight,
+ labels_chunk,
+ z_loss_weight,
+ compiled=compiled,
+ )
+ weighted_losses.append(chunk_loss * rows)
+ return torch.stack(weighted_losses).sum() / compact_rows
+
+ if flat_hidden.shape[0] <= _LOGITS_CE_MAX_POSITIONS:
+ fn = _get_compiled_lin_ce_small() if compiled else _lin_ce_small
+ return fn(flat_hidden, lm_head_weight, flat_labels, z_loss_weight)
+
+ # Keep Liger optional for inference-only installations.
+ from liger_kernel.transformers.fused_linear_cross_entropy import ( # noqa: PLC0415
+ LigerFusedLinearCrossEntropyLoss,
+ )
+
+ loss_fn = LigerFusedLinearCrossEntropyLoss(
+ ignore_index=-100,
+ lse_square_scale=float(z_loss_weight),
+ reduction="mean",
+ )
+ return loss_fn(lm_head_weight, flat_hidden, flat_labels)
+
+
+def _shifted_lin_ce(
+ hidden: Tensor,
+ lm_head_weight: Tensor,
+ labels: Tensor,
+ z_loss_weight: float = 0.0,
+ compiled: bool = False,
+ reduction: str = "mean",
+) -> Tensor:
+ """Compute next-token CE through the shape-aware linear-CE dispatcher."""
+ shift_hidden = hidden[:, :-1, :].contiguous()
+ shift_labels = labels[:, 1:].contiguous().long()
+ if reduction == "none":
+ return torch.stack(
+ [
+ _lin_ce_flat(
+ sample_hidden.to(lm_head_weight.dtype),
+ lm_head_weight,
+ sample_labels,
+ z_loss_weight,
+ compiled=compiled,
+ )
+ for sample_hidden, sample_labels in zip(shift_hidden, shift_labels, strict=True)
+ ]
+ )
+ batch_size, target_length, hidden_size = shift_hidden.shape
+ flat_hidden = shift_hidden.reshape(batch_size * target_length, hidden_size)
+ flat_labels = shift_labels.reshape(batch_size * target_length)
+ # Match the dtype the eager path used: cast hidden to the lm_head's
+ # weight dtype so bf16 weights see bf16 activations.
+ flat_hidden = flat_hidden.to(lm_head_weight.dtype)
+ return _lin_ce_flat(flat_hidden, lm_head_weight, flat_labels, z_loss_weight, compiled=compiled)
+
+
+def _mark_target_span_causal(
+ prefix_att_masks: Tensor, text_labels: Tensor, lang_start: int, lang_end: int
+) -> Tensor:
+ """Make supervised language targets causal while leaving prompts bidirectional."""
+ att = prefix_att_masks.clone()
+ n = min(text_labels.shape[1], lang_end - lang_start)
+ if n <= 0:
+ return att
+ target = text_labels[:, :n] != -100 # (B, n) bool
+ seg = att[:, lang_start : lang_start + n].bool()
+ att[:, lang_start : lang_start + n] = seg | target
+ return att
+
+
+def _apply_causal_language_marks(prefix_att_masks: Tensor, marks: Tensor) -> Tensor:
+ """OR per-token causal marks into the trailing language segment of a prefix."""
+ att = prefix_att_masks.clone()
+ n = min(marks.shape[1], att.shape[1])
+ if n <= 0:
+ return att
+ seg = att[:, -n:].bool()
+ att[:, -n:] = (seg | marks[:, -n:].bool()).to(att.dtype)
+ return att
+
+
+def _fast_lin_ce(
+ hidden: Tensor,
+ lm_head_weight: Tensor,
+ action_tokens: Tensor,
+ action_code_mask: Tensor,
+ predict_actions_t: Tensor | None,
+ compiled: bool = False,
+ reduction: str = "mean",
+) -> Tensor:
+ """Compute FAST token CE over the enabled action-code positions."""
+ shift_hidden = hidden[:, :-1, :].contiguous()
+ shift_targets = action_tokens[:, 1:].contiguous().long()
+ shift_valid = action_code_mask[:, 1:].contiguous().bool()
+ if predict_actions_t is not None:
+ sample_mask = predict_actions_t[:, None].expand_as(shift_valid)
+ shift_valid = shift_valid & sample_mask
+ # Encode the mask with ignore_index to avoid a host sync and preserve graph capture.
+ shift_targets = torch.where(shift_valid, shift_targets, torch.full_like(shift_targets, -100))
+
+ if reduction == "none":
+ return torch.stack(
+ [
+ _lin_ce_flat(
+ sample_hidden.to(lm_head_weight.dtype),
+ lm_head_weight,
+ sample_labels,
+ compiled=compiled,
+ )
+ for sample_hidden, sample_labels in zip(shift_hidden, shift_targets, strict=True)
+ ]
+ )
+
+ valid_counts = shift_valid.sum(dim=1)
+ active_samples = valid_counts > 0
+ if not bool(active_samples.any().item()):
+ return shift_hidden.sum() * 0.0
+
+ weighted_losses = []
+ active_count = active_samples.sum()
+ for token_count in torch.unique(valid_counts[active_samples]).tolist():
+ group = active_samples & valid_counts.eq(token_count)
+ group_size = group.sum()
+ group_hidden = shift_hidden[group].reshape(-1, shift_hidden.shape[-1]).to(lm_head_weight.dtype)
+ group_labels = shift_targets[group].reshape(-1)
+ group_loss = _lin_ce_flat(
+ group_hidden,
+ lm_head_weight,
+ group_labels,
+ compiled=compiled,
+ )
+ weighted_losses.append(group_loss * group_size)
+ return torch.stack(weighted_losses).sum() / active_count
+
+
+# ----------------------------------------------------------------------
+# Knowledge insulation helpers
+# ----------------------------------------------------------------------
+# Action queries consume detached VLM K/V. Flow-only callers may additionally
+# suppress the now-dead prefix graph without changing forward values.
+
+
+# Consumer GPUs need smaller FlexAttention backward tiles at head_dim=256.
+_FLEX_SHRUNK_TILES = {"BLOCK_M1": 32, "BLOCK_N1": 64, "BLOCK_M2": 64, "BLOCK_N2": 32}
+_flex_kernel_options: dict[int, dict | None] = {}
+_flex_fns: tuple | None | bool = None
+
+
+def _get_flex_kernel_options(device: torch.device) -> dict | None:
+ if device.type != "cuda" or not torch.cuda.is_available():
+ return None
+ device_index = device.index if device.index is not None else torch.cuda.current_device()
+ if device_index not in _flex_kernel_options:
+ smem = torch.cuda.get_device_properties(
+ device_index
+ ).shared_memory_per_block_optin # spellchecker:disable-line
+ _flex_kernel_options[device_index] = _FLEX_SHRUNK_TILES if smem < 128 * 1024 else None
+ return _flex_kernel_options[device_index]
+
+
+def _get_flex_fns(device: torch.device):
+ """Return compiled FlexAttention helpers when available."""
+ global _flex_fns
+ if device.type != "cuda" or not torch.cuda.is_available():
+ return None
+ if _flex_fns is None:
+ try:
+ from torch.nn.attention.flex_attention import ( # noqa: PLC0415
+ create_block_mask,
+ flex_attention,
+ )
+
+ _flex_fns = (
+ torch.compile(flex_attention, dynamic=False),
+ torch.compile(create_block_mask, dynamic=False),
+ )
+ _get_flex_kernel_options(device)
+ except Exception as exc:
+ logger.warning("PI052: FlexAttention unavailable (%s); using SDPA.", exc)
+ _flex_fns = False
+ return _flex_fns or None
+
+
+class _FlexMaskBuilder:
+ """Build KI masks while retaining stable compiled mask callables."""
+
+ def __init__(self):
+ self._key = None
+
+ def build(self, prefix_pad, prefix_att, non_fast_prefix_len, k, chunk):
+ _, create_bm = _get_flex_fns(prefix_pad.device)
+ b, p = prefix_pad.shape
+ a = k * chunk
+ device = prefix_pad.device
+ key = (b, p, a, int(non_fast_prefix_len), device)
+ if self._key != key:
+ self._key = key
+ self._pad = torch.empty(b, p, dtype=torch.bool, device=device)
+ self._cum = torch.empty(b, p, dtype=torch.long, device=device)
+ pad, cum = self._pad, self._cum
+ nf = int(non_fast_prefix_len)
+
+ def vlm_rows(bi, h, q_idx, kv_idx):
+ kv_p = kv_idx.clamp(max=p - 1)
+ ok = (cum[bi, kv_p] <= cum[bi, q_idx]) & pad[bi, kv_p] & pad[bi, q_idx]
+ return (kv_idx < p) & ok
+
+ def action_rows(bi, h, q_idx, kv_idx):
+ kv_p = kv_idx.clamp(max=p - 1)
+ to_prefix = (kv_idx < nf) & pad[bi, kv_p]
+ same_block = (q_idx // chunk) == ((kv_idx - p) // chunk)
+ return to_prefix | ((kv_idx >= p) & same_block)
+
+ self._vlm_mod, self._action_mod = vlm_rows, action_rows
+
+ self._pad.copy_(prefix_pad)
+ self._cum.copy_(torch.cumsum(prefix_att.to(torch.long), dim=1))
+ s = p + a
+ bm_vlm = create_bm(self._vlm_mod, B=b, H=None, Q_LEN=p, KV_LEN=s, device=device)
+ bm_action = create_bm(self._action_mod, B=b, H=None, Q_LEN=a, KV_LEN=s, device=device)
+ return bm_vlm, bm_action
+
+
+# Lazily loaded FlashRT AdaRMS backend; unsupported cases use eager PyTorch.
+_flashrt_adarms_cache = None
+
+
+def _get_adarms_backend():
+ global _flashrt_adarms_cache
+ if _flashrt_adarms_cache is None:
+ try:
+ from kernels import get_kernel # noqa: PLC0415
+
+ _flashrt_adarms_cache = get_kernel("flashrt/flashrt-adarms-train", revision="v1")
+ except Exception as exc:
+ logger.warning(
+ "PI052: flashrt-adarms-train unavailable (%s); using the eager norm path.",
+ exc,
+ )
+ _flashrt_adarms_cache = False
+ return _flashrt_adarms_cache or None
+
+
+def _adarms_norm(backend, norm, x, cond):
+ """PiGemmaRMSNorm through the fused kernel when a backend is present."""
+ if backend is not None:
+ if cond is not None and norm.dense is not None:
+ return backend.adarms(x, norm.dense(cond), norm.eps, True)
+ if norm.dense is None:
+ return backend.adarms(x, norm.weight, norm.eps, False)
+ return norm(x, cond=cond)
+
+
+def _manual_attention_part(qs, ks, vs, m, scale):
+ """Materialized-logits GQA with an FP32 softmax."""
+ batch_size, num_heads, query_length, head_dim = qs.shape
+ num_kv_heads = ks.shape[1]
+ if num_kv_heads != num_heads:
+ groups = num_heads // num_kv_heads
+ grouped_queries = qs.reshape(batch_size, num_kv_heads, groups * query_length, head_dim)
+ logits = (grouped_queries @ ks.transpose(-1, -2)).reshape(batch_size, num_heads, query_length, -1)
+ else:
+ logits = qs @ ks.transpose(-1, -2)
+ logits = logits * scale + m
+ p = logits.float().softmax(dim=-1).to(qs.dtype)
+ out = (
+ (p.reshape(batch_size, num_kv_heads, groups * query_length, -1) @ vs).reshape(
+ batch_size, num_heads, query_length, head_dim
+ )
+ if num_kv_heads != num_heads
+ else p @ vs
+ )
+ return out.transpose(1, 2).contiguous()
+
+
+# Knowledge insulation keeps the forward equivalent while detaching VLM K/V for action-query gradients.
+_manual_attention = None
+
+
+def _get_manual_attention():
+ """Load the Hub implementation, with the inline function as fallback."""
+ global _manual_attention
+ if _manual_attention is None:
+ part = _manual_attention_part
+ try:
+ from kernels import get_kernel # noqa: PLC0415
+
+ _hub = getattr(
+ get_kernel("flashrt/flashrt-flex-attention-train"),
+ "manual_attention_part",
+ None,
+ )
+ if _hub is not None:
+
+ def part(qs, ks, vs, m, scale, _hub=_hub):
+ return _hub(qs, ks, vs, m, scale).transpose(1, 2).contiguous()
+
+ logger.info("PI052: manual attention backed by flashrt-flex-attention-train (Hub).")
+ except Exception as exc:
+ logger.info(
+ "PI052: flashrt-flex-attention-train unavailable (%s); using the inline manual-attention path.",
+ exc,
+ )
+ _manual_attention = torch.compile(part, dynamic=False)
+ return _manual_attention
+
+
+def _compute_layer_ki(
+ layer_idx,
+ inputs_embeds,
+ attention_mask,
+ position_embeddings,
+ adarms_cond,
+ paligemma,
+ gemma_expert,
+ suppress_prefix_grads=False,
+ flex_masks=None,
+ adarms_backend=None,
+ manual_attention=False,
+):
+ from transformers.models.gemma import modeling_gemma # noqa: PLC0415
+
+ # ``_gated_residual`` is LeRobot's adaRMSNorm helper, not a Transformers symbol.
+ from ..pi_gemma import _gated_residual # noqa: PLC0415
+
+ def _vlm_ctx(i):
+ return torch.no_grad() if (i == 0 and suppress_prefix_grads) else nullcontext()
+
+ models = [paligemma.model.language_model, gemma_expert.model]
+ query_states, key_states, value_states, gates = [], [], [], []
+
+ vlm_len = inputs_embeds[0].shape[1]
+
+ for i, hidden_states in enumerate(inputs_embeds):
+ layer = models[i].layers[layer_idx]
+ with _vlm_ctx(i):
+ hidden_states, gate = _adarms_norm(
+ adarms_backend, layer.input_layernorm, hidden_states, adarms_cond[i]
+ )
+ gates.append(gate)
+ input_shape = hidden_states.shape[:-1]
+ hidden_shape = (*input_shape, -1, layer.self_attn.head_dim)
+ q = layer.self_attn.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ k = layer.self_attn.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ v = layer.self_attn.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ query_states.append(q)
+ key_states.append(k)
+ value_states.append(v)
+
+ query_states = torch.cat(query_states, dim=2)
+ key_states = torch.cat(key_states, dim=2)
+ value_states = torch.cat(value_states, dim=2)
+
+ cos, sin = position_embeddings
+ query_states, key_states = modeling_gemma.apply_rotary_pos_emb(
+ query_states, key_states, cos, sin, unsqueeze_dim=1
+ )
+
+ batch_size = query_states.shape[0]
+ scaling = paligemma.model.language_model.layers[layer_idx].self_attn.scaling
+
+ # Split queries / K / V at the VLM-vs-action boundary.
+ q_vlm = query_states[:, :, :vlm_len, :]
+ q_action = query_states[:, :, vlm_len:, :]
+ k_vlm = key_states[:, :, :vlm_len, :]
+ k_action = key_states[:, :, vlm_len:, :]
+ v_vlm = value_states[:, :, :vlm_len, :]
+ v_action = value_states[:, :, vlm_len:, :]
+
+ # Detach VLM K/V *only* on the path the action queries use.
+ k_for_vlm = key_states
+ v_for_vlm = value_states
+ k_for_action = torch.cat([k_vlm.detach(), k_action], dim=2)
+ v_for_action = torch.cat([v_vlm.detach(), v_action], dim=2)
+
+ if flex_masks is not None:
+ flex_attn, _ = _get_flex_fns(query_states.device)
+ bm_vlm, bm_action = flex_masks
+ n_rep = paligemma.model.language_model.layers[layer_idx].self_attn.num_key_value_groups
+ with _vlm_ctx(0):
+ att_vlm = flex_attn(
+ q_vlm,
+ k_for_vlm,
+ v_for_vlm,
+ block_mask=bm_vlm,
+ scale=scaling,
+ enable_gqa=n_rep > 1,
+ kernel_options=_get_flex_kernel_options(query_states.device),
+ ).transpose(1, 2)
+ att_action = flex_attn(
+ q_action,
+ k_for_action,
+ v_for_action,
+ block_mask=bm_action,
+ scale=scaling,
+ enable_gqa=n_rep > 1,
+ kernel_options=_get_flex_kernel_options(query_states.device),
+ ).transpose(1, 2)
+ else:
+ mask_for_vlm = attention_mask[:, :, :vlm_len, :]
+ mask_for_action = attention_mask[:, :, vlm_len:, :]
+ # SDPA requires the additive bias to match each query dtype.
+ if mask_for_vlm.dtype != q_vlm.dtype:
+ mask_for_vlm = mask_for_vlm.to(dtype=q_vlm.dtype)
+ if mask_for_action.dtype != q_action.dtype:
+ mask_for_action = mask_for_action.to(dtype=q_action.dtype)
+
+ if manual_attention:
+ manual_fn = _get_manual_attention()
+ if manual_attention == "action":
+ from ..pi_gemma import sdpa_attention_forward # noqa: PLC0415
+
+ with _vlm_ctx(0):
+ att_vlm, _ = sdpa_attention_forward(
+ paligemma.model.language_model.layers[layer_idx].self_attn,
+ q_vlm,
+ k_for_vlm,
+ v_for_vlm,
+ mask_for_vlm,
+ scaling,
+ )
+ else:
+ with _vlm_ctx(0):
+ att_vlm = manual_fn(q_vlm, k_for_vlm, v_for_vlm, mask_for_vlm, scaling)
+ att_action = manual_fn(q_action, k_for_action, v_for_action, mask_for_action, scaling)
+ else:
+ from ..pi_gemma import sdpa_attention_forward # noqa: PLC0415
+
+ with _vlm_ctx(0):
+ att_vlm, _ = sdpa_attention_forward(
+ paligemma.model.language_model.layers[layer_idx].self_attn,
+ q_vlm,
+ k_for_vlm,
+ v_for_vlm,
+ mask_for_vlm,
+ scaling,
+ )
+ att_action, _ = sdpa_attention_forward(
+ paligemma.model.language_model.layers[layer_idx].self_attn,
+ q_action,
+ k_for_action,
+ v_for_action,
+ mask_for_action,
+ scaling,
+ )
+ att = torch.cat([att_vlm, att_action], dim=1)
+
+ head_dim = paligemma.model.language_model.layers[layer_idx].self_attn.head_dim
+ att = att.reshape(batch_size, -1, 1 * 8 * head_dim)
+
+ outputs_embeds = []
+ start = 0
+ for i, hidden_states in enumerate(inputs_embeds):
+ layer = models[i].layers[layer_idx]
+ end = start + hidden_states.shape[1]
+ if att.dtype != layer.self_attn.o_proj.weight.dtype:
+ att = att.to(layer.self_attn.o_proj.weight.dtype)
+ with _vlm_ctx(i):
+ out_emb = layer.self_attn.o_proj(att[:, start:end])
+ pa_norm = layer.post_attention_layernorm
+ if adarms_backend is not None:
+ if adarms_cond[i] is not None and pa_norm.dense is not None:
+ after_first, out_emb, gate = adarms_backend.resgate_adarms(
+ hidden_states,
+ out_emb,
+ gates[i],
+ pa_norm.dense(adarms_cond[i]),
+ pa_norm.eps,
+ True,
+ )
+ else:
+ after_first, out_emb, gate = adarms_backend.resgate_adarms(
+ hidden_states, out_emb, gates[i], pa_norm.weight, pa_norm.eps, False
+ )
+ else:
+ out_emb = _gated_residual(hidden_states, out_emb, gates[i])
+ after_first = out_emb
+ out_emb, gate = pa_norm(out_emb, cond=adarms_cond[i])
+ if layer.mlp.up_proj.weight.dtype == torch.bfloat16:
+ out_emb = out_emb.to(dtype=torch.bfloat16)
+ out_emb = layer.mlp(out_emb)
+ out_emb = _gated_residual(after_first, out_emb, gate)
+ outputs_embeds.append(out_emb)
+ start = end
+ return outputs_embeds
+
+
+def _paligemma_forward_ki(
+ self,
+ attention_mask=None,
+ position_ids=None,
+ past_key_values=None,
+ inputs_embeds=None,
+ use_cache=None,
+ adarms_cond=None,
+ suppress_prefix_grads=False,
+ flex_masks=None,
+ adarms_backend=None,
+ manual_attention=False,
+):
+ """Run dual-expert layers through KI and defer single-expert calls."""
+ from ..pi_gemma import layernorm_forward # noqa: PLC0415
+
+ if adarms_cond is None:
+ adarms_cond = [None, None]
+
+ # Single-expert paths: defer to the original forward saved in
+ # PI052Policy.__init__.
+ if inputs_embeds[0] is None or inputs_embeds[1] is None:
+ return self._pi052_orig_forward(
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ inputs_embeds=inputs_embeds,
+ use_cache=use_cache,
+ adarms_cond=adarms_cond,
+ )
+
+ models = [self.paligemma.model.language_model, self.gemma_expert.model]
+ num_layers = self.paligemma.config.text_config.num_hidden_layers
+
+ # RoPE values are shared by every layer.
+ position_embeddings = self.paligemma.model.language_model.rotary_emb(inputs_embeds[0], position_ids)
+
+ use_gc = (
+ hasattr(self.gemma_expert.model, "gradient_checkpointing")
+ and self.gemma_expert.model.gradient_checkpointing
+ and self.training
+ ) or (hasattr(self, "gradient_checkpointing") and self.gradient_checkpointing and self.training)
+
+ for layer_idx in range(num_layers):
+ if use_gc:
+ inputs_embeds = torch.utils.checkpoint.checkpoint(
+ _compute_layer_ki,
+ layer_idx,
+ inputs_embeds,
+ attention_mask,
+ position_embeddings,
+ adarms_cond,
+ use_reentrant=False,
+ preserve_rng_state=False,
+ paligemma=self.paligemma,
+ gemma_expert=self.gemma_expert,
+ suppress_prefix_grads=suppress_prefix_grads,
+ flex_masks=flex_masks,
+ adarms_backend=adarms_backend,
+ manual_attention=manual_attention,
+ )
+ else:
+ inputs_embeds = _compute_layer_ki(
+ layer_idx,
+ inputs_embeds,
+ attention_mask,
+ position_embeddings,
+ adarms_cond,
+ paligemma=self.paligemma,
+ gemma_expert=self.gemma_expert,
+ suppress_prefix_grads=suppress_prefix_grads,
+ flex_masks=flex_masks,
+ adarms_backend=adarms_backend,
+ manual_attention=manual_attention,
+ )
+
+ outputs_embeds = []
+ for i, hidden_states in enumerate(inputs_embeds):
+ with torch.no_grad() if (i == 0 and suppress_prefix_grads) else nullcontext():
+ out_emb, _ = layernorm_forward(models[i].norm, hidden_states, adarms_cond[i])
+ outputs_embeds.append(out_emb)
+ return [outputs_embeds[0], outputs_embeds[1]], None
+
+
+class PI052Policy(PI05Policy):
+ """π0.5 with the PaliGemma LM head re-enabled.
+
+ It inherits unchanged PI0.5 policy behavior and replaces the core model with
+ the joint flow/text implementation below.
+ """
+
+ config_class = PI052Config
+ name = "pi052"
+ model_class = PI05Pytorch
+ eval_after_pretrained_load = True
+ show_openpi_disclaimer = False
+
+ def __init__(self, config: PI052Config, **kwargs: Any) -> None:
+ # Patch before constructing Gemma/SigLIP layers; the operation is optional and idempotent.
+ _enable_hf_kernels()
+ super().__init__(config, **kwargs)
+
+ # Re-enable layers PI0.5 freezes when text supervision is requested.
+ if config.text_loss_weight > 0 and config.unfreeze_lm_head:
+ self._unfreeze_lm_head()
+
+ # Bind knowledge insulation per instance so stock PI0.5 policies remain unchanged.
+ if getattr(config, "knowledge_insulation", False):
+ backbone = self.model.paligemma_with_expert
+ backbone._pi052_orig_forward = backbone.forward
+ backbone.forward = types.MethodType(_paligemma_forward_ki, backbone)
+ logger.info(
+ "PI052: knowledge insulation enabled — action→VLM K/V gradients are blocked in attention."
+ )
+ if config.use_flashrt_adarms:
+ self._flashrt_adarms = _get_adarms_backend()
+ if self._flashrt_adarms is not None:
+ logger.info("PI052: FlashRT adaRMS training kernels enabled.")
+
+ if config.use_compiled_vision:
+ _tower = self.model.paligemma_with_expert.paligemma.model.vision_tower
+ _tower_eager_fwd = _tower.forward
+ _tower_compiled_fwd = torch.compile(_tower_eager_fwd, dynamic=False)
+
+ def _tower_dispatch(*args, _e=_tower_eager_fwd, _c=_tower_compiled_fwd, **kwargs):
+ if torch.is_grad_enabled():
+ return _e(*args, **kwargs)
+ return _c(*args, **kwargs)
+
+ _tower.forward = _tower_dispatch
+ logger.info("PI052: SigLIP vision tower compiled for no-grad passes.")
+
+ # Cache the fixed K-repeat action mask outside the training step.
+ if config.flow_num_repeats > 1:
+ self.register_buffer(
+ "_flow_block_diag",
+ torch.block_diag(
+ *[
+ torch.ones(
+ config.chunk_size, config.chunk_size, dtype=torch.bool, device=config.device
+ )
+ for _ in range(config.flow_num_repeats)
+ ]
+ ),
+ persistent=False,
+ )
+
+ def reset(self):
+ """Reset action and high-level inference state."""
+ super().reset()
+ self.last_subtasks = None
+ self.last_subtasks_raw = None
+ self.last_subtasks_source = None
+ self._last_good_subtasks = None
+ # Counts action chunks since the last subtask (re)generation, so the
+ # subtask can be held across several chunks (see subtask_replan_steps).
+ self._subtask_chunk_counter = 0
+
+ def apply_flashrt_fp8_mlp(self, batch: dict[str, Tensor], *, safety: float = 1.05) -> bool:
+ """Opt-in: swap every Gemma + SigLIP MLP to FlashRT fused FP8 kernels.
+
+ Calibrates static activation scales once on ``batch`` (one representative
+ observation, already through the preprocessor) and swaps the MLP modules
+ in place. Returns False (no-op, BF16 kept) if the kernels are missing.
+ Gated by ``config.use_flashrt_fp8_mlp`` — see flashrt_fp8.py.
+ """
+ from .flashrt_fp8 import apply_fp8_mlp # noqa: PLC0415
+
+ return apply_fp8_mlp(self, batch, safety=safety)
+
+ def _unfreeze_lm_head(self) -> None:
+ """Walk the PaliGemma submodules and re-enable gradients on
+ ``lm_head`` + the immediately preceding norm / last text-model
+ layer that ``PI05Policy`` typically freezes."""
+ backbone = self.model.paligemma_with_expert.paligemma
+ if hasattr(backbone, "lm_head"):
+ for p in backbone.lm_head.parameters():
+ p.requires_grad_(True)
+ # Discover terminal text layers dynamically across Transformers versions.
+ text_model = getattr(backbone, "model", None)
+ text_model = getattr(text_model, "language_model", text_model)
+ if text_model is None:
+ return
+ norm = getattr(text_model, "norm", None)
+ if norm is not None:
+ for p in norm.parameters():
+ p.requires_grad_(True)
+ layers = getattr(text_model, "layers", None)
+ if isinstance(layers, (list, torch.nn.ModuleList)) and len(layers) > 0:
+ for p in layers[-1].parameters():
+ p.requires_grad_(True)
+
+ def forward(
+ self,
+ batch: dict[str, Tensor],
+ reduction: str = "mean",
+ ) -> tuple[Tensor, dict]:
+ """Compute the enabled flow, text and FAST training losses."""
+ if reduction not in {"mean", "none"}:
+ raise ValueError(f"Unsupported loss reduction: {reduction!r}")
+ text_labels = batch.get("text_labels")
+ predict_actions_t = batch.get("predict_actions")
+
+ # Delegate only unannotated batches; PI0.5 ignores recipe action-routing masks.
+ if (
+ text_labels is None
+ and predict_actions_t is None
+ and not getattr(self.config, "enable_fast_action_loss", False)
+ ):
+ return super().forward(batch, reduction=reduction)
+
+ # Compute the host-side action-routing decision once for both flow and FAST.
+ predict_any = predict_actions_t is None or bool(predict_actions_t.any().item())
+ run_flow = self.config.flow_loss_weight > 0 and predict_any
+ run_text = self.config.text_loss_weight > 0 and text_labels is not None
+
+ loss_dict: dict[str, Any] = {}
+ total: Tensor | None = None
+
+ # Decide which losses fire this step.
+ run_fast = (
+ getattr(self.config, "enable_fast_action_loss", False)
+ and self.config.fast_action_loss_weight > 0
+ and predict_any
+ )
+ action_tokens = action_mask = action_code_mask = None
+ if run_fast:
+ from lerobot.utils.constants import ( # noqa: PLC0415
+ ACTION_CODE_TOKEN_MASK,
+ ACTION_TOKEN_MASK,
+ ACTION_TOKENS,
+ )
+
+ action_tokens = batch.get(ACTION_TOKENS)
+ action_mask = batch.get(ACTION_TOKEN_MASK)
+ action_code_mask = batch.get(ACTION_CODE_TOKEN_MASK)
+ if action_tokens is None or action_mask is None or action_code_mask is None:
+ missing = [
+ key
+ for key, value in (
+ (ACTION_TOKENS, action_tokens),
+ (ACTION_TOKEN_MASK, action_mask),
+ (ACTION_CODE_TOKEN_MASK, action_code_mask),
+ )
+ if value is None
+ ]
+ raise ValueError(
+ "PI052 FAST action loss is enabled, but the preprocessor did not produce "
+ f"required batch keys: {missing}."
+ )
+
+ # Flow uses one fused prefix/suffix pass; text-only batches skip the suffix.
+ if run_flow:
+ flow_loss, text_loss, fast_loss = self._compute_all_losses_fused(
+ batch,
+ text_labels=text_labels if run_text else None,
+ action_tokens=action_tokens if run_fast else None,
+ action_mask=action_mask if run_fast else None,
+ action_code_mask=action_code_mask if run_fast else None,
+ predict_actions_t=predict_actions_t,
+ reduction=reduction,
+ )
+ loss_dict["flow_loss"] = flow_loss.detach().mean()
+ total = self.config.flow_loss_weight * flow_loss
+ if text_loss is not None:
+ loss_dict["text_loss"] = text_loss.detach().mean()
+ total = total + self.config.text_loss_weight * text_loss
+ if fast_loss is not None:
+ loss_dict["fast_action_loss"] = fast_loss.detach().mean()
+ total = total + self.config.fast_action_loss_weight * fast_loss
+ elif run_text or run_fast:
+ text_loss, fast_loss = self._compute_text_and_fast_loss(
+ batch,
+ text_labels=text_labels if run_text else None,
+ action_tokens=action_tokens if run_fast else None,
+ action_mask=action_mask if run_fast else None,
+ action_code_mask=action_code_mask if run_fast else None,
+ predict_actions_t=predict_actions_t,
+ reduction=reduction,
+ )
+ if text_loss is not None:
+ loss_dict["text_loss"] = text_loss.detach().mean()
+ weighted = self.config.text_loss_weight * text_loss
+ total = weighted if total is None else total + weighted
+ if fast_loss is not None:
+ loss_dict["fast_action_loss"] = fast_loss.detach().mean()
+ weighted = self.config.fast_action_loss_weight * fast_loss
+ total = weighted if total is None else total + weighted
+
+ if total is None:
+ # Both flow and text disabled — make this an obvious bug
+ # rather than a silent zero loss.
+ raise RuntimeError(
+ "PI052Policy.forward: both flow_loss_weight and "
+ "text_loss_weight are 0 (or text_labels missing) — "
+ "nothing to train."
+ )
+
+ # Keep metrics detached on-device until logging to avoid extra CUDA synchronization.
+ loss_dict["loss"] = total.detach().mean()
+ return total, loss_dict
+
+ def _compute_all_losses_fused(
+ self,
+ batch: dict[str, Tensor],
+ text_labels: Tensor | None,
+ action_tokens: Tensor | None,
+ action_mask: Tensor | None,
+ action_code_mask: Tensor | None,
+ predict_actions_t: Tensor | None = None,
+ reduction: str = "mean",
+ ) -> tuple[Tensor, Tensor | None, Tensor | None]:
+ """Compute flow, text and FAST losses from one shared prefix."""
+ # ---- preamble (mirrors PI05Pytorch.forward) ------------------
+ actions = self.prepare_action(batch)
+
+ # Flow-only KI steps have no live gradient path through the prefix.
+ suppress_prefix_grads = (
+ text_labels is None
+ and action_tokens is None
+ and getattr(self.config, "knowledge_insulation", False)
+ )
+
+ # ---- prefix: images + language + (optional FAST) -------------
+ images, img_masks = self._preprocess_images(batch)
+ lang_tokens = batch[OBS_LANGUAGE_TOKENS]
+ lang_masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
+ with torch.no_grad() if suppress_prefix_grads else nullcontext():
+ prefix_embs, prefix_pad, prefix_att = self.model.embed_prefix(
+ images, img_masks, lang_tokens, lang_masks
+ )
+ non_fast_prefix_len = prefix_embs.shape[1] # images + language only
+
+ # Make supervised text causal rather than a bidirectional copy task.
+ if text_labels is not None:
+ lang_start = non_fast_prefix_len - text_labels.shape[1]
+ if lang_start >= 0:
+ prefix_att = _mark_target_span_causal(
+ prefix_att, text_labels, lang_start, non_fast_prefix_len
+ )
+
+ fast_len = 0
+ if action_tokens is not None and action_mask is not None:
+ # Gemma embedding already applies its hidden-size scale.
+ fast_emb = self.model.paligemma_with_expert.embed_language_tokens(action_tokens)
+ fast_len = action_tokens.shape[1]
+ ones_att = torch.ones(
+ (action_tokens.shape[0], fast_len),
+ dtype=torch.bool,
+ device=prefix_embs.device,
+ )
+ prefix_embs = torch.cat([prefix_embs, fast_emb], dim=1)
+ prefix_pad = torch.cat([prefix_pad, action_mask.to(prefix_pad.dtype)], dim=1)
+ prefix_att = torch.cat([prefix_att, ones_att], dim=1)
+
+ # Amortized flow reuses one VLM prefix across fresh denoising targets.
+ num_repeats = int(getattr(self.config, "flow_num_repeats", 1))
+ if num_repeats > 1:
+ prefix_out, flow_loss = self._amortized_prefix_and_flow(
+ actions,
+ prefix_embs,
+ prefix_pad,
+ prefix_att,
+ non_fast_prefix_len,
+ fast_len,
+ predict_actions_t,
+ num_repeats,
+ suppress_prefix_grads=suppress_prefix_grads,
+ reduction=reduction,
+ )
+ else:
+ prefix_out, flow_loss = self._combined_prefix_and_flow(
+ actions,
+ prefix_embs,
+ prefix_pad,
+ prefix_att,
+ non_fast_prefix_len,
+ fast_len,
+ predict_actions_t,
+ suppress_prefix_grads=suppress_prefix_grads,
+ reduction=reduction,
+ )
+
+ text_loss, fast_loss = self._prefix_ce_losses(
+ prefix_out,
+ text_labels,
+ action_tokens,
+ action_code_mask,
+ fast_len,
+ predict_actions_t,
+ reduction,
+ )
+ return flow_loss, text_loss, fast_loss
+
+ def _combined_prefix_and_flow(
+ self,
+ actions: Tensor,
+ prefix_embs: Tensor,
+ prefix_pad: Tensor,
+ prefix_att: Tensor,
+ non_fast_prefix_len: int,
+ fast_len: int,
+ predict_actions_t: Tensor | None,
+ suppress_prefix_grads: bool = False,
+ reduction: str = "mean",
+ ) -> tuple[Tensor, Tensor]:
+ """Run the single-repeat combined prefix and action path."""
+
+ noise = self.model.sample_noise(actions.shape, actions.device)
+ time = self.model.sample_time(actions.shape[0], actions.device)
+ time_expanded = time[:, None, None]
+ x_t = time_expanded * noise + (1 - time_expanded) * actions
+ u_t = noise - actions
+
+ # ---- suffix: noisy actions ----------------------------------
+ suffix_embs, suffix_pad, suffix_att, adarms_cond = self.model.embed_suffix(x_t, time)
+
+ # ---- bf16 alignment (mirrors PI05Pytorch.forward) -----------
+ first_layer = self.model.paligemma_with_expert.paligemma.model.language_model.layers[0]
+ if first_layer.self_attn.q_proj.weight.dtype == torch.bfloat16:
+ suffix_embs = suffix_embs.to(dtype=torch.bfloat16)
+ prefix_embs = prefix_embs.to(dtype=torch.bfloat16)
+
+ pad_masks = torch.cat([prefix_pad, suffix_pad], dim=1)
+ att_masks = torch.cat([prefix_att, suffix_att], dim=1)
+ att_2d_masks = make_att_2d_masks(pad_masks, att_masks)
+
+ # Block suffix-to-FAST attention to prevent trivial action leakage.
+ if fast_len > 0:
+ fast_start = non_fast_prefix_len
+ fast_end = non_fast_prefix_len + fast_len # = prefix_pad.shape[1]
+ att_2d_masks[:, fast_end:, fast_start:fast_end] = False
+
+ position_ids = torch.cumsum(pad_masks, dim=1) - 1
+ if fast_len > 0:
+ # Position flow parallel to FAST so its RoPE offsets match inference without FAST.
+ non_fast_valid = prefix_pad[:, :non_fast_prefix_len].sum(dim=1, keepdim=True)
+ suffix_pos = non_fast_valid + torch.cumsum(suffix_pad, dim=1) - 1
+ position_ids = torch.cat([position_ids[:, : prefix_pad.shape[1]], suffix_pos], dim=1)
+ att_2d_masks_4d = self.model._prepare_attention_masks_4d(att_2d_masks, dtype=prefix_embs.dtype)
+
+ # ---- forward (capture BOTH expert outputs) ------------------
+ ki_kwargs = self._ki_forward_kwargs(suppress_prefix_grads=suppress_prefix_grads)
+ (prefix_out, suffix_out), _ = self.model.paligemma_with_expert.forward(
+ attention_mask=att_2d_masks_4d,
+ position_ids=position_ids,
+ past_key_values=None,
+ inputs_embeds=[prefix_embs, suffix_embs],
+ use_cache=False,
+ adarms_cond=[None, adarms_cond],
+ **ki_kwargs,
+ )
+
+ # ---- flow loss (mirrors PI05Pytorch.forward) ----------------
+ suffix_out_slice = suffix_out[:, -self.model.config.chunk_size :].to(dtype=torch.float32)
+ v_t = self.model.action_out_proj(suffix_out_slice)
+ flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")
+ # Truncate to the actual action dimensionality (PI05 pads
+ # internally to max_action_dim).
+ original_action_dim = self.config.output_features[ACTION].shape[0]
+ flow_per_dim = flow_per_dim[:, :, :original_action_dim]
+ per_sample_flow = flow_per_dim.mean(dim=(1, 2))
+ flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
+ return prefix_out, flow_loss
+
+ def _ki_forward_kwargs(self, suppress_prefix_grads: bool = False, flex_masks=None) -> dict[str, Any]:
+ kwargs: dict[str, Any] = {}
+ if suppress_prefix_grads:
+ kwargs["suppress_prefix_grads"] = True
+ if flex_masks is not None:
+ kwargs["flex_masks"] = flex_masks
+ adarms_backend = getattr(self, "_flashrt_adarms", None)
+ if adarms_backend is not None:
+ kwargs["adarms_backend"] = adarms_backend
+ if self.config.use_manual_attention:
+ kwargs["manual_attention"] = self.config.manual_attention_scope
+ return kwargs
+
+ def _amortized_prefix_and_flow(
+ self,
+ actions: Tensor,
+ prefix_embs: Tensor,
+ prefix_pad: Tensor,
+ prefix_att: Tensor,
+ non_fast_prefix_len: int,
+ fast_len: int,
+ predict_actions_t: Tensor | None,
+ num_repeats: int,
+ suppress_prefix_grads: bool = False,
+ reduction: str = "mean",
+ ) -> tuple[Tensor, Tensor]:
+ """Run K independent action draws against one shared VLM prefix."""
+
+ model = self.model
+ k = num_repeats
+ chunk = self.config.chunk_size
+ batch_size, prefix_len = prefix_pad.shape
+
+ first_layer = model.paligemma_with_expert.paligemma.model.language_model.layers[0]
+ use_bf16 = first_layer.self_attn.q_proj.weight.dtype == torch.bfloat16
+ if use_bf16:
+ prefix_embs = prefix_embs.to(dtype=torch.bfloat16)
+
+ # ---- K suffix blocks: independent noise/time draws ----------
+ # Embed all independent K draws in one flattened batch.
+ noise = model.sample_noise((k * batch_size, *actions.shape[1:]), actions.device)
+ time = model.sample_time(k * batch_size, actions.device)
+ actions_rep = actions.repeat(k, 1, 1) # (k*B, chunk, motor_dim)
+ time_expanded = time[:, None, None]
+ x_t = time_expanded * noise + (1 - time_expanded) * actions_rep
+ u_t = (noise - actions_rep).view(k, batch_size, chunk, -1).transpose(0, 1) # (B, k, chunk, motor)
+ s_embs, suffix_pad, suffix_att, adarms = model.embed_suffix(x_t, time)
+ if use_bf16:
+ s_embs = s_embs.to(dtype=torch.bfloat16)
+ suffix_pad = suffix_pad[:batch_size]
+ suffix_att = suffix_att[:batch_size]
+ suffix_embs = (
+ s_embs.view(k, batch_size, chunk, -1).transpose(0, 1).reshape(batch_size, k * chunk, -1)
+ ) # (B, k*chunk, D)
+ # Broadcast each draw's AdaRMS condition over its action chunk.
+ adarms_cond = (
+ adarms.view(k, batch_size, 1, adarms.shape[-1])
+ .expand(k, batch_size, chunk, adarms.shape[-1])
+ .transpose(0, 1)
+ .reshape(batch_size, k * chunk, adarms.shape[-1])
+ ) # (B, k*chunk, cond_dim)
+
+ # Prefix rows cannot see action blocks; each action block sees only itself and the prefix.
+ use_flex = (
+ self.config.use_flex_attention
+ and getattr(self.config, "knowledge_insulation", False)
+ and not getattr(self, "_flex_attention_disabled", False)
+ and _get_flex_fns(prefix_pad.device) is not None
+ )
+ flex_masks = None
+ if use_flex:
+ try:
+ if not hasattr(self, "_flex_mask_builder"):
+ self._flex_mask_builder = _FlexMaskBuilder()
+ flex_masks = self._flex_mask_builder.build(
+ prefix_pad, prefix_att, non_fast_prefix_len, k, chunk
+ )
+ except Exception as exc:
+ logger.warning("PI052: FlexAttention initialization failed (%s); using SDPA.", exc)
+ self._flex_attention_disabled = True
+ if flex_masks is not None:
+ att_2d_4d = None
+ else:
+ device = prefix_pad.device
+ prefix_att_2d = make_att_2d_masks(prefix_pad, prefix_att) # (B, P, P)
+ prefix_rows = torch.cat(
+ [
+ prefix_att_2d,
+ torch.zeros(batch_size, prefix_len, k * chunk, dtype=torch.bool, device=device),
+ ],
+ dim=2,
+ )
+
+ action_to_prefix = prefix_pad[:, None, :].expand(batch_size, k * chunk, prefix_len).clone()
+ if fast_len > 0:
+ action_to_prefix[:, :, non_fast_prefix_len:prefix_len] = False
+ action_to_action = self._flow_block_diag[None].expand(batch_size, k * chunk, k * chunk)
+ action_rows = torch.cat([action_to_prefix, action_to_action], dim=2)
+
+ att_2d = torch.cat([prefix_rows, action_rows], dim=1) # (B, P + k*chunk, P + k*chunk)
+ att_2d_4d = model._prepare_attention_masks_4d(att_2d, dtype=prefix_embs.dtype)
+
+ # Restart every independent flow block after the non-FAST prefix to match inference RoPE.
+ if fast_len > 0:
+ prefix_offsets = prefix_pad[:, :non_fast_prefix_len].sum(dim=-1)[:, None]
+ else:
+ prefix_offsets = torch.sum(prefix_pad, dim=-1)[:, None]
+ block_positions = prefix_offsets + torch.cumsum(suffix_pad, dim=1) - 1 # (B, chunk)
+ position_ids = torch.cat([torch.cumsum(prefix_pad, dim=1) - 1, block_positions.repeat(1, k)], dim=1)
+
+ ki_kwargs = self._ki_forward_kwargs(suppress_prefix_grads, flex_masks)
+ (prefix_out, suffix_out), _ = model.paligemma_with_expert.forward(
+ attention_mask=att_2d_4d,
+ position_ids=position_ids,
+ past_key_values=None,
+ inputs_embeds=[prefix_embs, suffix_embs],
+ use_cache=False,
+ adarms_cond=[None, adarms_cond],
+ **ki_kwargs,
+ )
+
+ # ---- flow loss averaged over the K blocks -------------------
+ # Project all blocks together before averaging their losses.
+ original_action_dim = self.config.output_features[ACTION].shape[0]
+ v_t = model.action_out_proj(suffix_out.to(dtype=torch.float32))
+ v_t = v_t.view(batch_size, k, chunk, -1) # (B, k, chunk, motor)
+ flow_per_dim = functional.mse_loss(u_t, v_t, reduction="none")[..., :original_action_dim]
+ per_sample_flow = flow_per_dim.mean(dim=(1, 2, 3))
+ flow_loss = _reduce_action_loss(per_sample_flow, predict_actions_t, reduction)
+ return prefix_out, flow_loss
+
+ def _prefix_ce_losses(
+ self,
+ prefix_out: Tensor | None,
+ text_labels: Tensor | None,
+ action_tokens: Tensor | None,
+ action_code_mask: Tensor | None,
+ fast_len: int,
+ predict_actions_t: Tensor | None,
+ reduction: str = "mean",
+ ) -> tuple[Tensor | None, Tensor | None]:
+ """Compute enabled text and FAST losses from the shared prefix output."""
+ lm_head = self.model.paligemma_with_expert.paligemma.lm_head
+
+ text_loss: Tensor | None = None
+ if text_labels is not None and prefix_out is not None:
+ lang_len = text_labels.shape[1]
+ if fast_len > 0:
+ text_hidden = prefix_out[:, -(fast_len + lang_len) : -fast_len, :]
+ else:
+ text_hidden = prefix_out[:, -lang_len:, :]
+ # Liger avoids materializing the full vocabulary logits tensor.
+ text_loss = _shifted_lin_ce(
+ text_hidden,
+ lm_head.weight,
+ text_labels,
+ z_loss_weight=getattr(self.config, "text_ce_z_loss_weight", 0.0),
+ compiled=self.config.use_compiled_text_ce,
+ reduction=reduction,
+ )
+
+ fast_loss: Tensor | None = None
+ if fast_len > 0 and prefix_out is not None and action_code_mask is not None:
+ fast_hidden = prefix_out[:, -fast_len:, :]
+ fast_loss = _fast_lin_ce(
+ fast_hidden,
+ lm_head.weight,
+ action_tokens,
+ action_code_mask,
+ predict_actions_t,
+ compiled=self.config.use_compiled_text_ce,
+ reduction=reduction,
+ )
+
+ return text_loss, fast_loss
+
+ def _compute_text_and_fast_loss(
+ self,
+ batch: dict[str, Tensor],
+ text_labels: Tensor | None,
+ action_tokens: Tensor | None,
+ action_mask: Tensor | None,
+ action_code_mask: Tensor | None,
+ predict_actions_t: Tensor | None = None,
+ reduction: str = "mean",
+ ) -> tuple[Tensor | None, Tensor | None]:
+ """Single prefix forward → text CE + FAST CE.
+
+ Embed [images, language] (and FAST when requested) once, run
+ one backbone forward, then slice the resulting hidden states
+ at the language and FAST positions to compute both CE losses.
+ Bit-equivalent to running the two losses in separate forwards
+ because the segment-aware ``make_att_2d_masks`` keeps FAST
+ tokens invisible to language tokens, so adding FAST to the
+ prefix doesn't perturb the hidden states at language positions.
+
+ Returns ``(text_loss, fast_loss)``. Either can be ``None`` if
+ the caller doesn't want that head.
+ """
+
+ images, img_masks = self._preprocess_images(batch)
+ lang_tokens = batch[OBS_LANGUAGE_TOKENS]
+ lang_masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
+
+ prefix_embs, prefix_pad, prefix_att = self.model.embed_prefix(
+ images, img_masks, lang_tokens, lang_masks
+ )
+
+ # Make supervised text causal before appending FAST tokens.
+ if text_labels is not None:
+ lang_start = prefix_embs.shape[1] - text_labels.shape[1]
+ if lang_start >= 0:
+ prefix_att = _mark_target_span_causal(
+ prefix_att, text_labels, lang_start, prefix_embs.shape[1]
+ )
+
+ fast_len = 0
+ if action_tokens is not None and action_mask is not None:
+ # embed_language_tokens already applies the Gemma sqrt(hidden) scale (tf>=5.4.0);
+ # do not scale FAST action tokens again (would double-scale).
+ fast_emb = self.model.paligemma_with_expert.embed_language_tokens(action_tokens)
+
+ fast_len = action_tokens.shape[1]
+ ones_att = torch.ones(
+ (action_tokens.shape[0], fast_len),
+ dtype=torch.bool,
+ device=prefix_embs.device,
+ )
+ full_embs = torch.cat([prefix_embs, fast_emb], dim=1)
+ full_pad = torch.cat([prefix_pad, action_mask.to(prefix_pad.dtype)], dim=1)
+ full_att = torch.cat([prefix_att, ones_att], dim=1)
+ else:
+ full_embs = prefix_embs
+ full_pad = prefix_pad
+ full_att = prefix_att
+
+ att_2d = make_att_2d_masks(full_pad, full_att)
+ position_ids = torch.cumsum(full_pad, dim=1) - 1
+ att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=full_embs.dtype)
+
+ (vlm_out, _), _ = self.model.paligemma_with_expert.forward(
+ attention_mask=att_2d_4d,
+ position_ids=position_ids,
+ past_key_values=None,
+ inputs_embeds=[full_embs, None],
+ use_cache=False,
+ )
+ if vlm_out is None:
+ raise RuntimeError("PI052 text+fast loss: VLM forward returned no hidden states.")
+
+ lm_head = self.model.paligemma_with_expert.paligemma.lm_head
+
+ text_loss: Tensor | None = None
+ if text_labels is not None:
+ lang_len = text_labels.shape[1]
+ # embed_prefix lays out as [images, language]; with FAST
+ # appended the full sequence is [images, language, FAST].
+ if fast_len > 0:
+ text_hidden = vlm_out[:, -(fast_len + lang_len) : -fast_len, :]
+ else:
+ text_hidden = vlm_out[:, -lang_len:, :]
+ text_loss = _shifted_lin_ce(
+ text_hidden,
+ lm_head.weight,
+ text_labels,
+ z_loss_weight=getattr(self.config, "text_ce_z_loss_weight", 0.0),
+ compiled=self.config.use_compiled_text_ce,
+ reduction=reduction,
+ )
+
+ fast_loss: Tensor | None = None
+ if action_tokens is not None and action_code_mask is not None and fast_len > 0:
+ fast_hidden = vlm_out[:, -fast_len:, :]
+ fast_loss = _fast_lin_ce(
+ fast_hidden,
+ lm_head.weight,
+ action_tokens,
+ action_code_mask,
+ predict_actions_t,
+ compiled=self.config.use_compiled_text_ce,
+ reduction=reduction,
+ )
+
+ return text_loss, fast_loss
+
+ def select_message(
+ self,
+ batch: dict[str, Tensor],
+ *,
+ max_new_tokens: int = 128,
+ min_new_tokens: int = 0,
+ eos_token_id: int | None = None,
+ temperature: float = 0.0,
+ top_p: float = 1.0,
+ tokenizer: Any = None,
+ suppress_loc_tokens: bool = False,
+ use_kv_cache: bool = True,
+ ) -> str:
+ """Generate text continuation from a multimodal prefix (used by the runtime CLI).
+
+ ``suppress_loc_tokens=True`` masks PaliGemma's reserved ```` ids
+ ([256000, 257024)) before sampling — the pretraining prior drifts back to
+ them on small text-CE budgets. Pass ``True`` for subtask/memory/plan,
+ ``False`` for VQA (spatial answers legitimately emit ````).
+ """
+ self.eval()
+
+ if tokenizer is None:
+ from transformers import AutoTokenizer # noqa: PLC0415
+
+ from .inference.pi052_adapter import _get_loc_tokenizer # noqa: PLC0415
+ from .text_processor_pi052 import register_paligemma_loc_tokens # noqa: PLC0415
+
+ tok_name = getattr(self.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
+ tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
+ if eos_token_id is None:
+ eos_token_id = tokenizer.eos_token_id
+
+ special_ids: set[int] = set()
+ try:
+ for sid in tokenizer.all_special_ids or []:
+ if sid is not None:
+ special_ids.add(int(sid))
+ except Exception: # noqa: BLE001 # nosec B110
+ pass
+ if eos_token_id is not None:
+ special_ids.add(int(eos_token_id))
+
+ images, img_masks = self._preprocess_images(batch)
+ tokens = batch[OBS_LANGUAGE_TOKENS]
+ masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
+
+ prefix_embs, prefix_pad_masks, prefix_att_masks = self.model.embed_prefix(
+ images, img_masks, tokens, masks
+ )
+
+ device = prefix_embs.device
+ bsize = prefix_embs.shape[0]
+ ones_step = torch.ones((bsize, 1), dtype=torch.bool, device=device)
+
+ current_embs = prefix_embs
+ current_pad = prefix_pad_masks
+ current_att = prefix_att_masks
+ generated: list[int] = []
+ new_emb = None
+
+ # Cache the image-heavy prefix; disabling the cache retains the full-recompute parity path.
+ cache = None
+
+ backbone = self.model.paligemma_with_expert
+ lm_head = backbone.paligemma.lm_head
+
+ # Use q_proj's dtype because norms and embeddings may remain fp32 while SDPA queries are bf16.
+ backbone_dtype = backbone.paligemma.model.language_model.layers[0].self_attn.q_proj.weight.dtype
+
+ for _ in range(max_new_tokens):
+ if cache is None:
+ # Run the full bidirectional prefix initially or whenever caching is disabled.
+ step_embs = current_embs
+ att_2d = make_att_2d_masks(current_pad, current_att)
+ position_ids = torch.cumsum(current_pad, dim=1) - 1
+ att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=backbone_dtype)
+ else:
+ # Incremental decoding feeds only the last token while retaining prefix padding masks.
+ step_embs = new_emb
+ att_2d = current_pad[:, None, :]
+ att_2d_4d = self.model._prepare_attention_masks_4d(att_2d, dtype=backbone_dtype)
+ position_ids = (torch.cumsum(current_pad, dim=1) - 1)[:, -1:]
+ (vlm_out, _), new_cache = backbone.forward(
+ attention_mask=att_2d_4d,
+ position_ids=position_ids,
+ past_key_values=cache,
+ inputs_embeds=[step_embs, None],
+ use_cache=use_kv_cache,
+ )
+ if use_kv_cache:
+ cache = new_cache
+ if vlm_out is None:
+ break
+ last = vlm_out[:, -1:].to(lm_head.weight.dtype)
+ logits_step = lm_head(last)[:, -1] # (B, V)
+ if special_ids and len(generated) < min_new_tokens:
+ for sid in special_ids:
+ logits_step[..., sid] = float("-inf")
+ # Suppress FAST-only vocabulary that otherwise leaks into generated text.
+ vocab_size = logits_step.shape[-1]
+ fast_skip = int(getattr(self.config, "fast_skip_tokens", 128))
+ fast_lo = vocab_size - 1 - fast_skip - (_FAST_ACTION_VOCAB_SIZE - 1)
+ if 0 < fast_lo < 256000:
+ logits_step[..., fast_lo:256000] = float("-inf")
+ if suppress_loc_tokens:
+ logits_step[..., 256000:257024] = float("-inf")
+ next_ids = self._sample_next_token(logits_step, temperature, top_p)
+ tok_id = int(next_ids[0].item())
+ generated.append(tok_id)
+ if eos_token_id is not None and tok_id == eos_token_id:
+ break
+
+ # embed_language_tokens already applies the Gemma sqrt(hidden) scale (tf>=5.4.0).
+ new_emb = backbone.embed_language_tokens(next_ids.unsqueeze(0))
+ # Both paths track valid keys, but only recompute retains full embedding history.
+ current_pad = torch.cat([current_pad, ones_step], dim=1)
+ if not use_kv_cache:
+ current_embs = torch.cat([current_embs, new_emb], dim=1)
+ current_att = torch.cat([current_att, ones_step], dim=1)
+
+ decoded = tokenizer.decode(generated, skip_special_tokens=True).strip()
+ if not decoded and generated:
+ try:
+ self._last_select_message_debug = (
+ f"raw_ids={generated[:16]} "
+ f"decoded_w_special={tokenizer.decode(generated, skip_special_tokens=False)!r}"
+ )
+ except Exception: # noqa: BLE001
+ self._last_select_message_debug = f"raw_ids={generated[:16]}"
+ else:
+ self._last_select_message_debug = ""
+ return decoded
+
+ def _prepare_action_batch(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
+ from .inference.pi052_adapter import _build_text_batch, _get_loc_tokenizer # noqa: PLC0415
+ from .text_processor_pi052 import ( # noqa: PLC0415
+ discretize_state_str,
+ encode_prompt_with_targets,
+ register_paligemma_loc_tokens,
+ )
+
+ n = self._batch_size_from_observation(batch)
+ self._ensure_subtask_state(n)
+ tasks = self._tasks_from_batch(batch, n)
+ # Mirror training by appending the already normalized state to low-level prompts.
+ state_all = batch.get(OBS_STATE)
+
+ joint = bool(getattr(self.config, "joint_subtask_conditioning", False))
+ joint_tokenizer = None
+ if joint:
+ from transformers import AutoTokenizer # noqa: PLC0415
+
+ tok_name = getattr(self.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
+ joint_tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
+
+ # Hold subtasks for the configured interval to match training and avoid rapid replanning.
+ replan = int(getattr(self.config, "subtask_replan_steps", 0) or 0)
+ hold_chunks = max(1, round(replan / self.config.n_action_steps)) if replan > 0 else 1
+ regenerate = self._subtask_chunk_counter % hold_chunks == 0 or not any(self.last_subtasks or [])
+ self._subtask_chunk_counter += 1
+
+ # Generate and batch one independently conditioned subtask per environment.
+ rows: list[tuple[Tensor, Tensor | None, Tensor | None]] = []
+ tokenizer = None
+ for i in range(n):
+ if regenerate or not self.last_subtasks[i]:
+ obs_i = self._slice_observation(batch, i)
+ subtask = self._generate_low_level_subtask(obs_i, tasks[i], i)
+ else:
+ # Hold the previously generated subtask; only the state in the
+ # prompt below is refreshed to the current observation.
+ subtask = self.last_subtasks[i]
+ state_str = discretize_state_str(state_all[i]) if torch.is_tensor(state_all) else None
+ if joint:
+ # Joint sequences keep the task turn (with state) and render the
+ # subtask as a causal assistant turn, exactly as trained.
+ task_content = tasks[i]
+ if state_str is not None:
+ task_content = f"{task_content}, State: {state_str};"
+ ids, attn, marks = encode_prompt_with_targets(
+ joint_tokenizer,
+ [
+ {"role": "user", "content": task_content},
+ {"role": "assistant", "content": subtask},
+ ],
+ target_indices=[1],
+ )
+ device = getattr(self.config, "device", None)
+ if device is not None:
+ ids, attn, marks = ids.to(device), attn.to(device), marks.to(device)
+ rows.append((ids, attn, marks))
+ tokenizer = joint_tokenizer
+ else:
+ content = subtask if state_str is None else f"{subtask}, State: {state_str};"
+ text_batch = _build_text_batch(
+ self,
+ [{"role": "user", "content": content}],
+ add_generation_prompt=False,
+ )
+ rows.append((text_batch["lang_tokens"], text_batch["lang_masks"], None))
+ tokenizer = text_batch["tokenizer"]
+
+ tokens, masks, marks = self._stack_token_rows(rows, tokenizer)
+
+ out = dict(batch)
+ out[OBS_LANGUAGE_TOKENS] = tokens
+ out[OBS_LANGUAGE_ATTENTION_MASK] = masks
+ if marks is not None:
+ out[OBS_LANGUAGE_CAUSAL_MARKS] = marks
+ return out
+
+ def _generate_low_level_subtask(self, obs_i: dict[str, Tensor], task: str, i: int) -> str:
+ from .inference.pi052_adapter import _generate_with_policy # noqa: PLC0415
+ from .text_processor_pi052 import discretize_state_str # noqa: PLC0415
+
+ msg = ""
+ if task:
+ content = task
+ if getattr(self.config, "joint_subtask_conditioning", False):
+ # Joint samples carry state on the task turn, so the subtask
+ # must be generated from the same state-bearing prompt.
+ state = obs_i.get(OBS_STATE)
+ if torch.is_tensor(state) and state.numel() > 0:
+ state_row = state[0] if state.ndim > 1 else state
+ content = f"{task}, State: {discretize_state_str(state_row)};"
+ msg = _generate_with_policy(
+ self,
+ [{"role": "user", "content": content}],
+ observation=obs_i,
+ label=f"eval subtask gen[{i}]",
+ suppress_loc_tokens=True,
+ )
+ self.last_subtasks_raw[i] = msg or ""
+
+ # Feed the generated subtask verbatim, matching low-level training.
+ if msg:
+ subtask = " ".join(msg.strip().split())
+ self._last_good_subtasks[i] = subtask
+ self.last_subtasks[i] = subtask
+ self.last_subtasks_source[i] = "generated"
+ logger.info("PI052 eval subtask[%d]: %r (task=%r)", i, subtask, task)
+ return subtask
+
+ # Reuse the last valid subtask, or derive an initial imperative, when generation fails.
+ debug = getattr(self, "_last_select_message_debug", "") or ""
+ if not task:
+ reason = "No task string was available in the batch."
+ else:
+ reason = f"Empty generated subtask. {debug}".strip()
+ if self._last_good_subtasks[i]:
+ subtask = self._last_good_subtasks[i]
+ source = "reuse_last"
+ else:
+ subtask = self._fallback_subtask_from_task(task)
+ source = "fallback_task"
+ self.last_subtasks[i] = subtask
+ self.last_subtasks_source[i] = source
+ logger.info(
+ "PI052 eval subtask[%d] fallback (%s): %s | final=%r task=%r",
+ i,
+ source,
+ reason,
+ subtask,
+ task,
+ )
+ return subtask
+
+ def _ensure_subtask_state(self, n: int) -> None:
+ """(Re)allocate per-env subtask buffers when the env count is first seen."""
+ if self.last_subtasks is not None and len(self.last_subtasks) == n:
+ return
+ self.last_subtasks = ["" for _ in range(n)]
+ self.last_subtasks_raw = ["" for _ in range(n)]
+ self.last_subtasks_source = ["unset" for _ in range(n)]
+ self._last_good_subtasks = [None for _ in range(n)]
+
+ @staticmethod
+ def _slice_observation(batch: dict[str, Tensor], i: int) -> dict[str, Tensor]:
+ """Slice the per-env observation tensors for env ``i`` (images/state).
+
+ Language keys are excluded so high-level generation uses the freshly
+ tokenized task prompt, not the preprocessor's low-level fallback tokens.
+ """
+ out: dict[str, Tensor] = {}
+ for k, v in batch.items():
+ if not (isinstance(k, str) and k.startswith("observation.")):
+ continue
+ if k.startswith("observation.language"):
+ continue
+ if torch.is_tensor(v):
+ out[k] = v[i : i + 1]
+ return out
+
+ @staticmethod
+ def _stack_token_rows(
+ rows: list[tuple[Tensor, Tensor | None, Tensor | None]], tokenizer: Any
+ ) -> tuple[Tensor, Tensor, Tensor | None]:
+ """Right-pad per-env ``(1, L_i)`` token/mask/marks rows and stack to ``(n, L)``.
+
+ Right-padding with a False attention mask matches the training-time
+ tokenizer (``padding_side="right"``), so the action expert treats pad
+ positions as masked. Causal marks (third element, optional) pad False.
+ """
+ max_len = max(t.shape[1] for t, _, _ in rows)
+ pad_id = getattr(tokenizer, "pad_token_id", None) or 0
+ has_marks = any(m is not None for _, _, m in rows)
+ tok_rows: list[Tensor] = []
+ mask_rows: list[Tensor] = []
+ marks_rows: list[Tensor] = []
+ for tokens, masks, marks in rows:
+ length = tokens.shape[1]
+ if masks is None:
+ masks = torch.ones((1, length), dtype=torch.bool, device=tokens.device)
+ if has_marks and marks is None:
+ marks = torch.zeros((1, length), dtype=torch.bool, device=tokens.device)
+ if length < max_len:
+ pad = max_len - length
+ tokens = torch.cat(
+ [tokens, torch.full((1, pad), pad_id, dtype=tokens.dtype, device=tokens.device)],
+ dim=1,
+ )
+ masks = torch.cat(
+ [masks, torch.zeros((1, pad), dtype=masks.dtype, device=masks.device)],
+ dim=1,
+ )
+ if has_marks:
+ marks = torch.cat(
+ [marks, torch.zeros((1, pad), dtype=marks.dtype, device=marks.device)],
+ dim=1,
+ )
+ tok_rows.append(tokens)
+ mask_rows.append(masks)
+ if has_marks:
+ marks_rows.append(marks)
+ stacked_marks = torch.cat(marks_rows, dim=0) if has_marks else None
+ return torch.cat(tok_rows, dim=0), torch.cat(mask_rows, dim=0), stacked_marks
+
+ @staticmethod
+ def _fallback_subtask_from_task(task: str) -> str:
+ target = PI052Policy._navigation_target_from_task(task)
+ if target:
+ return f"go to {target}"
+ if task.lower().startswith("open the stand mixer head"):
+ return "pull stand mixer head"
+ return task
+
+ @staticmethod
+ def _navigation_target_from_task(task: str) -> str:
+ prefix = "navigate to "
+ lower = task.lower().strip()
+ if not lower.startswith(prefix):
+ return ""
+ return lower[len(prefix) :].strip().rstrip(".")
+
+ @staticmethod
+ def _tasks_from_batch(batch: dict[str, Any], n: int) -> list[str]:
+ """Return one task string per env, padded/truncated to ``n``."""
+ task = batch.get("task")
+ if isinstance(task, list):
+ raw = list(task)
+ elif task is None:
+ raw = []
+ else:
+ raw = [task]
+ tasks: list[str] = []
+ for t in raw:
+ if hasattr(t, "item"):
+ t = t.item()
+ tasks.append(t if isinstance(t, str) else "")
+ if len(tasks) < n:
+ tasks += [tasks[-1] if tasks else ""] * (n - len(tasks))
+ return tasks[:n]
+
+ @staticmethod
+ def _batch_size_from_observation(batch: dict[str, Any]) -> int:
+ state = batch.get("observation.state")
+ if torch.is_tensor(state) and state.ndim > 0:
+ return int(state.shape[0])
+ for key, value in batch.items():
+ if isinstance(key, str) and key.startswith("observation.images.") and torch.is_tensor(value):
+ return int(value.shape[0])
+ return 1
+
+ @staticmethod
+ def _sample_next_token(logits: Tensor, temperature: float, top_p: float) -> Tensor:
+ if temperature <= 0.0:
+ return logits.argmax(dim=-1)
+ scaled = logits / max(temperature, 1e-6)
+ probs = torch.softmax(scaled, dim=-1)
+ if top_p < 1.0:
+ sorted_p, sorted_ix = torch.sort(probs, descending=True, dim=-1)
+ cum = torch.cumsum(sorted_p, dim=-1)
+ mask = cum > top_p
+ mask[..., 0] = False
+ sorted_p = sorted_p.masked_fill(mask, 0.0)
+ sorted_p = sorted_p / sorted_p.sum(dim=-1, keepdim=True).clamp_min(1e-8)
+ choice = torch.multinomial(sorted_p, num_samples=1)
+ return sorted_ix.gather(-1, choice).squeeze(-1)
+ return torch.multinomial(probs, num_samples=1).squeeze(-1)
+
+ def _prepare_pretrained_state_dict(self, remapped_state_dict: dict[str, Tensor]) -> dict[str, Tensor]:
+ lm_head_key = "model.paligemma_with_expert.paligemma.lm_head.weight"
+ embed_tokens_key = "model.paligemma_with_expert.paligemma.model.language_model.embed_tokens.weight"
+ if lm_head_key not in remapped_state_dict and embed_tokens_key in remapped_state_dict:
+ remapped_state_dict[lm_head_key] = remapped_state_dict[embed_tokens_key].clone().float()
+ elif lm_head_key in remapped_state_dict:
+ remapped_state_dict[lm_head_key] = remapped_state_dict[lm_head_key].float()
+ return remapped_state_dict
+
+ def get_optim_params(self):
+ """Return policy parameters, optionally split into LR-scaled groups.
+
+ Three orthogonal multipliers scale the base ``optimizer_lr``:
+ ``lm_head_lr_scale`` (PaliGemma ``lm_head`` + tied ``embed_tokens``),
+ ``backbone_lr_scale`` (the rest of the PaliGemma tower), and
+ ``action_expert_lr_scale`` (the Gemma expert + action/time projection
+ heads). The cosine scheduler multiplies every group by the same lambda
+ each step so the ratios are preserved across decay. When all three are
+ ``1.0`` this returns ``self.parameters()`` (back-compat with existing
+ checkpoints and configs).
+ """
+ head_scale = float(getattr(self.config, "lm_head_lr_scale", 1.0))
+ backbone_scale = float(getattr(self.config, "backbone_lr_scale", 1.0))
+ expert_scale = float(getattr(self.config, "action_expert_lr_scale", 1.0))
+ if head_scale == 1.0 and backbone_scale == 1.0 and expert_scale == 1.0:
+ return self.parameters()
+
+ # Keep the tied LM projection and embeddings in the same optimizer group.
+ head_substrings = (
+ "paligemma_with_expert.paligemma.lm_head.",
+ "paligemma_with_expert.paligemma.model.language_model.embed_tokens.",
+ )
+ backbone_substring = "paligemma_with_expert.paligemma."
+ head_params: list[torch.nn.Parameter] = []
+ backbone_params: list[torch.nn.Parameter] = []
+ expert_params: list[torch.nn.Parameter] = []
+ for name, p in self.named_parameters():
+ if not p.requires_grad:
+ continue
+ if any(s in name for s in head_substrings):
+ head_params.append(p)
+ elif backbone_substring in name:
+ backbone_params.append(p)
+ else:
+ expert_params.append(p)
+ base_lr = float(self.config.optimizer_lr)
+ groups: list[dict[str, object]] = []
+ if backbone_params:
+ groups.append({"params": backbone_params, "lr": base_lr * backbone_scale, "name": "backbone"})
+ if expert_params:
+ groups.append({"params": expert_params, "lr": base_lr * expert_scale, "name": "action_expert"})
+ if head_params:
+ groups.append({"params": head_params, "lr": base_lr * head_scale, "name": "lm_head"})
+ # Sanity: a non-trivial head scale that matches no params would silently
+ # do nothing — surface that fast.
+ if head_scale != 1.0 and not head_params:
+ raise RuntimeError(
+ "lm_head_lr_scale != 1.0 but no parameters matched the LM-head "
+ f"name patterns: {head_substrings!r}. Did the underlying PaliGemma "
+ "module rename?"
+ )
+ logging.info(
+ "PI052Policy LR groups (base=%.3g): backbone=%.3g (×%.3g, n=%d), "
+ "action_expert=%.3g (×%.3g, n=%d), lm_head=%.3g (×%.3g, n=%d)",
+ base_lr,
+ base_lr * backbone_scale,
+ backbone_scale,
+ len(backbone_params),
+ base_lr * expert_scale,
+ expert_scale,
+ len(expert_params),
+ base_lr * head_scale,
+ head_scale,
+ len(head_params),
+ )
+ return groups
+
+ @torch.no_grad()
+ def predict_action_chunk(self, batch: dict[str, Tensor], **kwargs: Unpack[ActionSelectKwargs]) -> Tensor:
+ # Guard before first-observation FP8 calibration to prevent recursive prediction.
+ if self.config.use_flashrt_fp8_mlp and not getattr(self, "_fp8_applied", False):
+ self._fp8_applied = True
+ self.apply_flashrt_fp8_mlp(batch)
+ marks = batch.get(OBS_LANGUAGE_CAUSAL_MARKS)
+ if marks is None:
+ return super().predict_action_chunk(batch, **kwargs)
+ return self._predict_action_chunk_with_marks(batch, marks, **kwargs)
+
+ @torch.no_grad()
+ def _predict_action_chunk_with_marks(
+ self, batch: dict[str, Tensor], marks: Tensor, **kwargs: Unpack[ActionSelectKwargs]
+ ) -> Tensor:
+ """Base ``predict_action_chunk`` plus causal marks on the generated-subtask span."""
+ self.eval()
+ images, img_masks = self._preprocess_images(batch)
+ tokens = batch[OBS_LANGUAGE_TOKENS]
+ masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
+ actions = self.model.sample_actions(
+ images, img_masks, tokens, masks, lang_causal_marks=marks, **kwargs
+ )
+ original_action_dim = self.config.output_features[ACTION].shape[0]
+ return actions[:, :, :original_action_dim]
diff --git a/src/lerobot/policies/pi052/processor_pi052.py b/src/lerobot/policies/pi052/processor_pi052.py
new file mode 100644
index 000000000..8202346e5
--- /dev/null
+++ b/src/lerobot/policies/pi052/processor_pi052.py
@@ -0,0 +1,164 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""PI052 processor factory with optional recipe rendering and text tokenization.
+
+Without a recipe it delegates to the standard PI0.5 pipeline.
+"""
+
+from __future__ import annotations
+
+from pathlib import Path
+from typing import Any
+
+import torch
+
+from lerobot.configs.recipe import TrainingRecipe
+from lerobot.processor import (
+ AbsoluteActionsProcessorStep,
+ ActionTokenizerProcessorStep,
+ AddBatchDimensionProcessorStep,
+ DeviceProcessorStep,
+ NormalizerProcessorStep,
+ PolicyAction,
+ PolicyProcessorPipeline,
+ RelativeActionsProcessorStep,
+ RenameObservationsProcessorStep,
+ UnnormalizerProcessorStep,
+ policy_action_to_transition,
+ transition_to_policy_action,
+)
+
+# Import directly to keep optional language dependencies out of ``lerobot.processor``.
+from lerobot.processor.render_messages_processor import RenderMessagesStep
+from lerobot.utils.constants import POLICY_POSTPROCESSOR_DEFAULT_NAME, POLICY_PREPROCESSOR_DEFAULT_NAME
+
+from ..pi05.processor_pi05 import make_pi05_pre_post_processors
+from .configuration_pi052 import PI052Config
+from .text_processor_pi052 import PI052TextTokenizerStep
+
+
+def make_pi052_pre_post_processors(
+ config: PI052Config,
+ dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
+ dataset_repo_id: str | None = None,
+ dataset_root: str | None = None,
+ dataset_revision: str | None = None,
+ episodes: list[int] | None = None,
+ exclude_episodes: list[int] | None = None,
+) -> tuple[
+ PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
+ PolicyProcessorPipeline[PolicyAction, PolicyAction],
+]:
+ """Build PI0.5-v2's pre/post-processor pipelines.
+
+ Falls through to π0.5's stock pipeline when ``recipe_path`` is unset.
+ """
+ if not config.recipe_path:
+ if getattr(config, "enable_fast_action_loss", False):
+ raise ValueError("PI052 FAST action loss requires recipe_path to build action supervision.")
+ return make_pi05_pre_post_processors(config, dataset_stats=dataset_stats)
+
+ recipe = _load_recipe(config.recipe_path)
+
+ relative_step = RelativeActionsProcessorStep(
+ enabled=config.use_relative_actions,
+ exclude_joints=getattr(config, "relative_exclude_joints", []),
+ action_names=getattr(config, "action_feature_names", None),
+ )
+
+ input_steps = [
+ RenameObservationsProcessorStep(rename_map={}),
+ AddBatchDimensionProcessorStep(),
+ relative_step,
+ NormalizerProcessorStep(
+ features={**config.input_features, **config.output_features},
+ norm_map=config.normalization_mapping,
+ stats=dataset_stats,
+ ),
+ RenderMessagesStep(recipe=recipe),
+ PI052TextTokenizerStep(
+ tokenizer_name="google/paligemma-3b-pt-224",
+ max_length=config.tokenizer_max_length,
+ plan_dropout_prob=getattr(config, "plan_dropout_prob", 0.0),
+ memory_dropout_prob=getattr(config, "memory_dropout_prob", 0.0),
+ subtask_dropout_prob=getattr(config, "subtask_dropout_prob", 0.0),
+ ),
+ ]
+
+ # Add FAST action-token supervision only when explicitly enabled.
+ if getattr(config, "enable_fast_action_loss", False):
+ from .fit_fast_tokenizer import resolve_fast_tokenizer # noqa: PLC0415
+
+ input_steps.append(
+ ActionTokenizerProcessorStep(
+ action_tokenizer_name=resolve_fast_tokenizer(
+ config,
+ dataset_repo_id,
+ dataset_root,
+ dataset_stats,
+ dataset_revision,
+ episodes,
+ exclude_episodes,
+ ),
+ max_action_tokens=config.max_action_tokens,
+ fast_skip_tokens=config.fast_skip_tokens,
+ paligemma_tokenizer_name="google/paligemma-3b-pt-224",
+ allow_truncation=False,
+ )
+ )
+
+ input_steps.append(DeviceProcessorStep(device=config.device))
+
+ output_steps = [
+ UnnormalizerProcessorStep(
+ features=config.output_features,
+ norm_map=config.normalization_mapping,
+ stats=dataset_stats,
+ ),
+ AbsoluteActionsProcessorStep(
+ enabled=config.use_relative_actions,
+ relative_step=relative_step,
+ ),
+ DeviceProcessorStep(device="cpu"),
+ ]
+ return (
+ PolicyProcessorPipeline[dict[str, Any], dict[str, Any]](
+ steps=input_steps,
+ name=POLICY_PREPROCESSOR_DEFAULT_NAME,
+ ),
+ PolicyProcessorPipeline[PolicyAction, PolicyAction](
+ steps=output_steps,
+ name=POLICY_POSTPROCESSOR_DEFAULT_NAME,
+ to_transition=policy_action_to_transition,
+ to_output=transition_to_policy_action,
+ ),
+ )
+
+
+def _load_recipe(path_str: str) -> TrainingRecipe:
+ """Resolve ``path_str`` to a ``TrainingRecipe``.
+
+ Accepts an absolute path or a path relative to
+ ``src/lerobot/configs/``.
+ """
+ p = Path(path_str)
+ if not p.is_absolute() and not p.exists():
+ from lerobot.configs import recipe as _recipe_module # noqa: PLC0415
+
+ configs_dir = Path(_recipe_module.__file__).resolve().parent
+ candidate = configs_dir / path_str
+ if candidate.exists():
+ p = candidate
+ return TrainingRecipe.from_yaml(p)
diff --git a/src/lerobot/policies/pi052/text_processor_pi052.py b/src/lerobot/policies/pi052/text_processor_pi052.py
new file mode 100644
index 000000000..9b6dd8bb2
--- /dev/null
+++ b/src/lerobot/policies/pi052/text_processor_pi052.py
@@ -0,0 +1,521 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tokenize PI052 messages and build text/action supervision masks."""
+
+from __future__ import annotations
+
+import json
+import logging
+from dataclasses import dataclass
+from typing import Any
+
+import numpy as np
+import torch
+from torch import Tensor
+
+from lerobot.configs import PipelineFeatureType, PolicyFeature
+from lerobot.processor.pipeline import ProcessorStep, ProcessorStepRegistry
+from lerobot.types import EnvTransition, TransitionKey
+from lerobot.utils.constants import OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE
+
+logger = logging.getLogger(__name__)
+
+
+def discretize_state_str(state_row: Any) -> str:
+ """Format one normalized state row with PI0.5's 256-bin convention."""
+ arr = state_row.detach().cpu().numpy() if hasattr(state_row, "detach") else np.asarray(state_row)
+ disc = np.digitize(arr, bins=np.linspace(-1, 1, 256 + 1)[:-1]) - 1
+ return " ".join(str(int(x)) for x in disc.reshape(-1).tolist())
+
+
+def _state_row_at(state_all: Any, pos: int) -> Any:
+ """Select the per-sample state row from a (possibly batched) state tensor."""
+ if state_all is None:
+ return None
+ if hasattr(state_all, "ndim") and state_all.ndim >= 2:
+ return state_all[pos]
+ return state_all
+
+
+def _content_to_text(content: Any) -> str:
+ """Collapse a message's ``content`` (string or multimodal blocks) to text."""
+ if isinstance(content, str):
+ return content
+ if isinstance(content, list):
+ parts = [
+ b["text"]
+ for b in content
+ if isinstance(b, dict) and b.get("type") == "text" and isinstance(b.get("text"), str)
+ ]
+ return "\n".join(parts)
+ return ""
+
+
+def _flatten_say_tool_calls(message: dict[str, Any]) -> dict[str, Any]:
+ """Move ``say`` tool calls into text markers that PaliGemma can learn."""
+ tool_calls = message.get("tool_calls")
+ if not tool_calls:
+ return message
+ say_texts: list[str] = []
+ for call in tool_calls:
+ if not isinstance(call, dict):
+ continue
+ fn = call.get("function") or {}
+ if fn.get("name") != "say":
+ continue
+ args = fn.get("arguments")
+ if isinstance(args, str):
+ try:
+ import json # noqa: PLC0415
+
+ args = json.loads(args)
+ except (ValueError, TypeError):
+ args = {}
+ text = args.get("text", "") if isinstance(args, dict) else ""
+ if text:
+ say_texts.append(str(text))
+ new = dict(message)
+ new.pop("tool_calls", None)
+ if not say_texts:
+ return new
+ base = _content_to_text(new.get("content")).strip()
+ marker = "".join(f"{t}" for t in say_texts)
+ new["content"] = f"{base}\n{marker}" if base else marker
+ return new
+
+
+def _strip_blocks(message: dict[str, Any]) -> dict[str, Any]:
+ """Flatten text blocks and drop image blocks handled by observation inputs."""
+ new = dict(message)
+ new.pop("stream", None)
+ new.pop("target", None)
+ content = new.get("content")
+ if content is None:
+ new["content"] = ""
+ elif isinstance(content, str):
+ pass
+ elif isinstance(content, list):
+ parts: list[str] = []
+ for block in content:
+ if not isinstance(block, dict):
+ continue
+ if block.get("type") == "text":
+ t = block.get("text", "")
+ if isinstance(t, str):
+ parts.append(t)
+ new["content"] = "\n".join(parts)
+ else:
+ new["content"] = str(content)
+ return new
+
+
+def _is_batched_messages(messages: Any) -> bool:
+ return isinstance(messages, list) and bool(messages) and isinstance(messages[0], list)
+
+
+def _sample_indices(value: Any, batch_size: int) -> list[int | None]:
+ if value is None:
+ return [None] * batch_size
+ if isinstance(value, torch.Tensor):
+ if value.numel() == 1:
+ return [int(value.item())] * batch_size
+ values = value.reshape(-1).tolist()
+ return [int(v) for v in values[:batch_size]]
+ if isinstance(value, (list, tuple)):
+ if len(value) == 1:
+ return _sample_indices(value[0], batch_size)
+ return [int(v.item() if hasattr(v, "item") else v) for v in value[:batch_size]]
+ return [int(value)] * batch_size
+
+
+_VQA_COORD_SCALE = 1000.0
+
+
+def register_paligemma_loc_tokens(tokenizer: Any) -> Any:
+ """Register PaliGemma's reserved ```` strings as single tokens.
+
+ Without registration, the stock tokenizer splits each location into generic text pieces.
+ """
+ if "" in getattr(tokenizer, "added_tokens_encoder", {}):
+ return tokenizer
+ tokenizer.add_tokens([f"" for i in range(1024)])
+ return tokenizer
+
+
+def _loc_token(coord: float, scale: float = _VQA_COORD_SCALE) -> str:
+ """PaliGemma ```` for a coord on a ``[0, scale]`` axis."""
+ idx = round(float(coord) / scale * 1023) if scale > 0 else 0
+ return f""
+
+
+def _vqa_answer_to_loc(answer: dict[str, Any]) -> str | None:
+ """Convert normalized bbox/keypoint answers to label-first PaliGemma locations.
+
+ Label-first targets prevent location tokens from dominating every assistant turn; non-spatial answers return ``None``.
+ """
+ point = answer.get("point")
+ if isinstance(point, list | tuple) and len(point) == 2 and "point_format" in answer:
+ try:
+ x, y = float(point[0]), float(point[1])
+ except (TypeError, ValueError):
+ return None
+ label = str(answer.get("label", "")).strip()
+ if not label:
+ return None
+ return f"{label} {_loc_token(y)}{_loc_token(x)}"
+
+ detections = answer.get("detections")
+ if isinstance(detections, list) and detections:
+ parts: list[str] = []
+ for det in detections:
+ if not isinstance(det, dict):
+ continue
+ box = det.get("bbox")
+ if not (isinstance(box, list | tuple) and len(box) == 4):
+ continue
+ try:
+ x1, y1, x2, y2 = (float(v) for v in box)
+ except (TypeError, ValueError):
+ continue
+ label = str(det.get("label", "")).strip()
+ if not label:
+ continue
+ toks = f"{_loc_token(y1)}{_loc_token(x1)}{_loc_token(y2)}{_loc_token(x2)}"
+ parts.append(f"{label} {toks}")
+ return " ; ".join(parts) if parts else None
+ return None
+
+
+def _messages_vqa_to_loc(
+ messages: list[dict[str, Any]],
+ target_indices: list[int],
+) -> list[dict[str, Any]]:
+ """Rewrite spatial VQA target JSON as camera-independent ```` text."""
+ if not target_indices:
+ return messages
+ out = list(messages)
+ for idx in target_indices:
+ if not (0 <= idx < len(out)):
+ continue
+ content = out[idx].get("content")
+ if not isinstance(content, str) or not content.strip():
+ continue
+ try:
+ answer = json.loads(content)
+ except (ValueError, TypeError):
+ continue
+ if not isinstance(answer, dict):
+ continue
+ loc_text = _vqa_answer_to_loc(answer)
+ if loc_text is not None:
+ out[idx] = {**out[idx], "content": loc_text}
+ return out
+
+
+def _format_messages(
+ messages: list[dict[str, Any]],
+ target_indices: list[int] | None = None,
+ eos_token: str | None = None,
+) -> tuple[str, list[tuple[int, int]]]:
+ """Build the flat PI0.5 prompt and each message's payload span.
+
+ Supervised targets include EOS so generation learns when to stop.
+ """
+ targets = set(target_indices or [])
+ parts: list[str] = []
+ spans: list[tuple[int, int]] = []
+ cursor = 0
+ for i, m in enumerate(messages):
+ role = m.get("role", "user")
+ content = m.get("content", "") or ""
+ header = f"{role.capitalize()}: "
+ body = content + eos_token if (eos_token and i in targets) else content
+ full = header + body + "\n"
+ start = cursor + len(header)
+ end = start + len(body)
+ parts.append(full)
+ spans.append((start, end))
+ cursor += len(full)
+ return "".join(parts), spans
+
+
+def encode_prompt_with_targets(
+ tokenizer: Any, messages: list[dict[str, Any]], target_indices: list[int]
+) -> tuple[Tensor, Tensor, Tensor]:
+ """Tokenize a flat prompt and mark the token positions of target spans.
+
+ Inference-side twin of ``PI052TextTokenizerStep._encode_messages``: same
+ serialization (role headers, target EOS) and the same offset-overlap span
+ arithmetic, but unpadded and returning a boolean target mask instead of
+ labels. Used to rebuild joint-sequence prompts whose target spans must be
+ attended causally, matching ``_mark_target_span_causal`` at train time.
+
+ Returns ``(input_ids, attention_mask, target_marks)``, each ``(1, L)``.
+ """
+ prompt, spans = _format_messages(messages, target_indices, getattr(tokenizer, "eos_token", None))
+ encoded = tokenizer(prompt, return_tensors="pt", return_offsets_mapping=True)
+ input_ids = encoded["input_ids"][0]
+ attention_mask = encoded.get("attention_mask")
+ if attention_mask is None:
+ attention_mask = torch.ones_like(input_ids, dtype=torch.bool)
+ else:
+ attention_mask = attention_mask[0].bool()
+ offsets = encoded["offset_mapping"][0]
+
+ marks = torch.zeros_like(input_ids, dtype=torch.bool)
+ for idx in target_indices:
+ if idx >= len(spans):
+ continue
+ char_start, char_end = spans[idx]
+ for token_pos in range(input_ids.shape[0]):
+ if not attention_mask[token_pos]:
+ continue
+ tok_start, tok_end = int(offsets[token_pos, 0]), int(offsets[token_pos, 1])
+ if tok_end <= char_start or tok_start >= char_end:
+ continue
+ marks[token_pos] = True
+ return input_ids.unsqueeze(0), attention_mask.unsqueeze(0), marks.unsqueeze(0)
+
+
+@dataclass
+@ProcessorStepRegistry.register(name="pi052_text_tokenizer")
+class PI052TextTokenizerStep(ProcessorStep):
+ """Convert flat role-delimited messages into tokens and supervision masks."""
+
+ tokenizer_name: str = "google/paligemma-3b-pt-224"
+ max_length: int = 200
+ padding: str = "max_length"
+ padding_side: str = "right"
+ plan_dropout_prob: float = 0.0
+ memory_dropout_prob: float = 0.0
+ subtask_dropout_prob: float = 0.0
+ interjection_dropout_prob: float = 0.0
+ dropout_seed: int | None = None
+
+ def __post_init__(self) -> None:
+ self._tokenizer: Any = None
+
+ def get_config(self) -> dict[str, Any]:
+ return {
+ "tokenizer_name": self.tokenizer_name,
+ "max_length": self.max_length,
+ "padding": self.padding,
+ "padding_side": self.padding_side,
+ "plan_dropout_prob": self.plan_dropout_prob,
+ "memory_dropout_prob": self.memory_dropout_prob,
+ "subtask_dropout_prob": self.subtask_dropout_prob,
+ "interjection_dropout_prob": self.interjection_dropout_prob,
+ "dropout_seed": self.dropout_seed,
+ }
+
+ def _ensure_tokenizer(self) -> Any:
+ if self._tokenizer is not None:
+ return self._tokenizer
+ from transformers import AutoTokenizer # noqa: PLC0415
+
+ self._tokenizer = register_paligemma_loc_tokens(AutoTokenizer.from_pretrained(self.tokenizer_name))
+ return self._tokenizer
+
+ def __call__(self, transition: EnvTransition) -> EnvTransition | None:
+ transition = transition.copy()
+ complementary = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) or {}
+ messages = complementary.get("messages") or []
+
+ if not messages:
+ return transition
+
+ tokenizer = self._ensure_tokenizer()
+ state_all = (transition.get(TransitionKey.OBSERVATION) or {}).get(OBS_STATE)
+ if _is_batched_messages(messages):
+ indices_iter = _sample_indices(complementary.get("index"), len(messages))
+ encoded = [
+ self._encode_messages(
+ tokenizer,
+ msg,
+ list(streams),
+ list(tgt_indices),
+ complementary,
+ sample_idx=int(s_idx) if s_idx is not None else None,
+ state_row=_state_row_at(state_all, pos),
+ )
+ for pos, (msg, streams, tgt_indices, s_idx) in enumerate(
+ zip(
+ messages,
+ complementary.get("message_streams") or [[] for _ in messages],
+ complementary.get("target_message_indices") or [[] for _ in messages],
+ indices_iter,
+ strict=False,
+ )
+ )
+ ]
+ else:
+ sample_idx = _sample_indices(complementary.get("index"), 1)[0]
+ encoded = [
+ self._encode_messages(
+ tokenizer,
+ messages,
+ list(complementary.get("message_streams") or []),
+ list(complementary.get("target_message_indices") or []),
+ complementary,
+ sample_idx=sample_idx,
+ state_row=_state_row_at(state_all, 0),
+ )
+ ]
+
+ obs = dict(transition.get(TransitionKey.OBSERVATION) or {})
+ obs[OBS_LANGUAGE_TOKENS] = torch.stack([ids for ids, _, _, _, _ in encoded])
+ obs[OBS_LANGUAGE_ATTENTION_MASK] = torch.stack([attn for _, attn, _, _, _ in encoded])
+ transition[TransitionKey.OBSERVATION] = obs
+
+ transition[TransitionKey.COMPLEMENTARY_DATA] = {
+ **complementary,
+ "text_labels": torch.stack([labels for _, _, labels, _, _ in encoded]),
+ "predict_actions": torch.stack([pred for _, _, _, pred, _ in encoded]),
+ }
+ return transition
+
+ def _encode_messages(
+ self,
+ tokenizer: Any,
+ messages: list[dict[str, Any]],
+ message_streams: list[str | None],
+ target_indices: list[int],
+ complementary: dict[str, Any],
+ sample_idx: int | None = None,
+ state_row: Any = None,
+ ) -> tuple[Tensor, Tensor, Tensor, Tensor, str]:
+ if (
+ self.plan_dropout_prob
+ or self.memory_dropout_prob
+ or self.subtask_dropout_prob
+ or self.interjection_dropout_prob
+ ):
+ messages, target_indices = self._apply_prompt_dropout(
+ messages,
+ target_indices,
+ complementary,
+ sample_idx=sample_idx,
+ )
+
+ messages = _messages_vqa_to_loc(messages, target_indices)
+
+ messages = [_strip_blocks(_flatten_say_tool_calls(m)) for m in messages]
+ # Only low-level prompts carry PI0.5-style proprioception.
+ if state_row is not None and any(s == "low_level" for s in message_streams):
+ state_str = discretize_state_str(state_row)
+ for m in reversed(messages):
+ if m.get("role") == "user":
+ base = _content_to_text(m.get("content", ""))
+ m["content"] = f"{base}, State: {state_str};"
+ break
+ prompt, spans = _format_messages(messages, target_indices, getattr(tokenizer, "eos_token", None))
+
+ encoded = tokenizer(
+ prompt,
+ max_length=self.max_length,
+ padding=self.padding,
+ truncation=True,
+ return_tensors="pt",
+ return_offsets_mapping=True,
+ padding_side=self.padding_side,
+ )
+
+ input_ids = encoded["input_ids"][0]
+ attention_mask = encoded["attention_mask"][0].bool()
+ offsets = encoded["offset_mapping"][0]
+
+ labels = torch.full_like(input_ids, fill_value=-100)
+ for idx in target_indices:
+ if idx >= len(spans):
+ continue
+ char_start, char_end = spans[idx]
+ for token_pos in range(input_ids.shape[0]):
+ if not attention_mask[token_pos]:
+ continue
+ tok_start, tok_end = int(offsets[token_pos, 0]), int(offsets[token_pos, 1])
+ if tok_end <= char_start or tok_start >= char_end:
+ continue
+ labels[token_pos] = input_ids[token_pos]
+
+ predict_actions = torch.tensor(
+ bool(any(s == "low_level" for s in message_streams)),
+ dtype=torch.bool,
+ )
+ return input_ids, attention_mask, labels, predict_actions, prompt
+
+ def _apply_prompt_dropout(
+ self,
+ messages: list[dict[str, Any]],
+ target_indices: list[int],
+ complementary: dict[str, Any],
+ sample_idx: int | None = None,
+ ) -> tuple[list[dict[str, Any]], list[int]]:
+ """Drop sampled context messages and remap the retained target positions."""
+ import random # noqa: PLC0415
+
+ seed = self.dropout_seed
+ if seed is None:
+ seed_src = sample_idx if sample_idx is not None else complementary.get("index", 0)
+ try:
+ if hasattr(seed_src, "item"):
+ seed_src = seed_src.item()
+ seed = int(seed_src)
+ except (TypeError, ValueError):
+ seed = 0
+ rng = random.Random(seed)
+
+ keep_indices: list[int] = []
+ for idx, msg in enumerate(messages):
+ if idx in target_indices:
+ keep_indices.append(idx)
+ continue
+ kind = _classify_for_dropout(msg)
+ prob = {
+ "plan": self.plan_dropout_prob,
+ "memory": self.memory_dropout_prob,
+ "subtask": self.subtask_dropout_prob,
+ "interjection": self.interjection_dropout_prob,
+ }.get(kind, 0.0)
+ if prob > 0.0 and rng.random() < prob:
+ continue
+ keep_indices.append(idx)
+
+ new_messages = [messages[i] for i in keep_indices]
+ old_to_new = {old: new for new, old in enumerate(keep_indices)}
+ new_targets = [old_to_new[t] for t in target_indices if t in old_to_new]
+ return new_messages, new_targets
+
+ def transform_features(
+ self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
+ ) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
+ return features
+
+
+def _classify_for_dropout(message: dict[str, Any]) -> str | None:
+ """Classify context from its rendered text prefix."""
+ content = message.get("content")
+ if isinstance(content, list):
+ text_parts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"]
+ content = " ".join(text_parts)
+ elif content is None or not isinstance(content, str):
+ return None
+ s = content.strip()
+ if s.startswith("Plan:") or s.startswith("Previous plan"):
+ return "plan"
+ if s.startswith("Memory:") or s.startswith("Previous memory"):
+ return "memory"
+ if s.startswith("Current subtask") or s.startswith("Completed subtask"):
+ return "subtask"
+ return None
diff --git a/src/lerobot/policies/pi0_fast/configuration_pi0_fast.py b/src/lerobot/policies/pi0_fast/configuration_pi0_fast.py
index e5c6851f4..492f9b7b9 100644
--- a/src/lerobot/policies/pi0_fast/configuration_pi0_fast.py
+++ b/src/lerobot/policies/pi0_fast/configuration_pi0_fast.py
@@ -61,6 +61,9 @@ class PI0FastConfig(PreTrainedConfig):
tokenizer_max_length: int = 200 # see openpi `__post_init__`
text_tokenizer_name: str = "google/paligemma-3b-pt-224"
action_tokenizer_name: str = "lerobot/fast-action-tokenizer"
+ auto_fit_fast_tokenizer: bool = False
+ fast_tokenizer_cache_dir: str = "~/.cache/lerobot/fast_tokenizers"
+ fast_tokenizer_fit_samples: int = 1024
temperature: float = 0.0
max_decoding_steps: int = 256
fast_skip_tokens: int = 128
diff --git a/src/lerobot/policies/pi0_fast/processor_pi0_fast.py b/src/lerobot/policies/pi0_fast/processor_pi0_fast.py
index fdf759e29..20a1d6745 100644
--- a/src/lerobot/policies/pi0_fast/processor_pi0_fast.py
+++ b/src/lerobot/policies/pi0_fast/processor_pi0_fast.py
@@ -92,6 +92,11 @@ class Pi0FastPrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
def make_pi0_fast_pre_post_processors(
config: PI0FastConfig,
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
+ dataset_repo_id: str | None = None,
+ dataset_root: str | None = None,
+ dataset_revision: str | None = None,
+ episodes: list[int] | None = None,
+ exclude_episodes: list[int] | None = None,
) -> tuple[
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
PolicyProcessorPipeline[PolicyAction, PolicyAction],
@@ -136,6 +141,18 @@ def make_pi0_fast_pre_post_processors(
# state from the observation but does not change it. NormalizerProcessorStep still runs
# before Pi0FastPrepareStateAndLanguageTokenizerProcessorStep, so the state tokenizer
# continues to receive normalized state in [-1, 1] as expected.
+ from ..pi052.fit_fast_tokenizer import resolve_fast_tokenizer # noqa: PLC0415
+
+ action_tokenizer_path = resolve_fast_tokenizer(
+ config,
+ dataset_repo_id,
+ dataset_root,
+ dataset_stats,
+ dataset_revision,
+ episodes,
+ exclude_episodes,
+ )
+
input_steps: list[ProcessorStep] = [
steps.rename_observations, # To mimic the same processor as pretrained one
steps.add_batch_dim,
@@ -149,7 +166,7 @@ def make_pi0_fast_pre_post_processors(
padding="max_length",
),
ActionTokenizerProcessorStep(
- action_tokenizer_name=config.action_tokenizer_name,
+ action_tokenizer_name=action_tokenizer_path,
max_action_tokens=config.max_action_tokens,
fast_skip_tokens=config.fast_skip_tokens,
paligemma_tokenizer_name=config.text_tokenizer_name,
diff --git a/src/lerobot/policies/pi_gemma.py b/src/lerobot/policies/pi_gemma.py
index 9986f9b79..dcaf145dd 100644
--- a/src/lerobot/policies/pi_gemma.py
+++ b/src/lerobot/policies/pi_gemma.py
@@ -14,18 +14,27 @@
from __future__ import annotations
-from typing import TYPE_CHECKING
+from typing import TYPE_CHECKING, Literal
import torch
from torch import nn
+from torch.nn import functional as F # noqa: N812
from lerobot.utils.import_utils import _transformers_available
+# Default PaliGemma SigLIP input resolution. Mirrors
+# ``pi05.configuration_pi05.DEFAULT_IMAGE_SIZE``; duplicated as a plain constant
+# to avoid importing the pi05 package here (which would create an import cycle:
+# pi_gemma -> pi05.__init__ -> modeling_pi05 -> pi_gemma).
+DEFAULT_IMAGE_SIZE = 224
+
if TYPE_CHECKING or _transformers_available:
from transformers.cache_utils import DynamicCache
from transformers.masking_utils import create_causal_mask
from transformers.modeling_layers import GradientCheckpointingLayer
from transformers.modeling_outputs import BaseModelOutputWithPast
+ from transformers.models.auto import CONFIG_MAPPING
+ from transformers.models.gemma import modeling_gemma
from transformers.models.gemma.modeling_gemma import (
GemmaAttention,
GemmaConfig,
@@ -49,6 +58,8 @@ else:
GradientCheckpointingLayer = None
BaseModelOutputWithPast = None
create_causal_mask = None
+ CONFIG_MAPPING = None
+ modeling_gemma = None
def _gated_residual(
@@ -121,7 +132,10 @@ class PiGemmaRMSNorm(nn.Module):
if cond.shape[-1] != self.cond_dim:
raise ValueError(f"Expected cond dim {self.cond_dim}, got {cond.shape[-1]}")
modulation = self.dense(cond)
- if len(x.shape) == 3:
+ # Per-sample cond (B, cond_dim) → broadcast over the sequence. A
+ # per-token cond (B, T, cond_dim) is already aligned with x and must
+ # not be unsqueezed (used by pi052's amortized K_repeat path).
+ if len(x.shape) == 3 and modulation.dim() == 2:
modulation = modulation.unsqueeze(1)
scale, shift, gate = modulation.chunk(3, dim=-1)
normed = normed * (1 + scale.float()) + shift.float()
@@ -275,6 +289,8 @@ class PiGemmaModel(GemmaModel): # type: ignore[misc]
# Convert to bfloat16 if the first layer uses bfloat16
if len(self.layers) > 0 and self.layers[0].self_attn.q_proj.weight.dtype == torch.bfloat16:
hidden_states = hidden_states.to(torch.bfloat16)
+ if causal_mask is not None and torch.is_floating_point(causal_mask):
+ causal_mask = causal_mask.to(dtype=hidden_states.dtype)
# create position embeddings to be shared across the decoder layers
position_embeddings = self.rotary_emb(hidden_states, position_ids)
@@ -367,3 +383,374 @@ __all__ = [
"PaliGemmaModelWithPiGemma",
"PaliGemmaForConditionalGenerationWithPiGemma",
]
+
+
+# PI0.5 / PI052 dual-expert backbone: generic PaliGemma + Gemma action-expert
+# transformer machinery used by the pi052 policy. GemmaVariantConfig is openpi's
+# width/depth variant config (renamed from GemmaConfig to avoid clashing with
+# transformers' GemmaConfig).
+
+
+def sdpa_attention_forward(
+ module,
+ query: torch.Tensor,
+ key: torch.Tensor,
+ value: torch.Tensor,
+ attention_mask: torch.Tensor | None,
+ scaling: float,
+ dropout: float = 0.0,
+):
+ """Drop-in for ``modeling_gemma.eager_attention_forward`` using
+ ``torch.nn.functional.scaled_dot_product_attention``.
+
+ PyTorch SDPA picks the memory-efficient kernel for arbitrary additive
+ bias masks (the FA backend only accepts causal/sliding-window). On
+ H100 that is ~1.3-1.7x faster and uses ~30-40% less attention memory
+ than the eager softmax(QK^T)+matmul path. Mirrors eager's signature
+ and output shape (``(B, Lq, H, D)``) so call sites are unchanged.
+ """
+ n_rep = module.num_key_value_groups
+ if n_rep > 1:
+ key = key.repeat_interleave(n_rep, dim=1)
+ value = value.repeat_interleave(n_rep, dim=1)
+ if attention_mask is not None and attention_mask.dtype != query.dtype:
+ attention_mask = attention_mask.to(dtype=query.dtype)
+ attn_output = F.scaled_dot_product_attention(
+ query,
+ key,
+ value,
+ attn_mask=attention_mask,
+ dropout_p=dropout if module.training else 0.0,
+ is_causal=False,
+ scale=scaling,
+ )
+ return attn_output.transpose(1, 2).contiguous(), None
+
+
+# Define the complete layer computation function for gradient checkpointing
+def compute_layer_complete(
+ layer_idx, inputs_embeds, attention_mask, position_ids, adarms_cond, paligemma, gemma_expert
+):
+ models = [paligemma.model.language_model, gemma_expert.model]
+ query_states = []
+ key_states = []
+ value_states = []
+ gates = []
+ for i, hidden_states in enumerate(inputs_embeds):
+ layer = models[i].layers[layer_idx]
+ hidden_states, gate = layernorm_forward(layer.input_layernorm, hidden_states, adarms_cond[i])
+ gates.append(gate)
+ input_shape = hidden_states.shape[:-1]
+ hidden_shape = (*input_shape, -1, layer.self_attn.head_dim)
+ query_state = layer.self_attn.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ key_state = layer.self_attn.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ value_state = layer.self_attn.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
+ query_states.append(query_state)
+ key_states.append(key_state)
+ value_states.append(value_state)
+ # Concatenate and process attention
+ query_states = torch.cat(query_states, dim=2)
+ key_states = torch.cat(key_states, dim=2)
+ value_states = torch.cat(value_states, dim=2)
+ dummy_tensor = torch.zeros(
+ query_states.shape[0],
+ query_states.shape[2],
+ query_states.shape[-1],
+ device=query_states.device,
+ dtype=query_states.dtype,
+ )
+ cos, sin = paligemma.model.language_model.rotary_emb(dummy_tensor, position_ids)
+ query_states, key_states = modeling_gemma.apply_rotary_pos_emb(
+ query_states, key_states, cos, sin, unsqueeze_dim=1
+ )
+ batch_size = query_states.shape[0]
+ scaling = paligemma.model.language_model.layers[layer_idx].self_attn.scaling
+ att_output, _ = sdpa_attention_forward(
+ paligemma.model.language_model.layers[layer_idx].self_attn,
+ query_states,
+ key_states,
+ value_states,
+ attention_mask,
+ scaling,
+ )
+ # Get head_dim from the current layer, not from the model
+ head_dim = paligemma.model.language_model.layers[layer_idx].self_attn.head_dim
+ att_output = att_output.reshape(batch_size, -1, 1 * 8 * head_dim)
+ # Process layer outputs
+ outputs_embeds = []
+ start_pos = 0
+ for i, hidden_states in enumerate(inputs_embeds):
+ layer = models[i].layers[layer_idx]
+ end_pos = start_pos + hidden_states.shape[1]
+ if att_output.dtype != layer.self_attn.o_proj.weight.dtype:
+ att_output = att_output.to(layer.self_attn.o_proj.weight.dtype)
+ out_emb = layer.self_attn.o_proj(att_output[:, start_pos:end_pos])
+ # first residual
+ out_emb = _gated_residual(hidden_states, out_emb, gates[i])
+ after_first_residual = out_emb.clone()
+ out_emb, gate = layernorm_forward(layer.post_attention_layernorm, out_emb, adarms_cond[i])
+ # Convert to bfloat16 if the next layer (mlp) uses bfloat16
+ if layer.mlp.up_proj.weight.dtype == torch.bfloat16:
+ out_emb = out_emb.to(dtype=torch.bfloat16)
+ out_emb = layer.mlp(out_emb)
+ # second residual
+ out_emb = _gated_residual(after_first_residual, out_emb, gate)
+ outputs_embeds.append(out_emb)
+ start_pos = end_pos
+ return outputs_embeds
+
+
+class GemmaVariantConfig: # see openpi `gemma.py: Config`
+ """Configuration for Gemma model variants."""
+
+ def __init__(self, width, depth, mlp_dim, num_heads, num_kv_heads, head_dim):
+ self.width = width
+ self.depth = depth
+ self.mlp_dim = mlp_dim
+ self.num_heads = num_heads
+ self.num_kv_heads = num_kv_heads
+ self.head_dim = head_dim
+
+
+def get_gemma_config(variant: str) -> GemmaVariantConfig: # see openpi `gemma.py: get_config`
+ """Returns config for specified gemma variant."""
+ if variant == "gemma_300m":
+ return GemmaVariantConfig(
+ width=1024,
+ depth=18,
+ mlp_dim=4096,
+ num_heads=8,
+ num_kv_heads=1,
+ head_dim=256,
+ )
+ elif variant == "gemma_2b":
+ return GemmaVariantConfig(
+ width=2048,
+ depth=18,
+ mlp_dim=16_384,
+ num_heads=8,
+ num_kv_heads=1,
+ head_dim=256,
+ )
+ else:
+ raise ValueError(f"Unknown variant: {variant}")
+
+
+class PaliGemmaWithExpertModel(
+ nn.Module
+): # see openpi `gemma_pytorch.py: PaliGemmaWithExpertModel` this class is almost a exact copy of PaliGemmaWithExpertModel in openpi
+ """PaliGemma model with action expert for PI05."""
+
+ def __init__(
+ self,
+ vlm_config,
+ action_expert_config,
+ use_adarms=None,
+ precision: Literal["bfloat16", "float32"] = "bfloat16",
+ image_size: int = DEFAULT_IMAGE_SIZE,
+ freeze_vision_encoder: bool = False,
+ train_expert_only: bool = False,
+ ):
+ if use_adarms is None:
+ use_adarms = [False, False]
+ super().__init__()
+ self.freeze_vision_encoder = freeze_vision_encoder
+ self.train_expert_only = train_expert_only
+
+ vlm_config_hf = CONFIG_MAPPING["paligemma"]()
+ vlm_config_hf._vocab_size = 257152 # noqa: SLF001
+ vlm_config_hf.image_token_index = 257152
+ vlm_config_hf.text_config.hidden_size = vlm_config.width
+ vlm_config_hf.text_config.intermediate_size = vlm_config.mlp_dim
+ vlm_config_hf.text_config.num_attention_heads = vlm_config.num_heads
+ vlm_config_hf.text_config.head_dim = vlm_config.head_dim
+ vlm_config_hf.text_config.num_hidden_layers = vlm_config.depth
+ vlm_config_hf.text_config.num_key_value_heads = vlm_config.num_kv_heads
+ vlm_config_hf.text_config.hidden_activation = "gelu_pytorch_tanh"
+ vlm_config_hf.text_config.dtype = "float32"
+ vlm_config_hf.text_config.vocab_size = 257152
+ vlm_config_hf.text_config.use_adarms = use_adarms[0]
+ vlm_config_hf.text_config.adarms_cond_dim = vlm_config.width if use_adarms[0] else None
+ vlm_config_hf.vision_config.image_size = image_size
+ vlm_config_hf.vision_config.intermediate_size = 4304
+ vlm_config_hf.vision_config.projection_dim = 2048
+ vlm_config_hf.vision_config.projector_hidden_act = "gelu_fast"
+ vlm_config_hf.vision_config.dtype = "float32"
+
+ action_expert_config_hf = CONFIG_MAPPING["gemma"](
+ head_dim=action_expert_config.head_dim,
+ hidden_size=action_expert_config.width,
+ intermediate_size=action_expert_config.mlp_dim,
+ num_attention_heads=action_expert_config.num_heads,
+ num_hidden_layers=action_expert_config.depth,
+ num_key_value_heads=action_expert_config.num_kv_heads,
+ vocab_size=257152,
+ hidden_activation="gelu_pytorch_tanh",
+ dtype="float32",
+ use_adarms=use_adarms[1],
+ adarms_cond_dim=action_expert_config.width if use_adarms[1] else None,
+ )
+
+ self.paligemma = PaliGemmaForConditionalGenerationWithPiGemma(config=vlm_config_hf)
+ self.gemma_expert = PiGemmaForCausalLM(config=action_expert_config_hf)
+ self.gemma_expert.model.embed_tokens = None
+
+ self.to_bfloat16_for_selected_params(precision)
+ self._set_requires_grad()
+
+ def to_bfloat16_for_selected_params(self, precision: Literal["bfloat16", "float32"] = "bfloat16"):
+ if precision == "bfloat16":
+ self.to(dtype=torch.bfloat16)
+ elif precision == "float32":
+ self.to(dtype=torch.float32)
+ return
+ else:
+ raise ValueError(f"Invalid precision: {precision}")
+
+ # Keep full vision path in float32 so we never toggle (toggle causes optimizer
+ # "same dtype" error). Saves memory vs full float32; more memory than only 3 params.
+ params_to_keep_float32 = [
+ "vision_tower",
+ "multi_modal_projector",
+ "lm_head",
+ "input_layernorm",
+ "post_attention_layernorm",
+ "model.norm",
+ ]
+
+ for name, param in self.named_parameters():
+ if any(selector in name for selector in params_to_keep_float32):
+ param.data = param.data.to(dtype=torch.float32)
+
+ def _set_requires_grad(self):
+ if self.freeze_vision_encoder:
+ self.paligemma.model.vision_tower.eval()
+ for param in self.paligemma.model.vision_tower.parameters():
+ param.requires_grad = False
+ if self.train_expert_only:
+ self.paligemma.eval()
+ for param in self.paligemma.parameters():
+ param.requires_grad = False
+
+ def train(self, mode: bool = True):
+ super().train(mode)
+ if self.freeze_vision_encoder:
+ self.paligemma.model.vision_tower.eval()
+ if self.train_expert_only:
+ self.paligemma.eval()
+
+ def embed_image(self, image: torch.Tensor):
+ # Vision tower and multi_modal_projector are kept in float32 (params_to_keep_float32).
+ out_dtype = image.dtype
+ if image.dtype != torch.float32:
+ image = image.to(torch.float32)
+ image_outputs = self.paligemma.model.get_image_features(image)
+ # OpenPI / big_vision convention: image (soft) tokens are NOT scaled by the
+ # Gemma embedder normalizer (sqrt(hidden_size)) — only text tokens are. lerobot/pi05_base
+ # was trained in this regime, so scaling image features here over-scales them ~45x and
+ # breaks the pretrained vision-language alignment. Keep image features un-normalized.
+ features = image_outputs.pooler_output
+ if features.dtype != out_dtype:
+ features = features.to(out_dtype)
+ return features
+
+ def embed_language_tokens(self, tokens: torch.Tensor):
+ return self.paligemma.model.language_model.embed_tokens(tokens)
+
+ def forward(
+ self,
+ attention_mask: torch.Tensor | None = None,
+ position_ids: torch.LongTensor | None = None,
+ past_key_values: list[torch.FloatTensor] | None = None,
+ inputs_embeds: list[torch.FloatTensor] | None = None,
+ use_cache: bool | None = None,
+ adarms_cond: list[torch.Tensor] | None = None,
+ ):
+ if adarms_cond is None:
+ adarms_cond = [None, None]
+ if inputs_embeds[1] is None:
+ prefix_output = self.paligemma.model.language_model.forward(
+ inputs_embeds=inputs_embeds[0],
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ use_cache=use_cache,
+ adarms_cond=adarms_cond[0] if adarms_cond is not None else None,
+ )
+ prefix_past_key_values = prefix_output.past_key_values
+ prefix_output = prefix_output.last_hidden_state
+ suffix_output = None
+ elif inputs_embeds[0] is None:
+ suffix_output = self.gemma_expert.model.forward(
+ inputs_embeds=inputs_embeds[1],
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ use_cache=use_cache,
+ adarms_cond=adarms_cond[1] if adarms_cond is not None else None,
+ )
+ suffix_output = suffix_output.last_hidden_state
+ prefix_output = None
+ prefix_past_key_values = None
+ else:
+ models = [self.paligemma.model.language_model, self.gemma_expert.model]
+ num_layers = self.paligemma.config.text_config.num_hidden_layers
+
+ # Check if gradient checkpointing is enabled for any of the models
+ use_gradient_checkpointing = (
+ hasattr(self.gemma_expert.model, "gradient_checkpointing")
+ and self.gemma_expert.model.gradient_checkpointing
+ and self.training
+ ) or (hasattr(self, "gradient_checkpointing") and self.gradient_checkpointing and self.training)
+
+ # Process all layers with gradient checkpointing if enabled
+ for layer_idx in range(num_layers):
+ if use_gradient_checkpointing:
+ inputs_embeds = torch.utils.checkpoint.checkpoint(
+ compute_layer_complete,
+ layer_idx,
+ inputs_embeds,
+ attention_mask,
+ position_ids,
+ adarms_cond,
+ use_reentrant=False,
+ preserve_rng_state=False,
+ paligemma=self.paligemma,
+ gemma_expert=self.gemma_expert,
+ )
+ else:
+ inputs_embeds = compute_layer_complete(
+ layer_idx,
+ inputs_embeds,
+ attention_mask,
+ position_ids,
+ adarms_cond,
+ paligemma=self.paligemma,
+ gemma_expert=self.gemma_expert,
+ )
+
+ # final norm
+ def compute_final_norms(inputs_embeds, adarms_cond):
+ outputs_embeds = []
+ for i, hidden_states in enumerate(inputs_embeds):
+ out_emb, _ = layernorm_forward(models[i].norm, hidden_states, adarms_cond[i])
+ outputs_embeds.append(out_emb)
+ return outputs_embeds
+
+ # Apply gradient checkpointing to final norm if enabled
+ if use_gradient_checkpointing:
+ outputs_embeds = torch.utils.checkpoint.checkpoint(
+ compute_final_norms,
+ inputs_embeds,
+ adarms_cond,
+ use_reentrant=False,
+ preserve_rng_state=False,
+ )
+ else:
+ outputs_embeds = compute_final_norms(inputs_embeds, adarms_cond)
+
+ prefix_output = outputs_embeds[0]
+ suffix_output = outputs_embeds[1]
+ prefix_past_key_values = None
+
+ return [prefix_output, suffix_output], prefix_past_key_values
diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py
index 8bfa16a98..a12e12348 100644
--- a/src/lerobot/scripts/lerobot_train.py
+++ b/src/lerobot/scripts/lerobot_train.py
@@ -20,9 +20,11 @@ Requires: pip install 'lerobot[training]' (includes dataset + accelerate + wand
import dataclasses
import logging
+import os
import sys
import time
from contextlib import nullcontext
+from datetime import timedelta
from pprint import pformat
from typing import TYPE_CHECKING, Any
@@ -91,6 +93,7 @@ def update_policy(
lr_scheduler=None,
lock=None,
sample_weighter=None,
+ log_metrics: bool = True,
) -> tuple[MetricsTracker, dict | None]:
"""
Performs a single training step to update the policy's weights.
@@ -108,6 +111,7 @@ def update_policy(
lr_scheduler: An optional learning rate scheduler.
lock: An optional lock for thread-safe optimizer updates.
sample_weighter: Optional SampleWeighter instance for per-sample loss weighting.
+ log_metrics: Whether to synchronize and record GPU metrics this step.
Returns:
A tuple containing:
@@ -175,12 +179,20 @@ def update_policy(
if has_method(accelerator.unwrap_model(policy, keep_fp32_wrapper=True), "update"):
accelerator.unwrap_model(policy, keep_fp32_wrapper=True).update()
- train_metrics.loss = loss.item()
- train_metrics.grad_norm = grad_norm.item()
train_metrics.lr = optimizer.param_groups[0]["lr"]
- train_metrics.update_s = time.perf_counter() - start_time
if torch.cuda.is_available():
train_metrics.gpu_mem_gb = torch.cuda.max_memory_allocated() / (1024**3)
+ train_metrics.accumulate_tensor("loss", loss)
+ train_metrics.accumulate_tensor("grad_norm", grad_norm)
+ train_metrics.update_s = time.perf_counter() - start_time
+ # Synchronize accumulated GPU metrics only when logging.
+ if log_metrics:
+ train_metrics.materialize_tensors()
+ # Materialize detached loss components during the same logging synchronization.
+ if output_dict:
+ output_dict = {
+ k: (v.item() if isinstance(v, torch.Tensor) else v) for k, v in output_dict.items()
+ }
# Aggregate the policy's scalar outputs for logging and rank-reduction across the log window.
if output_dict:
train_metrics.update_metrics(output_dict)
@@ -211,7 +223,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
require_package("accelerate", extra="training")
from accelerate import Accelerator
- from accelerate.utils import DistributedDataParallelKwargs, DistributedType
+ from accelerate.utils import DistributedDataParallelKwargs, DistributedType, InitProcessGroupKwargs
cfg.validate()
@@ -220,7 +232,16 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
# We set step_scheduler_with_optimizer=False to prevent accelerate from adjusting the lr_scheduler steps based on the num_processes
# We set find_unused_parameters=True to handle models with conditional computation
if accelerator is None:
- ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
+ # Static graphs restore DDP overlap when conditional parameter usage is stable.
+ # Environment flags retain the existing defaults.
+ ddp_find_unused = os.environ.get("LEROBOT_DDP_FIND_UNUSED", "1") == "1"
+ ddp_static_graph = os.environ.get("LEROBOT_DDP_STATIC_GRAPH", "0") == "1"
+ ddp_kwargs = DistributedDataParallelKwargs(
+ find_unused_parameters=ddp_find_unused and not ddp_static_graph,
+ static_graph=ddp_static_graph,
+ )
+ # Allow rank 0 enough time to index large datasets before other ranks leave the barrier.
+ ipg_kwargs = InitProcessGroupKwargs(timeout=timedelta(hours=2))
# Accelerate auto-detects the device based on the available hardware and ignores the policy.device setting.
# Force the device to be CPU when the active config's device is set to CPU (works for both policy and reward model training).
force_cpu = cfg.trainable_config.device == "cpu"
@@ -230,7 +251,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
accelerator = Accelerator(
step_scheduler_with_optimizer=False,
mixed_precision=mixed_precision,
- kwargs_handlers=[ddp_kwargs],
+ kwargs_handlers=[ddp_kwargs, ipg_kwargs],
cpu=force_cpu,
)
@@ -326,6 +347,14 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
active_cfg = cfg.trainable_config
processor_pretrained_path = active_cfg.pretrained_path
+ # A weight checkpoint may contain PI05 or differently configured PI052 processors.
+ if cfg.policy.type == "pi052" and processor_pretrained_path is not None and not cfg.resume:
+ logging.warning(
+ "pi052 is loading pretrained weights from %s, but building processors from the current "
+ "pi052 config so recipe text labels and FAST action labels are generated.",
+ processor_pretrained_path,
+ )
+ processor_pretrained_path = None
processor_kwargs = {}
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
@@ -334,6 +363,13 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
if cfg.is_reward_model_training:
processor_kwargs["dataset_meta"] = dataset.meta
+ if cfg.policy.type in {"pi0_fast", "pi052"}:
+ processor_kwargs["dataset_repo_id"] = cfg.dataset.repo_id
+ processor_kwargs["dataset_revision"] = cfg.dataset.revision
+ processor_kwargs["dataset_episodes"] = cfg.dataset.episodes
+ processor_kwargs["dataset_exclude_episodes"] = cfg.dataset.exclude_episodes
+ processor_kwargs["dataset_root"] = cfg.dataset.root
+
if not cfg.is_reward_model_training and processor_pretrained_path is not None:
preprocessor_overrides = {
"device_processor": {"device": device.type},
@@ -430,13 +466,17 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
# same permutation. accelerate then shards it disjointly across ranks via BatchSamplerShard
# without needing a `generator` attribute to synchronize an RNG, and resume is sample-exact.
shuffle = False
+ from_indices = dataset.meta.episodes["dataset_from_index"]
+ to_indices = dataset.meta.episodes["dataset_to_index"]
+ seed = cfg.seed if cfg.seed is not None else 0
+
sampler = EpisodeAwareSampler(
- dataset.meta.episodes["dataset_from_index"],
- dataset.meta.episodes["dataset_to_index"],
+ from_indices,
+ to_indices,
episode_indices_to_use=dataset.episodes,
drop_n_last_frames=getattr(active_cfg, "drop_n_last_frames", 0),
shuffle=True,
- seed=cfg.seed if cfg.seed is not None else 0,
+ seed=seed,
absolute_to_relative_idx=dataset.absolute_to_relative_idx,
)
if cfg.resume and step > 0:
@@ -583,7 +623,10 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
batch = preprocessor(batch)
train_tracker.dataloading_s = time.perf_counter() - start_time
- train_tracker, _ = update_policy(
+ # Synchronize GPU metrics only for updates that will be logged.
+ log_metrics = cfg.log_freq > 0 and (step + 1) % cfg.log_freq == 0
+
+ train_tracker, output_dict = update_policy(
train_tracker,
policy,
batch,
@@ -592,6 +635,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
accelerator=accelerator,
lr_scheduler=lr_scheduler,
sample_weighter=sample_weighter,
+ log_metrics=log_metrics,
)
# Note: eval and checkpoint happens *after* the `step`th training update has completed, so we
@@ -692,10 +736,11 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
if is_main_process:
step_id = get_step_identifier(step, cfg.steps)
logging.info(f"Eval policy at step {step}")
+ eval_target_policy = accelerator.unwrap_model(policy)
with torch.no_grad(), accelerator.autocast():
eval_info = eval_policy_all(
envs=eval_env, # dict[suite][task_id] -> vec_env
- policy=accelerator.unwrap_model(policy),
+ policy=eval_target_policy,
env_preprocessor=env_preprocessor,
env_postprocessor=env_postprocessor,
preprocessor=preprocessor,
diff --git a/src/lerobot/utils/logging_utils.py b/src/lerobot/utils/logging_utils.py
index fd3021fea..014496348 100644
--- a/src/lerobot/utils/logging_utils.py
+++ b/src/lerobot/utils/logging_utils.py
@@ -105,6 +105,8 @@ class MetricsTracker:
"epochs",
"accelerator",
"_caller_metrics",
+ "_tensor_sums",
+ "_tensor_counts",
]
def __init__(
@@ -133,6 +135,8 @@ class MetricsTracker:
# Meter names the caller registered up front. update_metrics() leaves these untouched, so a
# policy that echoes e.g. "loss" in its output dict can't clobber the aggregated meter.
self._caller_metrics: set[str] = set(self.metrics)
+ self._tensor_sums: dict[str, torch.Tensor] = {}
+ self._tensor_counts: dict[str, int] = {}
def __getattr__(self, name: str) -> int | dict[str, AverageMeter] | AverageMeter | Any:
if name in self.__dict__:
@@ -160,6 +164,22 @@ class MetricsTracker:
self.episodes = self.samples / self._avg_samples_per_ep
self.epochs = self.samples / self._num_frames
+ def accumulate_tensor(self, name: str, value: torch.Tensor) -> None:
+ """Accumulate a detached metric on-device until the next logging step."""
+ if name not in self.metrics:
+ raise KeyError(f"Unknown metric {name!r}.")
+ value = value.detach()
+ self._tensor_sums[name] = self._tensor_sums.get(name, torch.zeros_like(value)) + value
+ self._tensor_counts[name] = self._tensor_counts.get(name, 0) + 1
+
+ def materialize_tensors(self) -> None:
+ """Transfer pending tensor averages to their meters with one sync per metric."""
+ for name, total in self._tensor_sums.items():
+ count = self._tensor_counts[name]
+ self.metrics[name].update((total / count).item(), n=count)
+ self._tensor_sums.clear()
+ self._tensor_counts.clear()
+
def update_metrics(self, values: dict[str, Any]) -> None:
"""Accumulate a dict of scalar metrics, auto-registering a meter for each new key.
@@ -167,7 +187,7 @@ class MetricsTracker:
Caller-registered metrics (those passed to the constructor) are never overridden.
"""
for name, value in values.items():
- if isinstance(value, bool) or not isinstance(value, (int, float)):
+ if isinstance(value, bool) or not isinstance(value, int | float):
continue
if name in self._caller_metrics:
continue
@@ -235,3 +255,5 @@ class MetricsTracker:
"""Resets average meters."""
for m in self.metrics.values():
m.reset()
+ self._tensor_sums.clear()
+ self._tensor_counts.clear()
diff --git a/tests/policies/pi052/test_pi052_attention_masking.py b/tests/policies/pi052/test_pi052_attention_masking.py
new file mode 100644
index 000000000..96c98fb91
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_attention_masking.py
@@ -0,0 +1,151 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Attention-masking tests for the PI052 (π0.5 v2) text head.
+
+Regression coverage for the text-CE collapse bug: PaliGemma's
+``embed_prefix`` flags every language token ``att=0``, which
+``make_att_2d_masks`` turns into one fully *bidirectional* block. Under
+that mask the text cross-entropy degenerates into a copy task — a
+supervised target token attends to the tokens it is trained to predict —
+and the LM head never learns causal generation, so ``select_message``
+collapses at inference.
+
+``_mark_target_span_causal`` sets ``att=1`` on the supervised target
+language positions so each target token attends causally among the
+targets while staying bidirectional to images + the user prompt. These
+tests pin that behaviour for the PaliGemma prefix layout.
+"""
+
+import pytest
+import torch
+
+# modeling_pi052 / modeling_pi05 import transformers transitively.
+pytest.importorskip("transformers")
+
+from lerobot.policies.pi05.modeling_pi05 import make_att_2d_masks # noqa: E402
+from lerobot.policies.pi052.modeling_pi052 import ( # noqa: E402
+ _mark_target_span_causal,
+ _shifted_lin_ce,
+)
+
+
+def _shifted_ce(logits, labels):
+ """Adapter: ``_shifted_lin_ce`` is Liger-fused (hidden @ lm_head_weightᵀ).
+
+ An identity ``lm_head_weight`` makes the computed logits equal ``logits``.
+ Liger's Triton kernel is GPU-only, so inputs run on CUDA; the loss is
+ returned on CPU so grad still flows back to the CPU ``logits`` leaf.
+ """
+ if not torch.cuda.is_available():
+ pytest.skip("Liger fused CE requires CUDA")
+ vocab_size = logits.shape[-1]
+ eye = torch.eye(vocab_size, dtype=logits.dtype, device="cuda")
+ return _shifted_lin_ce(logits.cuda(), eye, labels.cuda()).cpu()
+
+
+# Synthetic prefix: two image tokens, three prompt tokens, and four supervised target tokens.
+# Text labels mask the prompt with -100 and cover the target through the prefix end.
+N_IMAGE = 2
+N_PROMPT = 3
+N_TARGET = 4
+LANG_START = N_IMAGE
+LANG_END = N_IMAGE + N_PROMPT + N_TARGET # = prefix length
+PREFIX_LEN = LANG_END
+
+
+def _embed_prefix_att_masks() -> torch.Tensor:
+ """Mimic PaliGemma ``embed_prefix``: images + lang all att=0."""
+ return torch.zeros(1, PREFIX_LEN, dtype=torch.bool)
+
+
+def _text_labels() -> torch.Tensor:
+ """-100 over the prompt span, real ids over the target span."""
+ labels = torch.full((1, N_PROMPT + N_TARGET), -100, dtype=torch.long)
+ labels[0, N_PROMPT:] = torch.arange(10, 10 + N_TARGET)
+ return labels
+
+
+def _attends(prefix_att_masks: torch.Tensor) -> torch.Tensor:
+ """2D boolean attendance matrix; ``[i, j]`` True ⇒ i attends to j."""
+ pad = torch.ones(1, PREFIX_LEN, dtype=torch.bool)
+ return make_att_2d_masks(pad, prefix_att_masks)[0]
+
+
+def test_mark_sets_att_on_targets_only():
+ """Only the supervised target language positions flip to att=1."""
+ marked = _mark_target_span_causal(_embed_prefix_att_masks(), _text_labels(), LANG_START, LANG_END)
+ expected = [False] * PREFIX_LEN
+ for i in range(LANG_START + N_PROMPT, LANG_END): # target span
+ expected[i] = True
+ assert marked[0].tolist() == expected
+
+
+def test_target_tokens_attend_causally_among_themselves():
+ """A target token must NOT attend to later targets, but must attend
+ to earlier ones — genuine causal next-token prediction."""
+ marked = _mark_target_span_causal(_embed_prefix_att_masks(), _text_labels(), LANG_START, LANG_END)
+ attends = _attends(marked)
+ tgt = range(LANG_START + N_PROMPT, LANG_END)
+ for i in tgt:
+ for j in tgt:
+ if j > i:
+ assert not attends[i, j], f"target {i} must not see future target {j}"
+ else:
+ assert attends[i, j], f"target {i} must see earlier/self target {j}"
+
+
+def test_target_tokens_attend_prompt_and_images_bidirectionally():
+ """Targets keep full visibility of images + the user prompt."""
+ marked = _mark_target_span_causal(_embed_prefix_att_masks(), _text_labels(), LANG_START, LANG_END)
+ attends = _attends(marked)
+ context = list(range(0, LANG_START + N_PROMPT)) # images + prompt
+ for i in range(LANG_START + N_PROMPT, LANG_END):
+ for j in context:
+ assert attends[i, j], f"target {i} must attend context {j}"
+
+
+def test_non_target_subtask_stays_bidirectional():
+ """A flow-only / non-target language span (all -100 labels) leaves the
+ mask untouched — the action expert reads it bidirectionally."""
+ all_ignored = torch.full((1, N_PROMPT + N_TARGET), -100, dtype=torch.long)
+ marked = _mark_target_span_causal(_embed_prefix_att_masks(), all_ignored, LANG_START, LANG_END)
+ assert torch.equal(marked, _embed_prefix_att_masks())
+
+
+def test_unmarked_mask_is_bidirectional_the_bug():
+ """Documents the bug the fix prevents: without ``_mark_target_span_causal``
+ a target token attends *bidirectionally* to later targets — the
+ text-CE can copy the answer it is trained to predict."""
+ attends = _attends(_embed_prefix_att_masks())
+ first_tgt = LANG_START + N_PROMPT
+ last_tgt = LANG_END - 1
+ assert attends[first_tgt, last_tgt], (
+ "raw embed_prefix mask is bidirectional over language — the first "
+ "target token can see the last, which is the collapse bug"
+ )
+
+
+def test_shifted_ce_returns_zero_when_no_text_positions_are_supervised():
+ pytest.importorskip("liger_kernel")
+ logits = torch.randn(2, 4, 8, requires_grad=True)
+ labels = torch.full((2, 4), -100, dtype=torch.long)
+
+ loss = _shifted_ce(logits, labels)
+
+ assert loss.item() == 0
+ loss.backward()
+ assert logits.grad is not None
diff --git a/tests/policies/pi052/test_pi052_bucketed_ce.py b/tests/policies/pi052/test_pi052_bucketed_ce.py
new file mode 100644
index 000000000..717874783
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_bucketed_ce.py
@@ -0,0 +1,146 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import pytest
+import torch
+
+pytest.importorskip("transformers")
+
+from lerobot.policies.pi052.modeling_pi052 import _lin_ce_flat, _shifted_lin_ce
+
+
+def test_shifted_ce_none_retains_distinct_per_sample_losses():
+ hidden = torch.tensor(
+ [
+ [[8.0, 0.0], [0.0, 8.0], [0.0, 0.0]],
+ [[0.0, 8.0], [8.0, 0.0], [0.0, 0.0]],
+ ]
+ )
+ labels = torch.tensor([[0, 0, 1], [0, 0, 1]])
+ losses = _shifted_lin_ce(hidden, torch.eye(2), labels, reduction="none")
+
+ assert losses.shape == (2,)
+ assert losses[0] < losses[1]
+
+
+def test_checkpoint_resolution_forwards_explicit_hub_options(monkeypatch, tmp_path):
+ import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
+
+ checkpoint = tmp_path / "model.safetensors"
+ checkpoint.touch()
+ calls = []
+
+ def fake_cached_file(model_id, filename, **kwargs):
+ calls.append((model_id, filename, kwargs))
+ return None if filename.endswith("index.json") else str(checkpoint)
+
+ monkeypatch.setattr(modeling_pi05, "cached_file", fake_cached_file)
+ files = modeling_pi05._resolve_weight_files(
+ "org/model",
+ force_download=True,
+ resume_download=True,
+ proxies={"https": "proxy"},
+ token="secret",
+ cache_dir=tmp_path / "cache",
+ local_files_only=True,
+ revision="commit",
+ )
+
+ assert files == [checkpoint]
+ for _model_id, _filename, kwargs in calls:
+ assert kwargs["revision"] == "commit"
+ assert kwargs["cache_dir"] == tmp_path / "cache"
+ assert kwargs["force_download"] is True
+ assert kwargs["resume_download"] is True
+ assert kwargs["proxies"] == {"https": "proxy"}
+ assert kwargs["token"] == "secret"
+ assert kwargs["local_files_only"] is True
+
+
+def test_checkpoint_resolution_rejects_local_directory_without_weights(tmp_path):
+ import lerobot.policies.pi05.modeling_pi05 as modeling_pi05
+
+ with pytest.raises(FileNotFoundError, match="model.safetensors"):
+ modeling_pi05._resolve_weight_files(
+ tmp_path,
+ force_download=False,
+ resume_download=None,
+ proxies=None,
+ token=None,
+ cache_dir=None,
+ local_files_only=False,
+ revision=None,
+ )
+
+
+@pytest.mark.parametrize("z_loss_weight", [0.0, 1e-4])
+@pytest.mark.parametrize("rows,valid_rows", [(24, 9), (48, 25)])
+def test_bucketed_ce_matches_dense_loss_and_gradients(z_loss_weight, rows, valid_rows):
+ generator = torch.Generator().manual_seed(23)
+ hidden_size, vocab_size = 7, 19
+ hidden_ref = torch.randn(rows, hidden_size, generator=generator, dtype=torch.float64, requires_grad=True)
+ weight_ref = torch.randn(
+ vocab_size, hidden_size, generator=generator, dtype=torch.float64, requires_grad=True
+ )
+ labels = torch.full((rows,), -100, dtype=torch.long)
+ valid_indices = torch.randperm(rows, generator=generator)[:valid_rows]
+ labels[valid_indices] = torch.randint(0, vocab_size, (valid_rows,), generator=generator)
+ hidden_bucketed = hidden_ref.detach().clone().requires_grad_(True)
+ weight_bucketed = weight_ref.detach().clone().requires_grad_(True)
+
+ import lerobot.policies.pi052.modeling_pi052 as modeling_pi052
+
+ loss_ref = _lin_ce_flat(hidden_ref, weight_ref, labels, z_loss_weight=z_loss_weight)
+ old_limit = modeling_pi052._LOGITS_CE_MAX_POSITIONS
+ modeling_pi052._LOGITS_CE_MAX_POSITIONS = 16
+ try:
+ loss_bucketed = _lin_ce_flat(
+ hidden_bucketed,
+ weight_bucketed,
+ labels,
+ z_loss_weight=z_loss_weight,
+ )
+ finally:
+ modeling_pi052._LOGITS_CE_MAX_POSITIONS = old_limit
+
+ loss_ref.backward()
+ loss_bucketed.backward()
+
+ torch.testing.assert_close(loss_bucketed, loss_ref, rtol=1e-6, atol=1e-6)
+ torch.testing.assert_close(hidden_bucketed.grad, hidden_ref.grad, rtol=1e-12, atol=1e-12)
+ torch.testing.assert_close(weight_bucketed.grad, weight_ref.grad, rtol=1e-12, atol=1e-12)
+
+
+def test_bucketed_ce_all_ignored_preserves_zero_gradients():
+ hidden = torch.randn(24, 7, dtype=torch.float64, requires_grad=True)
+ weight = torch.randn(19, 7, dtype=torch.float64, requires_grad=True)
+ labels = torch.full((24,), -100, dtype=torch.long)
+
+ import lerobot.policies.pi052.modeling_pi052 as modeling_pi052
+
+ old_limit = modeling_pi052._LOGITS_CE_MAX_POSITIONS
+ modeling_pi052._LOGITS_CE_MAX_POSITIONS = 16
+ try:
+ loss = _lin_ce_flat(hidden, weight, labels)
+ finally:
+ modeling_pi052._LOGITS_CE_MAX_POSITIONS = old_limit
+ loss.backward()
+
+ assert loss.item() == 0.0
+ assert hidden.grad is not None
+ assert weight.grad is not None
+ assert torch.count_nonzero(hidden.grad) == 0
+ assert torch.count_nonzero(weight.grad) == 0
diff --git a/tests/policies/pi052/test_pi052_checkpoint_portability.py b/tests/policies/pi052/test_pi052_checkpoint_portability.py
new file mode 100644
index 000000000..824dc5263
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_checkpoint_portability.py
@@ -0,0 +1,152 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import json
+import shutil
+from dataclasses import asdict
+from types import SimpleNamespace
+
+import numpy as np
+import pytest
+import torch
+
+from lerobot.configs import FeatureType, NormalizationMode, PolicyFeature
+from lerobot.configs.recipe import MessageTurn, TrainingRecipe
+from lerobot.policies import make_pre_post_processors
+from lerobot.processor import ActionTokenizerProcessorStep, DataProcessorPipeline, NormalizerProcessorStep
+from lerobot.processor.converters import identity_transition
+from lerobot.processor.render_messages_processor import RenderMessagesStep
+from lerobot.utils.constants import ACTION
+
+
+class _ActionTokenizer:
+ def __call__(self, actions):
+ return np.asarray(actions).round().astype(np.int64)
+
+ def save_pretrained(self, path):
+ path.mkdir(parents=True)
+ (path / "processor_config.json").write_text('{"processor_class": "_ActionTokenizer"}\n')
+
+
+class _PaligemmaTokenizer:
+ vocab_size = 4096
+ bos_token_id = 2
+
+ def encode(self, text, **kwargs):
+ return [10, 11] if text == "Action: " else [12]
+
+
+def _make_pipeline(action_tokenizer_path):
+ recipe = TrainingRecipe(
+ messages=[
+ MessageTurn(role="user", content="${task}", stream="high_level"),
+ MessageTurn(role="assistant", content="${subtask}", stream="low_level", target=True),
+ ]
+ )
+ stats = {ACTION: {"min": torch.tensor([-1.0, -2.0]), "max": torch.tensor([1.0, 2.0])}}
+ normalizer = NormalizerProcessorStep(
+ features={ACTION: PolicyFeature(type=FeatureType.ACTION, shape=(2,))},
+ norm_map={FeatureType.ACTION: NormalizationMode.MIN_MAX},
+ stats=stats,
+ )
+ action_tokenizer = ActionTokenizerProcessorStep(
+ action_tokenizer_name=str(action_tokenizer_path),
+ max_action_tokens=16,
+ fast_skip_tokens=128,
+ )
+ return DataProcessorPipeline(
+ [normalizer, RenderMessagesStep(recipe), action_tokenizer],
+ name="policy_preprocessor",
+ to_transition=identity_transition,
+ to_output=identity_transition,
+ )
+
+
+def test_pi052_pipeline_embeds_and_loads_fitted_action_tokenizer(tmp_path, monkeypatch):
+ original_cache = tmp_path / "original_fast_cache"
+ original_cache.mkdir()
+ tokenizer = _ActionTokenizer()
+ monkeypatch.setattr(
+ "lerobot.processor.tokenizer_processor.AutoProcessor.from_pretrained",
+ lambda path, **kwargs: tokenizer,
+ )
+ monkeypatch.setattr(
+ "lerobot.processor.tokenizer_processor.AutoTokenizer.from_pretrained",
+ lambda *args, **kwargs: _PaligemmaTokenizer(),
+ )
+ monkeypatch.setattr(
+ "lerobot.policies.pi052.fit_fast_tokenizer.fit_fast_tokenizer",
+ lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("FAST fitting must not run")),
+ )
+
+ pipeline = _make_pipeline(original_cache)
+ expected_tokens = pipeline.steps[-1]._tokenize_action(torch.tensor([[[0.2, 0.8]]]))[0]
+ expected_recipe = asdict(pipeline.steps[1].recipe)
+ expected_state = pipeline.steps[0].state_dict()
+ checkpoint = tmp_path / "checkpoint"
+ pipeline.save_pretrained(checkpoint)
+ DataProcessorPipeline(
+ [],
+ name="policy_postprocessor",
+ to_transition=identity_transition,
+ to_output=identity_transition,
+ ).save_pretrained(checkpoint)
+
+ saved_config = json.loads((checkpoint / "policy_preprocessor.json").read_text())
+ tokenizer_step = saved_config["steps"][2]
+ assert tokenizer_step["config"]["action_tokenizer_name"] == "action_tokenizer"
+ assert tokenizer_step["artifacts"] == {"action_tokenizer_name": "action_tokenizer"}
+ assert (checkpoint / "action_tokenizer" / "processor_config.json").is_file()
+
+ shutil.rmtree(original_cache)
+ loaded, _ = make_pre_post_processors(
+ SimpleNamespace(type="pi052", auto_fit_fast_tokenizer=True),
+ pretrained_path=str(checkpoint),
+ dataset_repo_id="org/dataset-that-must-not-be-read",
+ )
+
+ assert asdict(loaded.steps[1].recipe) == expected_recipe
+ for key, tensor in expected_state.items():
+ torch.testing.assert_close(loaded.steps[0].state_dict()[key], tensor)
+ torch.testing.assert_close(
+ loaded.steps[-1]._tokenize_action(torch.tensor([[[0.2, 0.8]]]))[0],
+ expected_tokens,
+ )
+
+
+def test_pi052_pipeline_rejects_missing_fitted_action_tokenizer(tmp_path, monkeypatch):
+ tokenizer = _ActionTokenizer()
+ monkeypatch.setattr(
+ "lerobot.processor.tokenizer_processor.AutoProcessor.from_pretrained",
+ lambda path, **kwargs: tokenizer,
+ )
+ monkeypatch.setattr(
+ "lerobot.processor.tokenizer_processor.AutoTokenizer.from_pretrained",
+ lambda *args, **kwargs: _PaligemmaTokenizer(),
+ )
+
+ pipeline = _make_pipeline(tmp_path / "original_fast_cache")
+ checkpoint = tmp_path / "checkpoint"
+ pipeline.save_pretrained(checkpoint)
+ shutil.rmtree(checkpoint / "action_tokenizer")
+
+ with pytest.raises(FileNotFoundError, match="Checkpoint artifacts are incomplete"):
+ DataProcessorPipeline.from_pretrained(
+ checkpoint,
+ config_filename="policy_preprocessor.json",
+ to_transition=identity_transition,
+ to_output=identity_transition,
+ )
diff --git a/tests/policies/pi052/test_pi052_fast_action_loss.py b/tests/policies/pi052/test_pi052_fast_action_loss.py
new file mode 100644
index 000000000..5250aeb90
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_fast_action_loss.py
@@ -0,0 +1,162 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Regression tests for PI052 FAST action-code supervision."""
+
+from types import SimpleNamespace
+
+import pytest
+import torch
+from torch import nn
+from torch.nn import functional as F # noqa: N812
+
+pytest.importorskip("transformers")
+pytest.importorskip("liger_kernel")
+
+from lerobot.policies.pi052.modeling_pi052 import PI052Policy, _fast_lin_ce # noqa: E402
+from lerobot.policies.pi052.processor_pi052 import make_pi052_pre_post_processors # noqa: E402
+
+
+def _fast_ce(logits, action_tokens, action_code_mask, predict_actions_t):
+ """Adapter: ``_fast_lin_ce`` is Liger-fused (hidden @ lm_head_weightᵀ).
+
+ Feeding an identity ``lm_head_weight`` makes the computed logits equal the
+ provided ``logits``, so these regression tests exercise the masking/gating
+ logic exactly as before the fused-CE refactor. Liger's Triton kernel is
+ GPU-only, so inputs are moved to CUDA and the loss is returned on CPU
+ (keeping grad flowing back to the CPU ``logits`` leaf).
+ """
+ if not torch.cuda.is_available():
+ pytest.skip("Liger fused CE requires CUDA")
+ vocab_size = logits.shape[-1]
+ eye = torch.eye(vocab_size, dtype=logits.dtype, device="cuda")
+ predict = predict_actions_t.cuda() if predict_actions_t is not None else None
+ loss = _fast_lin_ce(logits.cuda(), eye, action_tokens.cuda(), action_code_mask.cuda(), predict)
+ return loss.cpu()
+
+
+def test_fast_ce_supervises_only_discrete_action_codes():
+ """Wrapper tokens can be wrong without affecting the FAST action-code loss."""
+ vocab_size = 8
+ action_tokens = torch.tensor([[1, 2, 3, 4, 5, 0]])
+ action_code_mask = torch.tensor([[False, False, True, True, False, False]])
+
+ logits = torch.zeros(1, action_tokens.shape[1], vocab_size)
+ # Deliberately bad wrapper-token predictions. These should be ignored.
+ logits[0, 0, 7] = 10.0 # target would be token 2
+ logits[0, 3, 7] = 10.0 # target would be delimiter token 5
+ # Correct action-code predictions: hidden t predicts target t + 1.
+ logits[0, 1, 3] = 10.0
+ logits[0, 2, 4] = 10.0
+
+ loss = _fast_ce(logits, action_tokens, action_code_mask, predict_actions_t=None)
+ expected = F.cross_entropy(
+ torch.stack([logits[0, 1], logits[0, 2]]),
+ torch.tensor([3, 4]),
+ reduction="mean",
+ )
+
+ # Allow the fused GPU kernel's ~1e-7 difference on small losses.
+ assert torch.allclose(loss, expected, atol=1e-5, rtol=1e-3)
+
+
+def test_fast_ce_masks_non_action_samples():
+ """Recipe samples with predict_actions=False do not contribute FAST loss."""
+ vocab_size = 8
+ action_tokens = torch.tensor([[1, 2, 3, 4], [1, 2, 5, 6]])
+ action_code_mask = torch.tensor([[False, False, True, True], [False, False, True, True]])
+ predict_actions = torch.tensor([True, False])
+
+ logits = torch.zeros(2, action_tokens.shape[1], vocab_size)
+ logits[0, 1, 3] = 10.0
+ logits[0, 2, 4] = 10.0
+ # Bad predictions in the masked sample should not matter.
+ logits[1, 1, 7] = 10.0
+ logits[1, 2, 7] = 10.0
+
+ loss = _fast_ce(logits, action_tokens, action_code_mask, predict_actions)
+ expected = F.cross_entropy(
+ torch.stack([logits[0, 1], logits[0, 2]]),
+ torch.tensor([3, 4]),
+ reduction="mean",
+ )
+
+ # Allow the fused GPU kernel's ~1e-7 difference on small losses.
+ assert torch.allclose(loss, expected, atol=1e-5, rtol=1e-3)
+
+
+def test_fast_ce_returns_zero_when_no_action_code_positions_are_valid():
+ logits = torch.randn(2, 4, 8, requires_grad=True)
+ action_tokens = torch.tensor([[1, 2, 3, 4], [1, 2, 5, 6]])
+ action_code_mask = torch.zeros_like(action_tokens, dtype=torch.bool)
+
+ loss = _fast_ce(logits, action_tokens, action_code_mask, predict_actions_t=None)
+
+ assert loss.item() == 0
+ loss.backward()
+ assert logits.grad is not None
+
+
+def test_fast_ce_averages_each_action_sample_equally():
+ torch.manual_seed(0)
+ hidden = torch.randn(2, 5, 8)
+ lm_head_weight = torch.eye(8)
+ action_tokens = torch.tensor([[1, 2, 0, 0, 0], [1, 3, 4, 5, 6]])
+ action_code_mask = torch.tensor([[False, True, False, False, False], [False, True, True, True, True]])
+
+ loss = _fast_lin_ce(
+ hidden,
+ lm_head_weight,
+ action_tokens,
+ action_code_mask,
+ predict_actions_t=None,
+ reduction="mean",
+ )
+ per_sample = _fast_lin_ce(
+ hidden,
+ lm_head_weight,
+ action_tokens,
+ action_code_mask,
+ predict_actions_t=None,
+ reduction="none",
+ )
+
+ assert torch.allclose(loss, per_sample.mean())
+
+
+def test_pi052_rejects_fast_loss_without_recipe():
+ config = SimpleNamespace(recipe_path=None, enable_fast_action_loss=True)
+
+ with pytest.raises(ValueError, match="recipe_path"):
+ make_pi052_pre_post_processors(config)
+
+
+def test_pi052_rejects_missing_fast_batch_keys():
+ policy = PI052Policy.__new__(PI052Policy)
+ nn.Module.__init__(policy)
+ policy.config = SimpleNamespace(
+ enable_fast_action_loss=True,
+ fast_action_loss_weight=1.0,
+ flow_loss_weight=0.0,
+ text_loss_weight=1.0,
+ )
+ batch = {
+ "text_labels": torch.tensor([[1, 2]]),
+ "predict_actions": torch.tensor([True]),
+ }
+
+ with pytest.raises(ValueError, match="FAST action loss is enabled"):
+ policy.forward(batch)
diff --git a/tests/policies/pi052/test_pi052_fit_fast_tokenizer.py b/tests/policies/pi052/test_pi052_fit_fast_tokenizer.py
new file mode 100644
index 000000000..b4c86d180
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_fit_fast_tokenizer.py
@@ -0,0 +1,122 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import numpy as np
+import pytest
+
+from lerobot.policies.pi052.fit_fast_tokenizer import (
+ _apply_relative_actions,
+ _dataset_signature,
+ _is_global_leader,
+ _normalize_actions,
+ _select_episode_indices,
+ _validate_fast_reconstruction,
+)
+
+
+def test_fast_tokenizer_fit_uses_training_mean_std_normalization():
+ actions = np.array([[[1.0, 7.0], [3.0, 3.0]]], dtype=np.float32)
+ stats = {"mean": [2.0, 5.0], "std": [0.5, 2.0]}
+
+ normalized = _normalize_actions(actions, "MEAN_STD", stats)
+
+ np.testing.assert_allclose(normalized, [[[-2.0, 1.0], [2.0, -1.0]]])
+
+
+def test_fast_tokenizer_fit_quantiles_match_training_without_clipping():
+ actions = np.array([[[-1.0], [3.0]]], dtype=np.float32)
+ stats = {"q01": [0.0], "q99": [2.0]}
+
+ normalized = _normalize_actions(actions, "QUANTILES", stats)
+
+ np.testing.assert_allclose(normalized, [[[-2.0], [2.0]]])
+
+
+def test_fast_tokenizer_cache_signature_tracks_stats_and_episode_selection():
+ kwargs = {
+ "dataset_repo_id": "org/dataset",
+ "base_tokenizer_name": "physical-intelligence/fast",
+ "n_samples": 100,
+ "chunk_size": 20,
+ "normalization_mode": "QUANTILES",
+ "dataset_revision": "main",
+ "episodes": [1, 2, 3],
+ "exclude_episodes": [2],
+ "use_relative_actions": False,
+ "relative_action_mask": None,
+ }
+
+ first = _dataset_signature(**kwargs, action_stats={"q01": [0.0], "q99": [1.0]})
+ changed_stats = _dataset_signature(**kwargs, action_stats={"q01": [0.0], "q99": [2.0]})
+ changed_selection = _dataset_signature(
+ **{**kwargs, "exclude_episodes": [2, 3]},
+ action_stats={"q01": [0.0], "q99": [1.0]},
+ )
+
+ assert first != changed_stats
+ assert first != changed_selection
+
+
+def test_fast_tokenizer_uses_only_global_rank_zero(monkeypatch):
+ monkeypatch.setenv("RANK", "8")
+ monkeypatch.setenv("LOCAL_RANK", "0")
+ assert not _is_global_leader()
+
+ monkeypatch.setenv("RANK", "0")
+ assert _is_global_leader()
+
+
+def test_fast_tokenizer_episode_selection_applies_allowlist_and_exclusions():
+ selected = _select_episode_indices([0, 1, 2, 3], episodes=[1, 2, 3], exclude_episodes=[2])
+
+ assert selected == [1, 3]
+
+
+def test_fast_tokenizer_relative_actions_match_training_transform():
+ actions = np.array([[[2.0, 10.0], [3.0, 11.0]]], dtype=np.float32)
+ states = np.array([[1.0, 4.0]], dtype=np.float32)
+
+ relative = _apply_relative_actions(actions, states, [True, False])
+
+ np.testing.assert_allclose(relative, [[[1.0, 10.0], [2.0, 11.0]]])
+
+
+class _RoundTripTokenizer:
+ def __init__(self, offset: float = 0.0):
+ self.offset = offset
+
+ def __call__(self, actions):
+ return actions
+
+ def decode(self, tokens):
+ return tokens + self.offset
+
+
+def test_fast_tokenizer_reconstruction_validation_reports_error():
+ actions = np.arange(24, dtype=np.float32).reshape(2, 3, 4) / 24
+
+ report, decoded = _validate_fast_reconstruction(_RoundTripTokenizer(0.05), actions, 0.1, 0.1)
+
+ np.testing.assert_allclose(decoded, actions + 0.05)
+ assert report["reconstruction_rmse"] == pytest.approx(0.05)
+ assert report["max_dim_rmse"] == pytest.approx(0.05)
+
+
+def test_fast_tokenizer_reconstruction_validation_rejects_large_error():
+ actions = np.arange(24, dtype=np.float32).reshape(2, 3, 4) / 24
+
+ with pytest.raises(RuntimeError, match="exceeds the configured limit"):
+ _validate_fast_reconstruction(_RoundTripTokenizer(0.25), actions, 0.1, 0.2)
diff --git a/tests/policies/pi052/test_pi052_flex_attention.py b/tests/policies/pi052/test_pi052_flex_attention.py
new file mode 100644
index 000000000..338692e36
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_flex_attention.py
@@ -0,0 +1,65 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import logging
+
+import pytest
+import torch
+
+pytest.importorskip("transformers")
+
+import lerobot.policies.pi052.modeling_pi052 as modeling_pi052 # noqa: E402
+from lerobot.policies.pi052.configuration_pi052 import PI052Config # noqa: E402
+
+
+def test_flex_backend_skips_non_cuda_without_initializing(monkeypatch):
+ monkeypatch.setattr(modeling_pi052, "_flex_fns", None)
+ monkeypatch.setattr(torch, "compile", lambda *args, **kwargs: pytest.fail("torch.compile was called"))
+ monkeypatch.setattr(
+ torch.cuda,
+ "get_device_properties",
+ lambda *args, **kwargs: pytest.fail("CUDA properties were queried"),
+ )
+
+ assert modeling_pi052._get_flex_fns(torch.device("cpu")) is None
+ assert modeling_pi052._get_flex_kernel_options(torch.device("cpu")) is None
+ assert modeling_pi052._flex_fns is None
+
+
+def test_flex_initialization_failure_falls_back(monkeypatch, caplog):
+ monkeypatch.setattr(modeling_pi052, "_flex_fns", None)
+ monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
+
+ def fail_compile(*args, **kwargs):
+ raise RuntimeError("compile failed")
+
+ monkeypatch.setattr(torch, "compile", fail_compile)
+
+ with caplog.at_level(logging.WARNING, logger=modeling_pi052.__name__):
+ assert modeling_pi052._get_flex_fns(torch.device("cuda", 0)) is None
+
+ assert modeling_pi052._flex_fns is False
+ assert "FlexAttention unavailable" in caplog.text
+
+
+def test_flex_rejects_single_repeat_configuration():
+ with pytest.raises(ValueError, match="use_flex_attention requires flow_num_repeats > 1"):
+ PI052Config(use_flex_attention=True, flow_num_repeats=1)
+
+
+def test_flex_accepts_amortized_repeat_configuration():
+ config = PI052Config(use_flex_attention=True, flow_num_repeats=5)
+ assert config.use_flex_attention
diff --git a/tests/policies/pi052/test_pi052_import.py b/tests/policies/pi052/test_pi052_import.py
new file mode 100644
index 000000000..b298d39fe
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_import.py
@@ -0,0 +1,27 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import subprocess
+import sys
+
+
+def test_pi052_config_import_does_not_load_model_or_dataset_processor():
+ code = """
+import sys
+from lerobot.policies import PI052Config
+assert PI052Config.__name__ == "PI052Config"
+assert "lerobot.policies.pi052.modeling_pi052" not in sys.modules
+assert "lerobot.policies.pi052.processor_pi052" not in sys.modules
+"""
+ subprocess.run([sys.executable, "-c", code], check=True)
diff --git a/tests/policies/pi052/test_pi052_joint_sequence.py b/tests/policies/pi052/test_pi052_joint_sequence.py
new file mode 100644
index 000000000..a0fa22892
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_joint_sequence.py
@@ -0,0 +1,148 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for PI052 joint-sequence (paper-style) subtask conditioning.
+
+Joint recipes train the subtask text and the action losses in one sequence,
+with the supervised subtask span attended causally. At inference the same
+layout is rebuilt around the *generated* subtask, so these tests pin:
+
+- the inference-side encoder produces the same token ids and target positions
+ as the training-time tokenizer step for the same messages;
+- OR-ing causal marks into a prefix reproduces the training-time attention
+ pattern (prompt cannot see the subtask; subtask is causal over itself);
+- the joint recipe file stays a valid message recipe;
+- the FAST id mapping with the default ``fast_skip_tokens`` stays clear of
+ PaliGemma's ```` range so VQA and FAST supervision never collide.
+"""
+
+from pathlib import Path
+
+import torch
+
+from lerobot.configs.recipe import TrainingRecipe
+from lerobot.policies.pi052.text_processor_pi052 import (
+ PI052TextTokenizerStep,
+ encode_prompt_with_targets,
+)
+
+
+class _CharTokenizer:
+ """Char-level stub: 1 char = 1 token, so offsets are trivially aligned."""
+
+ pad_token_id = 0
+ eos_token = "\x1f" # unit separator — a 1-char "EOS" for testing
+
+ def __call__(self, text, max_length=None, padding=None, return_tensors=None, **kwargs):
+ limit = max_length if max_length is not None else len(text)
+ ids = [ord(c) % 251 + 1 for c in text[:limit]]
+ offsets = [(i, i + 1) for i in range(len(ids))]
+ attention = [1] * len(ids)
+ if padding == "max_length" and max_length is not None and len(ids) < max_length:
+ pad = max_length - len(ids)
+ ids += [self.pad_token_id] * pad
+ offsets += [(0, 0)] * pad
+ attention += [0] * pad
+ return {
+ "input_ids": torch.tensor([ids], dtype=torch.long),
+ "attention_mask": torch.tensor([attention], dtype=torch.long),
+ "offset_mapping": torch.tensor([offsets], dtype=torch.long),
+ }
+
+
+_MESSAGES = [
+ {"role": "user", "content": "fold the towel"},
+ {"role": "assistant", "content": "grab the near corner"},
+]
+
+
+def test_encode_prompt_with_targets_matches_training_labels():
+ tokenizer = _CharTokenizer()
+
+ step = PI052TextTokenizerStep(max_length=120)
+ step._tokenizer = tokenizer
+ train_ids, train_attn, labels, predict_actions, _prompt = step._encode_messages(
+ tokenizer,
+ [dict(m) for m in _MESSAGES],
+ message_streams=["low_level", "low_level"],
+ target_indices=[1],
+ complementary={},
+ )
+ assert bool(predict_actions)
+
+ ids, attn, marks = encode_prompt_with_targets(tokenizer, [dict(m) for m in _MESSAGES], [1])
+
+ n = int(attn.sum())
+ assert n == int(train_attn.sum())
+ assert torch.equal(ids[0, :n], train_ids[:n])
+ # Causal marks at inference must cover exactly the supervised label span.
+ assert torch.equal(marks[0, :n], labels[:n] != -100)
+ assert marks.any(), "the assistant target span must be marked"
+ # The user turn must stay unmarked (bidirectional prompt).
+ user_len = len("User: fold the towel\n")
+ assert not marks[0, :user_len].any()
+
+
+def test_apply_causal_language_marks_reproduces_training_mask():
+ from lerobot.policies.pi05.modeling_pi05 import make_att_2d_masks
+ from lerobot.policies.pi052.modeling_pi052 import _apply_causal_language_marks
+
+ n_img, n_lang = 4, 8
+ prefix_len = n_img + n_lang
+ pad = torch.ones((1, prefix_len), dtype=torch.bool)
+ att = torch.zeros((1, prefix_len), dtype=torch.bool)
+ # Subtask span = language positions 5..7 (prefix positions 9..11).
+ marks = torch.zeros((1, n_lang), dtype=torch.bool)
+ marks[0, 5:8] = True
+
+ att_marked = _apply_causal_language_marks(att, marks)
+ att_2d = make_att_2d_masks(pad, att_marked)[0]
+
+ subtask = [n_img + 5, n_img + 6, n_img + 7]
+ # Prompt and images never see the subtask.
+ for q in range(n_img + 5):
+ for k in subtask:
+ assert not att_2d[q, k], f"prompt position {q} must not attend subtask position {k}"
+ # Subtask tokens see the full prompt and earlier subtask tokens only.
+ for qi, q in enumerate(subtask):
+ for k in range(n_img + 5):
+ assert att_2d[q, k]
+ for ki, k in enumerate(subtask):
+ assert bool(att_2d[q, k]) == (ki <= qi)
+
+
+def test_joint_recipe_is_a_valid_message_recipe():
+ recipe_path = Path(__file__).parents[3] / "src" / "lerobot" / "configs" / "recipes" / "subtask_joint.yaml"
+ recipe = TrainingRecipe.from_yaml(recipe_path)
+ assert recipe.messages is not None and len(recipe.messages) == 2
+ assert all(turn.stream == "low_level" for turn in recipe.messages)
+ assert not recipe.messages[0].target
+ assert recipe.messages[1].target
+ assert recipe.messages[1].if_present == "subtask"
+
+
+def test_default_fast_mapping_clears_loc_and_seg_ranges():
+ from lerobot.policies.pi052.configuration_pi052 import PI052Config
+ from lerobot.policies.pi052.modeling_pi052 import _FAST_ACTION_VOCAB_SIZE
+
+ skip = PI052Config.__dataclass_fields__["fast_skip_tokens"].default
+ assert skip == 1152
+
+ paligemma_vocab = 257152
+ fast_ids = paligemma_vocab - 1 - skip - torch.arange(_FAST_ACTION_VOCAB_SIZE)
+ # Below the range [256000, 257024) and the range [257024, 257152).
+ assert int(fast_ids.max()) < 256000
+ assert int(fast_ids.min()) >= 0
diff --git a/tests/policies/pi052/test_pi052_runtime_adapter.py b/tests/policies/pi052/test_pi052_runtime_adapter.py
new file mode 100644
index 000000000..dddeb58c1
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_runtime_adapter.py
@@ -0,0 +1,85 @@
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from types import SimpleNamespace
+
+from lerobot.policies.pi052.inference.pi052_adapter import PI052PolicyAdapter
+from lerobot.runtime import RuntimeState
+from lerobot.runtime.adapter import split_plan_and_say
+
+
+def test_pi052_adapter_builds_recipe_prompts_from_runtime_state():
+ adapter = PI052PolicyAdapter(policy=object())
+ state = RuntimeState(
+ task="clean the kitchen",
+ language_context={"memory": "cup moved", "plan": "pick then place"},
+ extra={"prior_subtask": "pick the cup"},
+ )
+
+ assert adapter.build_messages("subtask", state) == [{"role": "user", "content": "clean the kitchen"}]
+ assert adapter.build_messages("memory", state) == [
+ {"role": "user", "content": "clean the kitchen"},
+ {"role": "assistant", "content": "Previous memory: cup moved"},
+ {"role": "user", "content": "Completed subtask: pick the cup"},
+ ]
+ assert adapter.build_messages("interjection", state, user_text="wait") == [
+ {"role": "user", "content": "clean the kitchen"},
+ {"role": "assistant", "content": "Previous plan:\npick then place"},
+ {"role": "user", "content": "wait"},
+ ]
+
+
+def test_pi052_adapter_strips_say_markers_from_plan_text():
+ adapter = PI052PolicyAdapter(policy=object())
+ text = "Move to the sink. heading to the sink"
+
+ assert split_plan_and_say(text) == ("Move to the sink.", "heading to the sink")
+ assert adapter.plan_from_text(text) == "Move to the sink."
+
+
+def test_rollout_language_cli_smoke_does_not_load_model(monkeypatch):
+ """lerobot-rollout dispatches language flags to the adapter-based runtime."""
+ from lerobot.runtime import cli
+ from lerobot.scripts import lerobot_rollout
+
+ fake_policy = SimpleNamespace(config=SimpleNamespace(device="cpu", type="pi052"))
+
+ monkeypatch.setattr(
+ cli,
+ "_load_policy_and_preprocessor",
+ lambda policy_path, **kwargs: (fake_policy, None, None),
+ )
+ monkeypatch.setattr(cli, "_run_repl", lambda runtime, **kwargs: 0)
+
+ assert lerobot_rollout.main(["--policy.path=fake", "--no_robot", "--task=clean", "--max_ticks=0"]) == 0
+
+
+def test_rollout_language_dispatch_preserves_standard_molmoact2_path(monkeypatch):
+ """MolmoAct2 only opts into open prompting when a language flag is present."""
+ from lerobot.scripts import lerobot_rollout
+
+ standard = [
+ "--policy.path=lerobot/MolmoAct2-SO100_101-LeRobot",
+ "--robot.type=so101_follower",
+ "--task=pick up the cube",
+ ]
+ assert not lerobot_rollout._uses_language_runtime(standard)
+ assert lerobot_rollout._uses_language_runtime([*standard, "--direct_subtask"])
+ assert lerobot_rollout._uses_language_runtime(["--policy.path=lerobot/pi052_robocasa", "--sim"])
+
+ standard_calls = []
+ monkeypatch.setattr(lerobot_rollout, "register_third_party_plugins", lambda: None)
+ monkeypatch.setattr(lerobot_rollout, "rollout", lambda: standard_calls.append(True))
+ lerobot_rollout.main(standard)
+ assert standard_calls == [True]
diff --git a/tests/policies/pi052/test_pi052_sdpa_attention.py b/tests/policies/pi052/test_pi052_sdpa_attention.py
new file mode 100644
index 000000000..02fc6d9cb
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_sdpa_attention.py
@@ -0,0 +1,147 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Numerical-parity tests for the SDPA attention port.
+
+``pi05`` / ``pi052`` replaced the per-layer call from
+``modeling_gemma.eager_attention_forward`` with
+``sdpa_attention_forward`` (PyTorch SDPA + GQA repeat). The forward
+output must be bit-equivalent (within bf16 tolerance) on the masks
+this model actually uses — block-bidirectional with an arbitrary
+additive bias — otherwise we silently change training behaviour.
+"""
+
+from types import SimpleNamespace
+
+import pytest
+import torch
+
+pytest.importorskip("transformers")
+
+from transformers.models.gemma import modeling_gemma # noqa: E402
+
+from lerobot.policies.pi052.modeling_pi052 import make_att_2d_masks # noqa: E402
+from lerobot.policies.pi_gemma import sdpa_attention_forward # noqa: E402
+from lerobot.utils.constants import OPENPI_ATTENTION_MASK_VALUE # noqa: E402
+
+
+def _mock_self_attn(num_kv_groups: int, training: bool = False):
+ """Bare module surface that both forwards read."""
+ return SimpleNamespace(
+ num_key_value_groups=num_kv_groups,
+ training=training,
+ )
+
+
+def _build_inputs(
+ bsize: int,
+ num_heads: int,
+ num_kv_heads: int,
+ seq_len: int,
+ head_dim: int,
+ dtype: torch.dtype,
+ seed: int = 0,
+):
+ g = torch.Generator(device="cpu").manual_seed(seed)
+ q = torch.randn(bsize, num_heads, seq_len, head_dim, dtype=dtype, generator=g)
+ k = torch.randn(bsize, num_kv_heads, seq_len, head_dim, dtype=dtype, generator=g)
+ v = torch.randn(bsize, num_kv_heads, seq_len, head_dim, dtype=dtype, generator=g)
+ return q, k, v
+
+
+def _block_bidirectional_mask(
+ bsize: int, seq_len: int, block_sizes: list[int], dtype: torch.dtype
+) -> torch.Tensor:
+ """Mimic ``_prepare_attention_masks_4d`` on a block layout that
+ matches ``[images, language, suffix]`` from ``embed_prefix`` +
+ ``embed_suffix``: every block bidirectional internally, later
+ blocks visible to earlier ones via the cumulative-block rule.
+ """
+ assert sum(block_sizes) == seq_len
+ att_marks = []
+ for i, n in enumerate(block_sizes):
+ att_marks += [1 if i > 0 else 0] + [0] * (n - 1)
+ pad = torch.ones(bsize, seq_len, dtype=torch.bool)
+ att = torch.tensor(att_marks, dtype=torch.bool)[None].expand(bsize, seq_len)
+ att_2d = make_att_2d_masks(pad, att)
+ bias = torch.where(
+ att_2d[:, None, :, :],
+ torch.zeros((), dtype=dtype),
+ torch.tensor(OPENPI_ATTENTION_MASK_VALUE, dtype=dtype),
+ )
+ return bias
+
+
+@pytest.mark.parametrize(
+ "num_heads,num_kv_heads,head_dim",
+ [
+ (8, 1, 256), # gemma_2b / paligemma config
+ (8, 8, 64), # MHA control (no GQA repeat)
+ ],
+)
+def test_sdpa_parity_with_eager_block_bidirectional(num_heads, num_kv_heads, head_dim):
+ """SDPA forward output matches the eager softmax(QK^T)@V on the
+ block-bidirectional mask layout pi05 actually uses."""
+ bsize, seq_len = 2, 13
+ block_sizes = [4, 5, 4] # images, language, suffix-style blocks
+ dtype = torch.float32 # cpu math kernel — keep fp32 for tight tol
+ scaling = head_dim**-0.5
+
+ q, k, v = _build_inputs(bsize, num_heads, num_kv_heads, seq_len, head_dim, dtype)
+ mask = _block_bidirectional_mask(bsize, seq_len, block_sizes, dtype)
+
+ module = _mock_self_attn(num_heads // num_kv_heads)
+
+ out_eager, _ = modeling_gemma.eager_attention_forward(module, q, k, v, mask, scaling)
+ out_sdpa, _ = sdpa_attention_forward(module, q, k, v, mask, scaling)
+ assert out_eager.shape == out_sdpa.shape
+ torch.testing.assert_close(out_sdpa, out_eager, atol=1e-5, rtol=1e-4)
+
+
+def test_sdpa_parity_bf16():
+ """bf16 path — looser tolerance, must still match eager."""
+ bsize, num_heads, num_kv_heads, seq_len, head_dim = 2, 8, 1, 17, 256
+ scaling = head_dim**-0.5
+ q, k, v = _build_inputs(bsize, num_heads, num_kv_heads, seq_len, head_dim, torch.bfloat16)
+ mask = _block_bidirectional_mask(bsize, seq_len, [5, 6, 6], torch.bfloat16)
+ module = _mock_self_attn(num_heads // num_kv_heads)
+
+ out_eager, _ = modeling_gemma.eager_attention_forward(module, q, k, v, mask, scaling)
+ out_sdpa, _ = sdpa_attention_forward(module, q, k, v, mask, scaling)
+ torch.testing.assert_close(out_sdpa, out_eager, atol=2e-2, rtol=2e-2)
+
+
+def test_sdpa_parity_backward():
+ """Gradients flow through SDPA and match the eager path within
+ bf16 tolerance — critical for any training-side parity claim."""
+ bsize, num_heads, num_kv_heads, seq_len, head_dim = 1, 4, 2, 9, 32
+ scaling = head_dim**-0.5
+ q, k, v = _build_inputs(bsize, num_heads, num_kv_heads, seq_len, head_dim, torch.float32)
+ q.requires_grad_(True)
+ k.requires_grad_(True)
+ v.requires_grad_(True)
+ mask = _block_bidirectional_mask(bsize, seq_len, [3, 3, 3], torch.float32)
+ module = _mock_self_attn(num_heads // num_kv_heads)
+
+ out_e, _ = modeling_gemma.eager_attention_forward(module, q, k, v, mask, scaling)
+ g_q_e, g_k_e, g_v_e = torch.autograd.grad(out_e.sum(), [q, k, v])
+
+ out_s, _ = sdpa_attention_forward(module, q, k, v, mask, scaling)
+ g_q_s, g_k_s, g_v_s = torch.autograd.grad(out_s.sum(), [q, k, v])
+
+ torch.testing.assert_close(g_q_s, g_q_e, atol=1e-5, rtol=1e-4)
+ torch.testing.assert_close(g_k_s, g_k_e, atol=1e-5, rtol=1e-4)
+ torch.testing.assert_close(g_v_s, g_v_e, atol=1e-5, rtol=1e-4)
diff --git a/tests/policies/pi052/test_pi052_text_processor.py b/tests/policies/pi052/test_pi052_text_processor.py
new file mode 100644
index 000000000..8e1ea02c6
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_text_processor.py
@@ -0,0 +1,223 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+"""Tests for PI052's text tokenizer.
+
+Covers ``say`` tool-call flattening (PaliGemma's flat prompt has no
+structured tool calls, so a ``say`` call must be serialized into a
+``...`` text marker) and EOS-termination supervision (the
+supervised target span must end with an EOS token so the LM head learns
+to stop instead of rambling to ``max_length`` at inference).
+"""
+
+import torch
+
+from lerobot.configs.recipe import MessageTurn, TrainingRecipe
+from lerobot.policies.pi052.text_processor_pi052 import (
+ PI052TextTokenizerStep,
+ _flatten_say_tool_calls,
+ _format_messages,
+)
+from lerobot.processor import PolicyProcessorPipeline
+from lerobot.processor.render_messages_processor import RenderMessagesStep
+from lerobot.types import TransitionKey
+from lerobot.utils.constants import (
+ OBS_LANGUAGE_ATTENTION_MASK,
+ OBS_LANGUAGE_TOKENS,
+ POLICY_PREPROCESSOR_DEFAULT_NAME,
+)
+
+
+def _say_call(text):
+ return {"type": "function", "function": {"name": "say", "arguments": {"text": text}}}
+
+
+def test_flatten_appends_say_marker_and_drops_tool_calls():
+ msg = {"role": "assistant", "content": "Heading to the cube.", "tool_calls": [_say_call("On it!")]}
+ out = _flatten_say_tool_calls(msg)
+ assert "tool_calls" not in out
+ assert out["content"] == "Heading to the cube.\nOn it!"
+
+
+def test_flatten_marker_only_when_content_empty_or_none():
+ out = _flatten_say_tool_calls({"role": "assistant", "tool_calls": [_say_call("hi")]})
+ assert out["content"] == "hi"
+
+
+def test_flatten_accepts_json_string_arguments():
+ call = {"type": "function", "function": {"name": "say", "arguments": '{"text": "hello there"}'}}
+ out = _flatten_say_tool_calls({"role": "assistant", "content": "p", "tool_calls": [call]})
+ assert out["content"] == "p\nhello there"
+
+
+def test_flatten_leaves_messages_without_tool_calls_untouched():
+ msg = {"role": "assistant", "content": "just a plan"}
+ assert _flatten_say_tool_calls(msg) == msg
+
+
+def test_flatten_drops_non_say_tool_calls_but_keeps_content():
+ weather = {"type": "function", "function": {"name": "check_weather", "arguments": {}}}
+ out = _flatten_say_tool_calls({"role": "assistant", "content": "plan only", "tool_calls": [weather]})
+ assert out["content"] == "plan only"
+ assert "tool_calls" not in out
+
+
+def test_format_messages_appends_eos_to_target_turns_only():
+ msgs = [
+ {"role": "user", "content": "pick cube"},
+ {"role": "assistant", "content": "move to cube"},
+ ]
+ prompt, spans = _format_messages(msgs, target_indices=[1], eos_token="")
+ # EOS is appended to the supervised target (assistant) turn only.
+ assert prompt == "User: pick cube\nAssistant: move to cube\n"
+ # The user span is unchanged; the target span covers content + EOS.
+ assert prompt[spans[0][0] : spans[0][1]] == "pick cube"
+ assert prompt[spans[1][0] : spans[1][1]] == "move to cube"
+
+
+def test_format_messages_without_eos_args_is_unchanged():
+ """Inference callers omit target_indices / eos_token — no EOS baked in."""
+ prompt, spans = _format_messages([{"role": "user", "content": "hi"}])
+ assert prompt == "User: hi\n"
+ assert prompt[spans[0][0] : spans[0][1]] == "hi"
+
+
+def test_pi052_steps_roundtrip_through_standard_pipeline_loader(tmp_path):
+ recipe = TrainingRecipe(messages=[MessageTurn(role="user", content="${task}", stream="low_level")])
+ pipeline = PolicyProcessorPipeline(
+ steps=[
+ RenderMessagesStep(recipe),
+ PI052TextTokenizerStep(
+ tokenizer_name="custom-tokenizer",
+ max_length=77,
+ plan_dropout_prob=0.2,
+ dropout_seed=3,
+ ),
+ ],
+ name=POLICY_PREPROCESSOR_DEFAULT_NAME,
+ )
+ pipeline.save_pretrained(tmp_path)
+
+ loaded = PolicyProcessorPipeline.from_pretrained(
+ tmp_path, config_filename=f"{POLICY_PREPROCESSOR_DEFAULT_NAME}.json"
+ )
+
+ assert loaded.steps[0].recipe == recipe
+ assert loaded.steps[1].tokenizer_name == "custom-tokenizer"
+ assert loaded.steps[1].max_length == 77
+ assert loaded.steps[1].plan_dropout_prob == 0.2
+ assert loaded.steps[1].dropout_seed == 3
+
+
+def _eos_char_id() -> int:
+ """Token id _CharTokenizer assigns to its 1-char EOS."""
+ return ord("\x1f") % 251 + 1
+
+
+def test_pi052_text_tokenizer_supervises_eos_at_target_end():
+ """The appended EOS is the last supervised label on a target turn —
+ that's the signal that teaches the LM head to stop. The trailing
+ newline right after it stays unsupervised (-100)."""
+ step = PI052TextTokenizerStep(max_length=64)
+ step._tokenizer = _CharTokenizer()
+ transition = {
+ TransitionKey.OBSERVATION: {},
+ TransitionKey.COMPLEMENTARY_DATA: {
+ "messages": [
+ {"role": "user", "content": "pick cube"},
+ {"role": "assistant", "content": "move to cube"},
+ ],
+ "target_message_indices": [1],
+ "message_streams": ["high_level", "high_level"],
+ "index": torch.tensor(10),
+ },
+ }
+ out = step(transition)
+ ids = out[TransitionKey.OBSERVATION][OBS_LANGUAGE_TOKENS][0]
+ labels = out[TransitionKey.COMPLEMENTARY_DATA]["text_labels"][0]
+
+ supervised = (labels != -100).nonzero().flatten().tolist()
+ assert supervised, "target turn produced no supervised labels"
+ last = supervised[-1]
+ # The last supervised token is the appended EOS.
+ assert int(ids[last]) == _eos_char_id()
+ assert int(labels[last]) == _eos_char_id()
+ # The token right after the EOS (the trailing newline) is NOT supervised.
+ assert int(labels[last + 1]) == -100
+
+
+class _CharTokenizer:
+ pad_token_id = 0
+ eos_token = "\x1f" # unit separator — a 1-char "EOS" for testing
+
+ def __call__(
+ self,
+ text,
+ max_length,
+ padding,
+ truncation,
+ return_tensors,
+ return_offsets_mapping,
+ padding_side,
+ ):
+ ids = [ord(c) % 251 + 1 for c in text[:max_length]]
+ offsets = [(i, i + 1) for i in range(len(ids))]
+ attention = [1] * len(ids)
+ if padding == "max_length" and len(ids) < max_length:
+ pad = max_length - len(ids)
+ ids += [self.pad_token_id] * pad
+ offsets += [(0, 0)] * pad
+ attention += [0] * pad
+ return {
+ "input_ids": torch.tensor([ids], dtype=torch.long),
+ "attention_mask": torch.tensor([attention], dtype=torch.long),
+ "offset_mapping": torch.tensor([offsets], dtype=torch.long),
+ }
+
+ def decode(self, token_ids, skip_special_tokens=False):
+ return "".join(chr(max(int(i) - 1, 0)) for i in token_ids if int(i) != self.pad_token_id)
+
+
+def test_pi052_text_tokenizer_handles_batched_rendered_messages():
+ step = PI052TextTokenizerStep(max_length=64)
+ step._tokenizer = _CharTokenizer()
+
+ transition = {
+ TransitionKey.OBSERVATION: {},
+ TransitionKey.COMPLEMENTARY_DATA: {
+ "messages": [
+ [
+ {"role": "user", "content": "pick cube"},
+ {"role": "assistant", "content": "move to cube"},
+ ],
+ [{"role": "user", "content": "open drawer"}],
+ ],
+ "target_message_indices": [[1], []],
+ "message_streams": [["high_level", "high_level"], ["low_level"]],
+ "index": torch.tensor([10, 11]),
+ },
+ }
+
+ out = step(transition)
+ obs = out[TransitionKey.OBSERVATION]
+ comp = out[TransitionKey.COMPLEMENTARY_DATA]
+
+ assert obs[OBS_LANGUAGE_TOKENS].shape == (2, 64)
+ assert obs[OBS_LANGUAGE_ATTENTION_MASK].shape == (2, 64)
+ assert comp["text_labels"].shape == (2, 64)
+ assert comp["predict_actions"].tolist() == [False, True]
+ assert (comp["text_labels"][0] != -100).any()
+ assert not (comp["text_labels"][1] != -100).any()
diff --git a/tests/policies/pi052/test_pi052_vision_checkpointing.py b/tests/policies/pi052/test_pi052_vision_checkpointing.py
new file mode 100644
index 000000000..430b8afb4
--- /dev/null
+++ b/tests/policies/pi052/test_pi052_vision_checkpointing.py
@@ -0,0 +1,141 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from types import MethodType, SimpleNamespace
+
+import pytest
+import torch
+from torch import nn
+
+pytest.importorskip("transformers")
+
+from lerobot.policies.pi052.modeling_pi052 import PI05Pytorch
+
+
+class _MockVisionTower:
+ def __init__(self):
+ self.enable_kwargs = None
+ self.disable_calls = 0
+
+ def gradient_checkpointing_enable(self, **kwargs):
+ self.enable_kwargs = kwargs
+
+ def gradient_checkpointing_disable(self):
+ self.disable_calls += 1
+
+
+def _checkpoint_model():
+ tower = _MockVisionTower()
+ language_model = SimpleNamespace(gradient_checkpointing=False)
+ expert_model = SimpleNamespace(gradient_checkpointing=False)
+ model = PI05Pytorch.__new__(PI05Pytorch)
+ nn.Module.__init__(model)
+ model.gradient_checkpointing_enabled = False
+ model.paligemma_with_expert = SimpleNamespace(
+ paligemma=SimpleNamespace(model=SimpleNamespace(language_model=language_model, vision_tower=tower)),
+ gemma_expert=SimpleNamespace(model=expert_model),
+ )
+ return model, tower, language_model, expert_model
+
+
+def test_gradient_checkpointing_uses_vision_tower_layer_api():
+ model, tower, language_model, expert_model = _checkpoint_model()
+
+ PI05Pytorch.gradient_checkpointing_enable(model)
+
+ assert model.gradient_checkpointing_enabled
+ assert language_model.gradient_checkpointing
+ assert expert_model.gradient_checkpointing
+ assert tower.enable_kwargs == {"gradient_checkpointing_kwargs": {"use_reentrant": False}}
+
+ PI05Pytorch.gradient_checkpointing_disable(model)
+
+ assert not model.gradient_checkpointing_enabled
+ assert not language_model.gradient_checkpointing
+ assert not expert_model.gradient_checkpointing
+ assert tower.disable_calls == 1
+
+
+def test_siglip_layers_recompute_individually():
+ from transformers.models.siglip.configuration_siglip import SiglipVisionConfig
+ from transformers.models.siglip.modeling_siglip import SiglipVisionModel
+
+ config = SiglipVisionConfig(
+ hidden_size=16,
+ intermediate_size=32,
+ num_hidden_layers=2,
+ num_attention_heads=2,
+ num_channels=3,
+ image_size=16,
+ patch_size=8,
+ )
+ tower = SiglipVisionModel(config).train()
+ tower.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
+ calls = [0] * config.num_hidden_layers
+
+ for index, layer in enumerate(tower.vision_model.encoder.layers):
+ original_forward = layer.forward
+
+ def counted_forward(self, *args, _index=index, _forward=original_forward, **kwargs):
+ calls[_index] += 1
+ return _forward(*args, **kwargs)
+
+ layer.forward = MethodType(counted_forward, layer)
+
+ pixels = torch.randn(2, config.num_channels, config.image_size, config.image_size)
+ tower(pixels).last_hidden_state.sum().backward()
+
+ assert calls == [2] * config.num_hidden_layers
+
+
+def test_embed_prefix_does_not_wrap_the_whole_vision_tower_checkpoint():
+ model = PI05Pytorch.__new__(PI05Pytorch)
+ nn.Module.__init__(model)
+ model.config = SimpleNamespace()
+ model.gradient_checkpointing_enabled = True
+ model.train()
+
+ image_calls = []
+
+ def embed_image(image):
+ image_calls.append(image.shape)
+ return image[:, :1, 0, :2]
+
+ def embed_language_tokens(tokens):
+ return tokens.to(torch.float32).unsqueeze(-1).expand(*tokens.shape, 2)
+
+ model.paligemma_with_expert = SimpleNamespace(
+ embed_image=embed_image,
+ embed_language_tokens=embed_language_tokens,
+ )
+ outer_checkpoint_calls = []
+
+ def apply_checkpoint(func, value):
+ outer_checkpoint_calls.append(value.shape)
+ return func(value)
+
+ model._apply_checkpoint = apply_checkpoint
+
+ images = [torch.randn(2, 3, 4, 4), torch.randn(2, 3, 4, 4)]
+ image_masks = [torch.ones(2, dtype=torch.bool) for _ in images]
+ tokens = torch.ones(2, 3, dtype=torch.long)
+ token_masks = torch.ones_like(tokens, dtype=torch.bool)
+
+ embeddings, _, _ = model.embed_prefix(images, image_masks, tokens, token_masks)
+
+ assert image_calls == [image.shape for image in images]
+ assert outer_checkpoint_calls == [tokens.shape]
+ assert embeddings.shape == (2, 5, 2)
diff --git a/tests/policies/pi0_fast/test_pi0_fast_tokenizer_fit.py b/tests/policies/pi0_fast/test_pi0_fast_tokenizer_fit.py
new file mode 100644
index 000000000..5ec59fa69
--- /dev/null
+++ b/tests/policies/pi0_fast/test_pi0_fast_tokenizer_fit.py
@@ -0,0 +1,106 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from types import SimpleNamespace
+
+import pytest
+
+from lerobot.policies import factory
+from lerobot.policies.pi0_fast.configuration_pi0_fast import PI0FastConfig
+from lerobot.policies.pi052 import fit_fast_tokenizer as fit_module
+
+
+def test_pi0_fast_resolves_dataset_specific_tokenizer(monkeypatch, tmp_path):
+ config = PI0FastConfig(
+ auto_fit_fast_tokenizer=True,
+ action_tokenizer_name="base-tokenizer",
+ fast_tokenizer_cache_dir=str(tmp_path),
+ fast_tokenizer_fit_samples=17,
+ chunk_size=12,
+ n_action_steps=12,
+ )
+ received = {}
+
+ def fake_fit(**kwargs):
+ received.update(kwargs)
+ return "/cache/fitted-tokenizer"
+
+ monkeypatch.setattr(fit_module, "fit_fast_tokenizer", fake_fit)
+
+ assert fit_module.resolve_fast_tokenizer(config, "user/dataset") == "/cache/fitted-tokenizer"
+ assert received == {
+ "dataset_repo_id": "user/dataset",
+ "cache_dir": tmp_path,
+ "base_tokenizer_name": "base-tokenizer",
+ "n_samples": 17,
+ "chunk_size": 12,
+ "dataset_root": None,
+ "dataset_revision": None,
+ "episodes": None,
+ "exclude_episodes": None,
+ "normalization_mode": config.normalization_mapping["ACTION"],
+ "action_stats": None,
+ "use_relative_actions": False,
+ "relative_action_mask": None,
+ }
+
+
+def test_fast_fit_failure_is_not_silently_replaced(monkeypatch, tmp_path):
+ config = PI0FastConfig(auto_fit_fast_tokenizer=True, fast_tokenizer_cache_dir=str(tmp_path))
+ monkeypatch.setattr(
+ fit_module,
+ "fit_fast_tokenizer",
+ lambda **kwargs: (_ for _ in ()).throw(RuntimeError("fit failed")),
+ )
+
+ with pytest.raises(RuntimeError, match="fit failed"):
+ fit_module.resolve_fast_tokenizer(config, "user/dataset")
+
+
+def test_only_global_rank_zero_fits_shared_tokenizer(monkeypatch):
+ monkeypatch.setenv("RANK", "8")
+ monkeypatch.setenv("LOCAL_RANK", "0")
+ assert not fit_module._is_global_leader()
+
+ monkeypatch.setenv("RANK", "0")
+ assert fit_module._is_global_leader()
+
+
+def test_pretrained_pi0_fast_overrides_only_fitted_tokenizer(monkeypatch):
+ config = PI0FastConfig(auto_fit_fast_tokenizer=True)
+ calls = []
+
+ monkeypatch.setattr(
+ fit_module,
+ "resolve_fast_tokenizer",
+ lambda config, dataset_repo_id, *args: "/cache/fitted-tokenizer",
+ )
+
+ def fake_from_pretrained(cls, *args, **kwargs):
+ calls.append(kwargs)
+ return SimpleNamespace(steps=[])
+
+ monkeypatch.setattr(factory.PolicyProcessorPipeline, "from_pretrained", classmethod(fake_from_pretrained))
+
+ factory.make_pre_post_processors(
+ config,
+ pretrained_path="checkpoint",
+ dataset_repo_id="user/dataset",
+ )
+
+ assert calls[0]["overrides"] == {
+ "action_tokenizer_processor": {"action_tokenizer_name": "/cache/fitted-tokenizer"}
+ }
diff --git a/tests/policies/pi0_pi05/test_pi05.py b/tests/policies/pi0_pi05/test_pi05.py
index f0da2971b..44327fb6c 100644
--- a/tests/policies/pi0_pi05/test_pi05.py
+++ b/tests/policies/pi0_pi05/test_pi05.py
@@ -16,8 +16,12 @@
"""Test script to verify PI0.5 (pi05) support in PI0 policy"""
+from types import SimpleNamespace
+
import pytest
import torch
+from safetensors.torch import save_file
+from torch import nn
pytest.importorskip("transformers")
@@ -31,6 +35,26 @@ from lerobot.utils.random_utils import set_seed
from tests.utils import require_cuda, require_hf_token # noqa: E402
+class _CheckpointPolicy(PI05Policy):
+ def __init__(self, config, **kwargs):
+ nn.Module.__init__(self)
+ self.config = config
+ self.loaded_state_dict = None
+
+ def load_state_dict(self, state_dict, strict=True, assign=False):
+ self.loaded_state_dict = state_dict
+ return [], []
+
+
+def test_from_pretrained_loads_existing_single_file_checkpoint(tmp_path):
+ save_file({"weight": torch.tensor([1.0])}, tmp_path / "model.safetensors")
+
+ policy = _CheckpointPolicy.from_pretrained(tmp_path, config=SimpleNamespace())
+
+ assert policy.loaded_state_dict is not None
+ torch.testing.assert_close(policy.loaded_state_dict["model.weight"], torch.tensor([1.0]))
+
+
@require_cuda
@require_hf_token
def test_policy_instantiation():
diff --git a/tests/scripts/test_backfill_pi052_action_tokenizer.py b/tests/scripts/test_backfill_pi052_action_tokenizer.py
new file mode 100644
index 000000000..d3d664785
--- /dev/null
+++ b/tests/scripts/test_backfill_pi052_action_tokenizer.py
@@ -0,0 +1,67 @@
+#!/usr/bin/env python
+
+# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+from scripts.backfill_pi052_action_tokenizer import (
+ CHECKPOINT_DIRECTORIES,
+ DEFAULT_REPOSITORIES,
+ artifact_fingerprint,
+ make_portable_preprocessor,
+)
+
+
+def test_atomic4_backfill_covers_every_repository_and_checkpoint():
+ assert len(DEFAULT_REPOSITORIES) == 6
+ assert CHECKPOINT_DIRECTORIES == (
+ "",
+ "checkpoints/003000/pretrained_model",
+ "checkpoints/006000/pretrained_model",
+ "checkpoints/009000/pretrained_model",
+ "checkpoints/012000/pretrained_model",
+ )
+
+
+def test_backfill_embeds_recipe_and_declares_relative_tokenizer():
+ recipe = {"messages": [{"role": "user", "content": "${task}", "stream": "low_level"}]}
+ preprocessor = {
+ "name": "policy_preprocessor",
+ "steps": [
+ {
+ "registry_name": "normalizer_processor",
+ "config": {},
+ "state_file": "normalizer.safetensors",
+ },
+ {"registry_name": "render_messages_processor", "config": {"recipe": recipe}},
+ {
+ "registry_name": "action_tokenizer_processor",
+ "config": {"action_tokenizer_name": "/fsx/original/tokenizer"},
+ },
+ ],
+ }
+
+ portable = make_portable_preprocessor(preprocessor)
+
+ assert portable["steps"][1]["config"]["recipe"] == recipe
+ assert portable["steps"][2]["config"]["action_tokenizer_name"] == "action_tokenizer"
+ assert portable["steps"][2]["artifacts"] == {"action_tokenizer_name": "action_tokenizer"}
+ assert preprocessor["steps"][2]["config"]["action_tokenizer_name"].startswith("/fsx/")
+
+
+def test_artifact_fingerprint_includes_paths_and_contents():
+ first = artifact_fingerprint([("a/file", b"same"), ("b/file", b"content")])
+
+ assert first == artifact_fingerprint([("b/file", b"content"), ("a/file", b"same")])
+ assert first != artifact_fingerprint([("a/renamed", b"same"), ("b/file", b"content")])
+ assert first != artifact_fingerprint([("a/file", b"changed"), ("b/file", b"content")])
diff --git a/tests/utils/test_logging_utils.py b/tests/utils/test_logging_utils.py
index 111bd0f01..6700acaf8 100644
--- a/tests/utils/test_logging_utils.py
+++ b/tests/utils/test_logging_utils.py
@@ -37,6 +37,12 @@ class MockAccelerator:
return self._reduce_fn(tensor, reduction)
return tensor
+ def gather(self, tensor):
+ if self._reduce_fn is None:
+ return tensor.repeat(self.num_processes)
+ reduced = self._reduce_fn(tensor, "max")
+ return torch.cat([tensor.repeat(self.num_processes - 1), reduced])
+
def test_average_meter_initialization():
meter = AverageMeter("loss", ":.2f")
@@ -168,6 +174,18 @@ def test_metrics_tracker_reset_averages(mock_metrics):
assert tracker.accuracy.avg == 0.0
+def test_metrics_tracker_materializes_full_tensor_window(mock_metrics):
+ tracker = MetricsTracker(batch_size=2, num_frames=10, num_episodes=2, metrics=mock_metrics)
+ tracker.accumulate_tensor("loss", torch.tensor(1.0))
+ tracker.accumulate_tensor("loss", torch.tensor(3.0))
+
+ assert tracker.loss.count == 0
+ tracker.materialize_tensors()
+
+ assert tracker.loss.avg == pytest.approx(2.0)
+ assert tracker.loss.count == 2
+
+
def test_average_meter_invalid_reduction():
with pytest.raises(ValueError):
AverageMeter("loss", reduction="median")
diff --git a/uv.lock b/uv.lock
index a7055011e..8ef6d68fe 100644
--- a/uv.lock
+++ b/uv.lock
@@ -402,10 +402,10 @@ name = "bddl"
version = "1.0.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "jupytext" },
- { name = "networkx" },
- { name = "numpy" },
- { name = "pytest" },
+ { name = "jupytext", marker = "sys_platform == 'linux'" },
+ { name = "networkx", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "pytest", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/5c/37/0211f82891a9f14efcfd2b2096f8d9e4351398ad637fdd1ee59cfc580b0e/bddl-1.0.1.tar.gz", hash = "sha256:1fa4e6e5050b93888ff6fd8455c39bfb29d3864ce06b4c37c0f781f513a2ae26", size = 164809, upload-time = "2022-03-08T01:48:23.564Z" }
@@ -1010,7 +1010,7 @@ name = "cuda-bindings"
version = "12.9.7"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "cuda-pathfinder" },
+ { name = "cuda-pathfinder", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/32/45/557d4ed1fa54f0c7db8aee083229f624990d69f7d00f55477eed5c7e169a/cuda_bindings-12.9.7-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0666d3c082ef8f4b2d670950589373550e9f3bf564d635dd883f24a0b40402ff", size = 7071026, upload-time = "2026-05-27T18:44:13.356Z" },
@@ -1043,37 +1043,37 @@ wheels = [
[package.optional-dependencies]
cublas = [
- { name = "nvidia-cublas-cu12" },
+ { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" },
]
cudart = [
- { name = "nvidia-cuda-runtime-cu12" },
+ { name = "nvidia-cuda-runtime-cu12", marker = "sys_platform == 'linux'" },
]
cufft = [
- { name = "nvidia-cufft-cu12" },
+ { name = "nvidia-cufft-cu12", marker = "sys_platform == 'linux'" },
]
cufile = [
- { name = "nvidia-cufile-cu12" },
+ { name = "nvidia-cufile-cu12", marker = "sys_platform == 'linux'" },
]
cupti = [
- { name = "nvidia-cuda-cupti-cu12" },
+ { name = "nvidia-cuda-cupti-cu12", marker = "sys_platform == 'linux'" },
]
curand = [
- { name = "nvidia-curand-cu12" },
+ { name = "nvidia-curand-cu12", marker = "sys_platform == 'linux'" },
]
cusolver = [
- { name = "nvidia-cusolver-cu12" },
+ { name = "nvidia-cusolver-cu12", marker = "sys_platform == 'linux'" },
]
cusparse = [
- { name = "nvidia-cusparse-cu12" },
+ { name = "nvidia-cusparse-cu12", marker = "sys_platform == 'linux'" },
]
nvjitlink = [
- { name = "nvidia-nvjitlink-cu12" },
+ { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
]
nvrtc = [
- { name = "nvidia-cuda-nvrtc-cu12" },
+ { name = "nvidia-cuda-nvrtc-cu12", marker = "sys_platform == 'linux'" },
]
nvtx = [
- { name = "nvidia-nvtx-cu12" },
+ { name = "nvidia-nvtx-cu12", marker = "sys_platform == 'linux'" },
]
[[package]]
@@ -1145,7 +1145,7 @@ name = "decord"
version = "0.6.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "numpy" },
+ { name = "numpy", marker = "(platform_machine != 'arm64' and sys_platform == 'darwin') or (platform_machine == 'AMD64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (sys_platform != 'darwin' and sys_platform != 'linux')" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/11/79/936af42edf90a7bd4e41a6cac89c913d4b47fa48a26b042d5129a9242ee3/decord-0.6.0-py3-none-manylinux2010_x86_64.whl", hash = "sha256:51997f20be8958e23b7c4061ba45d0efcd86bffd5fe81c695d0befee0d442976", size = 13602299, upload-time = "2021-06-14T21:30:55.486Z" },
@@ -1283,10 +1283,10 @@ resolution-markers = [
"python_full_version == '3.14.*' and sys_platform == 'win32'",
]
dependencies = [
- { name = "absl-py" },
- { name = "attrs" },
- { name = "numpy" },
- { name = "wrapt" },
+ { name = "absl-py", marker = "python_full_version >= '3.14'" },
+ { name = "attrs", marker = "python_full_version >= '3.14'" },
+ { name = "numpy", marker = "python_full_version >= '3.14'" },
+ { name = "wrapt", marker = "python_full_version >= '3.14'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/a6/83/ce29720ccf934c6cfa9b9c95ebbe96558386e66886626066632b5e44afed/dm_tree-0.1.9.tar.gz", hash = "sha256:a4c7db3d3935a5a2d5e4b383fc26c6b0cd6f78c6d4605d3e7b518800ecd5342b", size = 35623, upload-time = "2025-01-30T20:45:37.13Z" }
wheels = [
@@ -1324,10 +1324,10 @@ resolution-markers = [
"python_full_version < '3.13' and sys_platform == 'win32'",
]
dependencies = [
- { name = "absl-py" },
- { name = "attrs" },
- { name = "numpy" },
- { name = "wrapt" },
+ { name = "absl-py", marker = "python_full_version < '3.14'" },
+ { name = "attrs", marker = "python_full_version < '3.14'" },
+ { name = "numpy", marker = "python_full_version < '3.14'" },
+ { name = "wrapt", marker = "python_full_version < '3.14'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/5a/66/a3ec619d22b6baffa5ab853e8dc6ec9d0c837127948af59bb15b988d7312/dm_tree-0.1.10.tar.gz", hash = "sha256:22f37b599e01cc3402a17f79c257a802aebd8d326de05b54657650845956208a", size = 35748, upload-time = "2026-03-31T17:35:39.03Z" }
wheels = [
@@ -1912,7 +1912,7 @@ name = "h5py"
version = "3.16.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "numpy" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/db/33/acd0ce6863b6c0d7735007df01815403f5589a21ff8c2e1ee2587a38f548/h5py-3.16.0.tar.gz", hash = "sha256:a0dbaad796840ccaa67a4c144a0d0c8080073c34c76d5a6941d6818678ef2738", size = 446526, upload-time = "2026-03-06T13:49:08.07Z" }
wheels = [
@@ -1956,23 +1956,23 @@ name = "hf-libero"
version = "0.1.4"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "bddl" },
- { name = "cloudpickle" },
- { name = "easydict" },
- { name = "einops" },
- { name = "future" },
- { name = "gymnasium" },
- { name = "hf-egl-probe" },
- { name = "hydra-core" },
- { name = "matplotlib" },
- { name = "mujoco" },
- { name = "numpy" },
- { name = "opencv-python" },
- { name = "robomimic" },
- { name = "robosuite" },
- { name = "thop" },
- { name = "transformers" },
- { name = "wandb" },
+ { name = "bddl", marker = "sys_platform == 'linux'" },
+ { name = "cloudpickle", marker = "sys_platform == 'linux'" },
+ { name = "easydict", marker = "sys_platform == 'linux'" },
+ { name = "einops", marker = "sys_platform == 'linux'" },
+ { name = "future", marker = "sys_platform == 'linux'" },
+ { name = "gymnasium", marker = "sys_platform == 'linux'" },
+ { name = "hf-egl-probe", marker = "sys_platform == 'linux'" },
+ { name = "hydra-core", marker = "sys_platform == 'linux'" },
+ { name = "matplotlib", marker = "sys_platform == 'linux'" },
+ { name = "mujoco", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "opencv-python", marker = "sys_platform == 'linux'" },
+ { name = "robomimic", marker = "sys_platform == 'linux'" },
+ { name = "robosuite", marker = "sys_platform == 'linux'" },
+ { name = "thop", marker = "sys_platform == 'linux'" },
+ { name = "transformers", marker = "sys_platform == 'linux'" },
+ { name = "wandb", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/af/aa/4e9eb8715e0bff9cb6553db563a35d253393097d446f82bd53575e8b253d/hf_libero-0.1.4.tar.gz", hash = "sha256:c058d67ad5a2b589529c14d614282ef4cca3a7763dafa134f58a6c9039657e34", size = 2961319, upload-time = "2026-06-10T09:56:13.994Z" }
wheels = [
@@ -2123,9 +2123,9 @@ name = "hydra-core"
version = "1.3.4"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "antlr4-python3-runtime" },
- { name = "omegaconf" },
- { name = "packaging" },
+ { name = "antlr4-python3-runtime", marker = "sys_platform == 'linux'" },
+ { name = "omegaconf", marker = "sys_platform == 'linux'" },
+ { name = "packaging", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/10/dd/220f0e91743136725352497e98540772a01fc7c3ab96ff16c3c74424e984/hydra_core-1.3.4.tar.gz", hash = "sha256:ad0f7b05a0242255a8984d5a4ed2f6847f7b783ed727368a2c0155ec52d6c34c", size = 3263348, upload-time = "2026-07-04T16:25:38.891Z" }
wheels = [
@@ -2678,11 +2678,11 @@ name = "jupytext"
version = "1.19.5"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "markdown-it-py" },
- { name = "mdit-py-plugins" },
- { name = "nbformat" },
- { name = "packaging" },
- { name = "pyyaml" },
+ { name = "markdown-it-py", marker = "sys_platform == 'linux'" },
+ { name = "mdit-py-plugins", marker = "sys_platform == 'linux'" },
+ { name = "nbformat", marker = "sys_platform == 'linux'" },
+ { name = "packaging", marker = "sys_platform == 'linux'" },
+ { name = "pyyaml", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/a1/ca/473f8ebb101553fb2ea6ab1d34324d6677844c968947ac050c759d539f2c/jupytext-1.19.5.tar.gz", hash = "sha256:605026446d605aa54fd7f7fc69df6ae51c7a46053d4cebf05afdc64d66de3df0", size = 4600916, upload-time = "2026-07-21T22:00:29.198Z" }
wheels = [
@@ -2903,6 +2903,7 @@ all = [
{ name = "ruff" },
{ name = "scikit-image" },
{ name = "scipy" },
+ { name = "sentencepiece" },
{ name = "teleop" },
{ name = "timm" },
{ name = "torchcodec", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'AMD64' and sys_platform == 'linux') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'arm64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or sys_platform == 'win32'" },
@@ -3151,6 +3152,7 @@ phone = [
]
pi = [
{ name = "scipy" },
+ { name = "sentencepiece" },
{ name = "transformers" },
]
placo-dep = [
@@ -3211,6 +3213,9 @@ sarm = [
scipy-dep = [
{ name = "scipy" },
]
+sentencepiece-dep = [
+ { name = "sentencepiece" },
+]
smolvla = [
{ name = "accelerate" },
{ name = "num2words" },
@@ -3411,6 +3416,7 @@ requires-dist = [
{ name = "lerobot", extras = ["scipy-dep"], marker = "extra == 'phone'" },
{ name = "lerobot", extras = ["scipy-dep"], marker = "extra == 'pi'" },
{ name = "lerobot", extras = ["scipy-dep"], marker = "extra == 'wallx'" },
+ { name = "lerobot", extras = ["sentencepiece-dep"], marker = "extra == 'pi'" },
{ name = "lerobot", extras = ["smolvla"], marker = "extra == 'all'" },
{ name = "lerobot", extras = ["test"], marker = "extra == 'all'" },
{ name = "lerobot", extras = ["timm-dep"], marker = "extra == 'groot'" },
@@ -3486,6 +3492,7 @@ requires-dist = [
{ name = "scikit-image", marker = "extra == 'video-benchmark'", specifier = ">=0.23.2,<0.26.0" },
{ name = "scipy", marker = "extra == 'all'", specifier = ">=1.14.0,<2.0.0" },
{ name = "scipy", marker = "extra == 'scipy-dep'", specifier = ">=1.14.0,<2.0.0" },
+ { name = "sentencepiece", marker = "extra == 'sentencepiece-dep'", specifier = ">=0.2.0,<0.3.0" },
{ name = "setuptools", specifier = ">=71.0.0,<81.0.0" },
{ name = "teleop", marker = "extra == 'phone'", specifier = ">=0.1.0,<0.2.0" },
{ name = "termcolor", specifier = ">=2.4.0,<4.0.0" },
@@ -3502,7 +3509,7 @@ requires-dist = [
{ name = "transformers", marker = "extra == 'transformers-dep'", specifier = ">=5.4.0,<5.6.0" },
{ name = "wandb", marker = "extra == 'training'", specifier = ">=0.24.0,<0.28.0" },
]
-provides-extras = ["dataset", "training", "hardware", "viz", "core-scripts", "evaluation", "dataset-viz", "av-dep", "pygame-dep", "placo-dep", "transformers-dep", "grpcio-dep", "accelerate-dep", "can-dep", "peft-dep", "scipy-dep", "diffusers-dep", "qwen-vl-utils-dep", "matplotlib-dep", "pyserial-dep", "deepdiff-dep", "pynput-dep", "pyzmq-dep", "motorbridge-dep", "motorbridge-smart-servo-dep", "timm-dep", "feetech", "dynamixel", "damiao", "robstride", "openarms", "gamepad", "hopejr", "lekiwi", "unitree-g1", "reachy2", "rebot", "kinematics", "intelrealsense", "phone", "diffusion", "wallx", "pi", "molmoact2", "smolvla", "multi-task-dit", "groot", "sarm", "robometer", "topreward", "xvla", "eo1", "fastwam", "evo1", "hilserl", "vla-jepa", "lingbot-va", "async", "peft", "annotations", "dev", "notebook", "test", "video-benchmark", "aloha", "pusht", "libero", "metaworld", "all"]
+provides-extras = ["dataset", "training", "hardware", "viz", "core-scripts", "evaluation", "dataset-viz", "av-dep", "pygame-dep", "placo-dep", "transformers-dep", "sentencepiece-dep", "grpcio-dep", "accelerate-dep", "can-dep", "peft-dep", "scipy-dep", "diffusers-dep", "qwen-vl-utils-dep", "matplotlib-dep", "pyserial-dep", "deepdiff-dep", "pynput-dep", "pyzmq-dep", "motorbridge-dep", "motorbridge-smart-servo-dep", "timm-dep", "feetech", "dynamixel", "damiao", "robstride", "openarms", "gamepad", "hopejr", "lekiwi", "unitree-g1", "reachy2", "rebot", "kinematics", "intelrealsense", "phone", "diffusion", "wallx", "pi", "molmoact2", "smolvla", "multi-task-dit", "groot", "sarm", "robometer", "topreward", "xvla", "eo1", "fastwam", "evo1", "hilserl", "vla-jepa", "lingbot-va", "async", "peft", "annotations", "dev", "notebook", "test", "video-benchmark", "aloha", "pusht", "libero", "metaworld", "all"]
[[package]]
name = "librt"
@@ -3817,7 +3824,7 @@ name = "mdit-py-plugins"
version = "0.6.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "markdown-it-py" },
+ { name = "markdown-it-py", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/59/fc/f8d0863f8862f25602c0404d75568e89fb6b4109804645e5cdfb1be5cf56/mdit_py_plugins-0.6.1.tar.gz", hash = "sha256:a2bca0f039f39dbd35fb74ae1b5f998608c437463371f0ff7f49a19a17a114d0", size = 56114, upload-time = "2026-05-13T09:03:38.91Z" }
wheels = [
@@ -4296,8 +4303,8 @@ name = "numba"
version = "0.66.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "llvmlite" },
- { name = "numpy" },
+ { name = "llvmlite", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/ae/a0/570e3dc53e5602b49108f62a13e529f1eec8bfc7ef37d49c825924dcf546/numba-0.66.0.tar.gz", hash = "sha256:b900e63a0e26c05ea9a6d5a3a5a0a177cb64c5011887bf43edb8c3ed2c38d363", size = 2806181, upload-time = "2026-07-01T23:12:46.36Z" }
wheels = [
@@ -4390,7 +4397,7 @@ name = "nvidia-cudnn-cu12"
version = "9.19.0.56"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "nvidia-cublas-cu12" },
+ { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/09/b8/277c51962ee46fa3e5b203ac5f76107c650f781d6891e681e28e6f3e9fe6/nvidia_cudnn_cu12-9.19.0.56-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:08caaf27fe556aca82a3ee3b5aa49a77e7de0cfcb7ff4e5c29da426387a8267e", size = 656910700, upload-time = "2026-02-03T20:40:25.508Z" },
@@ -4402,7 +4409,7 @@ name = "nvidia-cufft-cu12"
version = "11.3.3.83"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "nvidia-nvjitlink-cu12" },
+ { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/60/bc/7771846d3a0272026c416fbb7e5f4c1f146d6d80704534d0b187dd6f4800/nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:848ef7224d6305cdb2a4df928759dca7b1201874787083b6e7550dd6765ce69a", size = 193109211, upload-time = "2025-03-07T01:44:56.873Z" },
@@ -4432,9 +4439,9 @@ name = "nvidia-cusolver-cu12"
version = "11.7.3.90"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "nvidia-cublas-cu12" },
- { name = "nvidia-cusparse-cu12" },
- { name = "nvidia-nvjitlink-cu12" },
+ { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-cusparse-cu12", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/c8/32/f7cd6ce8a7690544d084ea21c26e910a97e077c9b7f07bf5de623ee19981/nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:db9ed69dbef9715071232caa9b69c52ac7de3a95773c2db65bdba85916e4e5c0", size = 267229841, upload-time = "2025-03-07T01:46:54.356Z" },
@@ -4446,7 +4453,7 @@ name = "nvidia-cusparse-cu12"
version = "12.5.8.93"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "nvidia-nvjitlink-cu12" },
+ { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/bc/f7/cd777c4109681367721b00a106f491e0d0d15cfa1fd59672ce580ce42a97/nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:9b6c161cb130be1a07a27ea6923df8141f3c295852f4b260c65f18f3e0a091dc", size = 288117129, upload-time = "2025-03-07T01:47:40.407Z" },
@@ -4503,8 +4510,8 @@ name = "omegaconf"
version = "2.3.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "antlr4-python3-runtime" },
- { name = "pyyaml" },
+ { name = "antlr4-python3-runtime", marker = "sys_platform == 'linux'" },
+ { name = "pyyaml", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/ce/3d/e4b57b8d9008c6ebe0d5eff901f91d5700cf7bdb8c8863df817463a7fd5e/omegaconf-2.3.1.tar.gz", hash = "sha256:e5e7de64aeebeddaf8e6d3f7a783b32ac2a01c0fbd9c878012caecb891a1f42a", size = 3298472, upload-time = "2026-06-11T05:05:12.885Z" }
wheels = [
@@ -4743,7 +4750,7 @@ name = "pexpect"
version = "4.9.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "ptyprocess" },
+ { name = "ptyprocess", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/42/92/cc564bf6381ff43ce1f4d06852fc19a2f11d180f23dc32d9588bee2f149d/pexpect-4.9.0.tar.gz", hash = "sha256:ee7d41123f3c9911050ea2c2dac107568dc43b2d3b0c7557a33212c398ead30f", size = 166450, upload-time = "2023-11-25T09:07:26.339Z" }
wheels = [
@@ -5317,10 +5324,10 @@ name = "pyobjc-framework-applicationservices"
version = "12.2.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "pyobjc-core" },
- { name = "pyobjc-framework-cocoa" },
- { name = "pyobjc-framework-coretext" },
- { name = "pyobjc-framework-quartz" },
+ { name = "pyobjc-core", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-cocoa", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-coretext", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-quartz", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/5e/4d/0ebdd8144aba94b8fe9828ccee5616a4bf53d1f8bc51cff55f3cce86d695/pyobjc_framework_applicationservices-12.2.1.tar.gz", hash = "sha256:048ea663c9ac75c44a15dc7d5b8d78cbb4c97bf1c76e83835e8d5498e184001f", size = 109342, upload-time = "2026-06-19T16:19:46.149Z" }
wheels = [
@@ -5338,7 +5345,7 @@ name = "pyobjc-framework-cocoa"
version = "12.2.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "pyobjc-core" },
+ { name = "pyobjc-core", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/51/34/fbe38a204643aa4e1b91391cdce07a34da565a69171ebcad08de7438a556/pyobjc_framework_cocoa-12.2.1.tar.gz", hash = "sha256:b94b37fe5730e5ae1fb0052912cd174e6ec329b0bfba4a012ae5db1014b5864b", size = 3125751, upload-time = "2026-06-19T16:20:05.159Z" }
wheels = [
@@ -5356,9 +5363,9 @@ name = "pyobjc-framework-coretext"
version = "12.2.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "pyobjc-core" },
- { name = "pyobjc-framework-cocoa" },
- { name = "pyobjc-framework-quartz" },
+ { name = "pyobjc-core", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-cocoa", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-quartz", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/5a/9c/4c7f452059dc1d3845b8e627b9113c247a997b9b07518e848c2ab7ff3149/pyobjc_framework_coretext-12.2.1.tar.gz", hash = "sha256:af740e784d7c592c34025ec7165f4f6c1a69b5a2d9075f06e41e4f77c212aed2", size = 97349, upload-time = "2026-06-19T16:20:22.508Z" }
wheels = [
@@ -5376,8 +5383,8 @@ name = "pyobjc-framework-quartz"
version = "12.2.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "pyobjc-core" },
- { name = "pyobjc-framework-cocoa" },
+ { name = "pyobjc-core", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
+ { name = "pyobjc-framework-cocoa", marker = "sys_platform != 'emscripten' and sys_platform != 'linux' and sys_platform != 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3b/f6/2a8b84dbf1fe7c04dd96ea73d991678d4e09a909f51971ecc51629bb2ab4/pyobjc_framework_quartz-12.2.1.tar.gz", hash = "sha256:b3b8b6f71e66147f8ff9e6213864cc8527e3a0b1ee90835b93ce221f4802d9b0", size = 3215521, upload-time = "2026-06-19T16:21:30.199Z" }
wheels = [
@@ -5952,18 +5959,18 @@ name = "robomimic"
version = "0.2.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "egl-probe" },
- { name = "h5py" },
- { name = "imageio" },
- { name = "imageio-ffmpeg" },
- { name = "numpy" },
- { name = "psutil" },
- { name = "tensorboard" },
- { name = "tensorboardx" },
- { name = "termcolor" },
- { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } },
- { name = "torchvision", version = "0.26.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } },
- { name = "tqdm" },
+ { name = "egl-probe", marker = "sys_platform == 'linux'" },
+ { name = "h5py", marker = "sys_platform == 'linux'" },
+ { name = "imageio", marker = "sys_platform == 'linux'" },
+ { name = "imageio-ffmpeg", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "psutil", marker = "sys_platform == 'linux'" },
+ { name = "tensorboard", marker = "sys_platform == 'linux'" },
+ { name = "tensorboardx", marker = "sys_platform == 'linux'" },
+ { name = "termcolor", marker = "sys_platform == 'linux'" },
+ { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux'" },
+ { name = "torchvision", version = "0.26.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux'" },
+ { name = "tqdm", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3d/c3/44b1d1ea4bcb4bbed43d19e09505f4142714451ded74020d4f679cdc89fb/robomimic-0.2.0.tar.gz", hash = "sha256:ee3bb5cf9c3e1feead6b57b43c5db738fd0a8e0c015fdf6419808af8fffdc463", size = 192919, upload-time = "2021-12-17T19:00:33.279Z" }
@@ -5972,12 +5979,12 @@ name = "robosuite"
version = "1.4.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "mujoco" },
- { name = "numba" },
- { name = "numpy" },
- { name = "opencv-python" },
- { name = "pillow" },
- { name = "scipy" },
+ { name = "mujoco", marker = "sys_platform == 'linux'" },
+ { name = "numba", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "opencv-python", marker = "sys_platform == 'linux'" },
+ { name = "pillow", marker = "sys_platform == 'linux'" },
+ { name = "scipy", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/25/a1/9dd07a9a5e09c6aa032faf531da985808b34437cbf6c8f358fe8f7c47118/robosuite-1.4.0.tar.gz", hash = "sha256:a8a6233d7458dbd91bf00a86cab15aa1c178bd9d1b28d515db2cf3d152cb48e6", size = 192182147, upload-time = "2022-12-01T07:31:55.791Z" }
wheels = [
@@ -6218,6 +6225,49 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/1c/78/504fdd027da3b84ff1aecd9f6957e65f35134534ccc6da8628eb71e76d3f/send2trash-2.1.0-py3-none-any.whl", hash = "sha256:0da2f112e6d6bb22de6aa6daa7e144831a4febf2a87261451c4ad849fe9a873c", size = 17610, upload-time = "2026-01-14T06:27:35.218Z" },
]
+[[package]]
+name = "sentencepiece"
+version = "0.2.2"
+source = { registry = "https://pypi.org/simple" }
+sdist = { url = "https://files.pythonhosted.org/packages/cc/33/ea3cb3839607eb175da835244a798f797f478c5ddf0e8ecdf57ea85a4c70/sentencepiece-0.2.2.tar.gz", hash = "sha256:3d2b5e824b5622038dc7b490897efe05ebbbb9e7350fc142f3ecc8789ef9bdf6", size = 8218435, upload-time = "2026-07-12T08:39:34.701Z" }
+wheels = [
+ { url = "https://files.pythonhosted.org/packages/b8/13/7a562289c8d5b49ebdf3f9c1e8ab67cf14a8743b1d90c8f406bfdec36b72/sentencepiece-0.2.2-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:1edb10e520e4bddf74d85b0f5ae74cc2d60c2b448885080bfb618bc2b3a49f6b", size = 2188384, upload-time = "2026-07-12T08:38:28.486Z" },
+ { url = "https://files.pythonhosted.org/packages/85/d1/912f14fd5eae168aba726ffb6a9a2dc1c71fe7676c53da6f5c442b886d4a/sentencepiece-0.2.2-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:f7c06c751c19d923435a54bff4f7e66e728fad160e8da28254f133abc9725820", size = 1441553, upload-time = "2026-07-12T08:38:30.552Z" },
+ { url = "https://files.pythonhosted.org/packages/bd/44/caa9cab5f261a019e2808bc5046152775dc57352ba9cbae7525e9e7a1ed4/sentencepiece-0.2.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:38111ed1f79268f399c505028023d5eaaf0ab4e5eafceb709468b0d3323e7838", size = 1347176, upload-time = "2026-07-12T08:38:32.211Z" },
+ { url = "https://files.pythonhosted.org/packages/19/90/cd798935668cff71d309d8ff10385844ecf216b1fe454f1993ed8bf2cb91/sentencepiece-0.2.2-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:cbce24284f51f71d10a42b7b9c964dcb9048b28f1c8e5db40bcbcb6f428cba6a", size = 1325200, upload-time = "2026-07-12T08:38:33.689Z" },
+ { url = "https://files.pythonhosted.org/packages/b6/2d/37e3da037318a70066ded0d51bc2a7f35491ae6338dd993d5eb1503fc3b5/sentencepiece-0.2.2-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c8a168b040bc61681293f79a949b5d911c8e25086f4260285b8d97ab5f1195da", size = 1397736, upload-time = "2026-07-12T08:38:35.771Z" },
+ { url = "https://files.pythonhosted.org/packages/8d/11/753fca2e6b109be3ab7867abf357dfe48677fe726ae5a5363d0b54ca9450/sentencepiece-0.2.2-cp312-cp312-win_amd64.whl", hash = "sha256:7c6e7bf684dc12145bfa685d3060beaea55139134ba848289bee514ed42e7383", size = 1248030, upload-time = "2026-07-12T08:38:37.604Z" },
+ { url = "https://files.pythonhosted.org/packages/e2/0a/70efbe861ca182d7d4b6e1a20f58e043400848fa9f2915229f082e221648/sentencepiece-0.2.2-cp312-cp312-win_arm64.whl", hash = "sha256:76ff5814db72e7462dece042d7593cdf102b8ec82c2b1cc201a2add34ee3050d", size = 1187325, upload-time = "2026-07-12T08:38:39.348Z" },
+ { url = "https://files.pythonhosted.org/packages/b9/a3/b3b05095c174d6e80d37d5ddc2f57c2c56237333e7bbd6079cf3243c2a8a/sentencepiece-0.2.2-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:77c3ce990b23441e5ecfa5bce181fd6f408b564aeb6d7e1d1e7de9c5612501c8", size = 2188346, upload-time = "2026-07-12T08:38:41.089Z" },
+ { url = "https://files.pythonhosted.org/packages/ca/f3/72ebc4acb10a06bcf7503fbc6091c8f5db68300f6aac4356c09e6c76e0e1/sentencepiece-0.2.2-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:fd523c4992041faa5c2b3cde62253d11a96c30d73a34afe48a486e8e2254cd1c", size = 1441434, upload-time = "2026-07-12T08:38:42.56Z" },
+ { url = "https://files.pythonhosted.org/packages/34/db/f9ea1a6844b4fa5dfe2312095cd866a1f724cd0905054ab9d5991778ba50/sentencepiece-0.2.2-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:201a8e0f55501a76e08dbf2c54bc45f4642b379271e89c667d517bfbc2191f2a", size = 1347267, upload-time = "2026-07-12T08:38:44.389Z" },
+ { url = "https://files.pythonhosted.org/packages/32/4f/31c1073314ad94466bca37d29581761d70110237ee3d46b0efece59a8c1e/sentencepiece-0.2.2-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8eed98514bffe5ecac37f493f91869c351fbb05629328bfdbc08502c6c094dc0", size = 1324980, upload-time = "2026-07-12T08:38:46.304Z" },
+ { url = "https://files.pythonhosted.org/packages/59/b4/a0356fa04d6a14337a6e0e443556785a0422c53ec58baae6b9568120eb0f/sentencepiece-0.2.2-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:64b656f025355cf8c51abe9fbe3848540756c6d7ca5e6791b1afa664bc24c7cb", size = 1397593, upload-time = "2026-07-12T08:38:48.302Z" },
+ { url = "https://files.pythonhosted.org/packages/09/fa/d2d6369257fd2f0de616b1c7110b73fab409ef61b14f1b9e0010ed325914/sentencepiece-0.2.2-cp313-cp313-win_amd64.whl", hash = "sha256:74f0ee601047c0c12a783088b51be4e6214a62ecd9e02278c477433cd16e0ed9", size = 1247987, upload-time = "2026-07-12T08:38:50.15Z" },
+ { url = "https://files.pythonhosted.org/packages/17/ee/2bb594da6fd95e32f29057f1aa7fa996701b8980090923c2d8711fdc0a24/sentencepiece-0.2.2-cp313-cp313-win_arm64.whl", hash = "sha256:b23fe17779834d3c27aaf2edac9486d04cca1a7deb8f5facda35150ac6263a91", size = 1187250, upload-time = "2026-07-12T08:38:52.246Z" },
+ { url = "https://files.pythonhosted.org/packages/58/9c/dfc82846460e7a712310f5613f23d8b553cabb4e2e648663c11d8382af56/sentencepiece-0.2.2-cp313-cp313t-macosx_10_13_universal2.whl", hash = "sha256:72b7825b331b1b7e7c45be2e674b3e3c65af608fa376bad2d851b20aaf0cdc78", size = 2223080, upload-time = "2026-07-12T08:38:54.391Z" },
+ { url = "https://files.pythonhosted.org/packages/8d/4e/3ff12cebe6d31662d9ceeabfb282de20bd0d6098fa282b4a3b8305abc7e8/sentencepiece-0.2.2-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:d795c4ac689a57f9d4ba2288126ec7901d389ad5827d2f8b8533c883974fe563", size = 1458511, upload-time = "2026-07-12T08:38:56.811Z" },
+ { url = "https://files.pythonhosted.org/packages/59/5a/16d51d05360be4cee3ebfe4837c184054c4eed16cabaeb3b039524e9a000/sentencepiece-0.2.2-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:3ab3f1ae98970b5590e2209341522718900ba19bcc2c207ffaa6bd417ad960c5", size = 1361138, upload-time = "2026-07-12T08:38:58.808Z" },
+ { url = "https://files.pythonhosted.org/packages/0f/af/c30ee2a9f99d51db9844acaa8fa0b611a97c2fa7116646fa43db3300b187/sentencepiece-0.2.2-cp313-cp313t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3ec27c152a1f1b24bc9168b55a5880f3c16e2334e697da6f55a1046a22405a3d", size = 1328625, upload-time = "2026-07-12T08:39:00.849Z" },
+ { url = "https://files.pythonhosted.org/packages/3e/1a/4c6b39d03f5ba8439509adbd5a23c9538088a3cb679e7a47b911e8442bc6/sentencepiece-0.2.2-cp313-cp313t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:59d6588712101ccfcae9b03692be3aaae1514c2078666d7b05f15ba3a702e41b", size = 1398595, upload-time = "2026-07-12T08:39:02.86Z" },
+ { url = "https://files.pythonhosted.org/packages/0f/bc/9eedddcec1fd57bc70200fa3ebf792d18fa63527a5369581cd416c81f97f/sentencepiece-0.2.2-cp313-cp313t-win_amd64.whl", hash = "sha256:89625fb43765cccaa1443b9adb61f283e5fe4cb1536728205d06bada730caa53", size = 1259346, upload-time = "2026-07-12T08:39:04.559Z" },
+ { url = "https://files.pythonhosted.org/packages/41/15/7e74c8533848866ff560b29f7d8719921b76c4ec7149592d6d28e0deee75/sentencepiece-0.2.2-cp313-cp313t-win_arm64.whl", hash = "sha256:4f0603267cd15b92b68c2c0e852a441507614b70dc7773659baa6b8c214a91fd", size = 1196596, upload-time = "2026-07-12T08:39:06.454Z" },
+ { url = "https://files.pythonhosted.org/packages/0b/7e/f5df63edb6bcb46c1343cfa5d9192d73a4eb61af2e800d9402efff387523/sentencepiece-0.2.2-cp314-cp314-macosx_10_15_universal2.whl", hash = "sha256:c62bd361cec1f5b556eb8210264ecfff37486cd990c3386cc00310f26c54090a", size = 2190240, upload-time = "2026-07-12T08:39:08.178Z" },
+ { url = "https://files.pythonhosted.org/packages/52/0a/095d183b453b2a2e20b016829029c58eca90adc1c9911113e5d26fff45ed/sentencepiece-0.2.2-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:46ba07b543add034de0ff47ac5f907e9a06682f91d85121a972764628933be6b", size = 1442220, upload-time = "2026-07-12T08:39:09.91Z" },
+ { url = "https://files.pythonhosted.org/packages/d1/18/823954c9c90e74eba09fb96752dc37a5555df00d69866cb9406d1725dc7e/sentencepiece-0.2.2-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:79bac5a251f23a7341e28fda9ce0d5319edf45328239ce037c0682936f137906", size = 1348056, upload-time = "2026-07-12T08:39:11.744Z" },
+ { url = "https://files.pythonhosted.org/packages/10/ca/1b6c251321901cbf8a2d2e48b8b70eb82a449011b766af52a228d0a90b6b/sentencepiece-0.2.2-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1402d8ee36f0d851cea8eee4dbb85fea14643b7503cf4d00d102eec0fe3ca719", size = 1325463, upload-time = "2026-07-12T08:39:13.413Z" },
+ { url = "https://files.pythonhosted.org/packages/24/b3/718847349da7b25c8220ed86d85b89080af94740b2d87a59198104ae5c51/sentencepiece-0.2.2-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8d44b20234905ff022b7d535f79d1f823ad7670c9851cc4f03cdc34787cdb3ab", size = 1398138, upload-time = "2026-07-12T08:39:15.564Z" },
+ { url = "https://files.pythonhosted.org/packages/33/fe/4906f12c458274edd96387e4baaad7c6f064a2b7c11a1cc2401c8a7bd483/sentencepiece-0.2.2-cp314-cp314-win_amd64.whl", hash = "sha256:63250cfab8b80a1ef82a614eb2b3cadfec2c405f870cedc139d08e2f063eb708", size = 1356144, upload-time = "2026-07-12T08:39:17.313Z" },
+ { url = "https://files.pythonhosted.org/packages/d3/eb/22f89b6542aba400b0007cf0b1697cc3f99be8fb682fdb4c05eec450e33f/sentencepiece-0.2.2-cp314-cp314-win_arm64.whl", hash = "sha256:65d84ec36888de4a848eee5f910e67fbc79b064685ef1e10a502e14520ead9c9", size = 1294351, upload-time = "2026-07-12T08:39:18.967Z" },
+ { url = "https://files.pythonhosted.org/packages/84/c4/7afe8c2315b76e46818851a057e50a378a0382aa00b970a1fa444181b6f6/sentencepiece-0.2.2-cp314-cp314t-macosx_10_15_universal2.whl", hash = "sha256:d254c98ca6387655400b3959c33c83efd807f5edeb608e3aca45800ceaa77151", size = 2223281, upload-time = "2026-07-12T08:39:20.978Z" },
+ { url = "https://files.pythonhosted.org/packages/98/42/fb678e472c554ef086be6375d20060ca610a2c4218854d4c091001fc6f91/sentencepiece-0.2.2-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:3fd9ce2ab4460c713cfdeb4aca693ca6732a11538e05fb332d5af42e3d7fde25", size = 1458779, upload-time = "2026-07-12T08:39:22.812Z" },
+ { url = "https://files.pythonhosted.org/packages/78/52/ffe402b13bce1889228a98dc6cd86ae8afac1112362236be3468be784441/sentencepiece-0.2.2-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:7fc14c1585139fa6b68775e616a6b90cf622ebf219f9558c0aeaf5d253ee6c9b", size = 1361736, upload-time = "2026-07-12T08:39:24.602Z" },
+ { url = "https://files.pythonhosted.org/packages/78/4a/2288f60e7283583ec0a0f16e72f9c8e68557d7e7a4b585d2cda4f9f47e64/sentencepiece-0.2.2-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:df88b0c34f2fa909d322f7b06b1398e1e81af4b2f42a7b8e3556f928b25d1811", size = 1328155, upload-time = "2026-07-12T08:39:26.422Z" },
+ { url = "https://files.pythonhosted.org/packages/26/31/5dd6882ebe899f741a5cfe40ff56c6efc06bc26ee287abdb723b671f409c/sentencepiece-0.2.2-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3f5851441ab1ef8634963a5100b733a8bbeefe623e0c5c005b1f1f3880e574cf", size = 1398307, upload-time = "2026-07-12T08:39:28.637Z" },
+ { url = "https://files.pythonhosted.org/packages/da/05/7d7780fa63f4b8c1821953b916e25f89ae8f14d4da6ba91e10f6d06dc2b4/sentencepiece-0.2.2-cp314-cp314t-win_amd64.whl", hash = "sha256:046b15ea22d8042e2e173561d464ec3b64a9c2081324df70ebce7bf7ebb3e497", size = 1367133, upload-time = "2026-07-12T08:39:30.546Z" },
+ { url = "https://files.pythonhosted.org/packages/49/a1/70007fef3f818c688de4a730f98024a671599ab67f20270f8efb03d69dcc/sentencepiece-0.2.2-cp314-cp314t-win_arm64.whl", hash = "sha256:fa9f5ef0e2a82233dd0b8b32ea3f5710e0c44afbc07ed3620219f32601e56090", size = 1302760, upload-time = "2026-07-12T08:39:32.457Z" },
+]
+
[[package]]
name = "sentry-sdk"
version = "2.66.1"
@@ -6398,16 +6448,16 @@ name = "tensorboard"
version = "2.20.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "absl-py" },
- { name = "grpcio" },
- { name = "markdown" },
- { name = "numpy" },
- { name = "packaging" },
- { name = "pillow" },
- { name = "protobuf" },
- { name = "setuptools" },
- { name = "tensorboard-data-server" },
- { name = "werkzeug" },
+ { name = "absl-py", marker = "sys_platform == 'linux'" },
+ { name = "grpcio", marker = "sys_platform == 'linux'" },
+ { name = "markdown", marker = "sys_platform == 'linux'" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "packaging", marker = "sys_platform == 'linux'" },
+ { name = "pillow", marker = "sys_platform == 'linux'" },
+ { name = "protobuf", marker = "sys_platform == 'linux'" },
+ { name = "setuptools", marker = "sys_platform == 'linux'" },
+ { name = "tensorboard-data-server", marker = "sys_platform == 'linux'" },
+ { name = "werkzeug", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/9c/d9/a5db55f88f258ac669a92858b70a714bbbd5acd993820b41ec4a96a4d77f/tensorboard-2.20.0-py3-none-any.whl", hash = "sha256:9dc9f978cb84c0723acf9a345d96c184f0293d18f166bb8d59ee098e6cfaaba6", size = 5525680, upload-time = "2025-07-17T19:20:49.638Z" },
@@ -6427,9 +6477,9 @@ name = "tensorboardx"
version = "2.6.5"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "numpy" },
- { name = "packaging" },
- { name = "protobuf" },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "packaging", marker = "sys_platform == 'linux'" },
+ { name = "protobuf", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/48/a9/fc520ea91ab1f3ba51cbf3fe24f2b6364ed3b49046969e0868d46d6da372/tensorboardx-2.6.5.tar.gz", hash = "sha256:ca176db3997ee8c07d2eb77381225956a3fd1c10c91beafab1f17069adc47017", size = 4770195, upload-time = "2026-04-03T15:40:23.803Z" }
wheels = [
@@ -6464,7 +6514,7 @@ name = "thop"
version = "0.1.1.post2209072238"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } },
+ { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/bb/0f/72beeab4ff5221dc47127c80f8834b4bcd0cb36f6ba91c0b1d04a1233403/thop-0.1.1.post2209072238-py3-none-any.whl", hash = "sha256:01473c225231927d2ad718351f78ebf7cffe6af3bed464c4f1ba1ef0f7cdda27", size = 15443, upload-time = "2022-09-07T14:38:37.211Z" },
@@ -6570,13 +6620,13 @@ resolution-markers = [
"python_full_version < '3.13' and sys_platform == 'win32'",
]
dependencies = [
- { name = "filelock" },
- { name = "fsspec" },
- { name = "jinja2" },
- { name = "networkx" },
- { name = "setuptools" },
- { name = "sympy" },
- { name = "typing-extensions" },
+ { name = "filelock", marker = "sys_platform != 'linux'" },
+ { name = "fsspec", marker = "sys_platform != 'linux'" },
+ { name = "jinja2", marker = "sys_platform != 'linux'" },
+ { name = "networkx", marker = "sys_platform != 'linux'" },
+ { name = "setuptools", marker = "sys_platform != 'linux'" },
+ { name = "sympy", marker = "sys_platform != 'linux'" },
+ { name = "typing-extensions", marker = "sys_platform != 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/6f/8b/69e3008d78e5cee2b30183340cc425081b78afc5eff3d080daab0adda9aa/torch-2.11.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:4b5866312ee6e52ea625cd211dcb97d6a2cdc1131a5f15cc0d87eec948f6dd34", size = 80606338, upload-time = "2026-03-23T18:11:34.781Z" },
@@ -6610,20 +6660,20 @@ resolution-markers = [
"python_full_version < '3.13' and platform_machine != 'AMD64' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'x86_64' and sys_platform == 'linux'",
]
dependencies = [
- { name = "cuda-bindings" },
- { name = "cuda-toolkit", extra = ["cublas", "cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"] },
- { name = "filelock" },
- { name = "fsspec" },
- { name = "jinja2" },
- { name = "networkx" },
- { name = "nvidia-cudnn-cu12" },
- { name = "nvidia-cusparselt-cu12" },
- { name = "nvidia-nccl-cu12" },
- { name = "nvidia-nvshmem-cu12" },
- { name = "setuptools" },
- { name = "sympy" },
- { name = "triton" },
- { name = "typing-extensions" },
+ { name = "cuda-bindings", marker = "sys_platform == 'linux'" },
+ { name = "cuda-toolkit", extra = ["cublas", "cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "sys_platform == 'linux'" },
+ { name = "filelock", marker = "sys_platform == 'linux'" },
+ { name = "fsspec", marker = "sys_platform == 'linux'" },
+ { name = "jinja2", marker = "sys_platform == 'linux'" },
+ { name = "networkx", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-cudnn-cu12", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-cusparselt-cu12", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-nccl-cu12", marker = "sys_platform == 'linux'" },
+ { name = "nvidia-nvshmem-cu12", marker = "sys_platform == 'linux'" },
+ { name = "setuptools", marker = "sys_platform == 'linux'" },
+ { name = "sympy", marker = "sys_platform == 'linux'" },
+ { name = "triton", marker = "sys_platform == 'linux'" },
+ { name = "typing-extensions", marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://download-r2.pytorch.org/whl/cu128/torch-2.11.0%2Bcu128-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:9c8f38efee365cb9d334de8a83ce52fc7e5fc9e5a7b0853285efa1b69e00b0f2", upload-time = "2026-04-27T17:41:30Z" },
@@ -6694,9 +6744,9 @@ resolution-markers = [
"python_full_version < '3.13' and sys_platform == 'win32'",
]
dependencies = [
- { name = "numpy" },
- { name = "pillow" },
- { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" } },
+ { name = "numpy", marker = "sys_platform != 'linux'" },
+ { name = "pillow", marker = "sys_platform != 'linux'" },
+ { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform != 'linux'" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/ae/e7/56b47cc3b132aea90ccce22bcb8975dec688b002150012acc842846039d0/torchvision-0.26.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:c409e1c3fdebec7a3834465086dbda8bf7680eff79abf7fd2f10c6b59520a7a4", size = 1863502, upload-time = "2026-03-23T18:12:57.326Z" },
@@ -6730,9 +6780,9 @@ resolution-markers = [
"python_full_version < '3.13' and platform_machine != 'AMD64' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'x86_64' and sys_platform == 'linux'",
]
dependencies = [
- { name = "numpy" },
- { name = "pillow" },
- { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } },
+ { name = "numpy", marker = "sys_platform == 'linux'" },
+ { name = "pillow", marker = "sys_platform == 'linux'" },
+ { name = "torch", version = "2.11.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux'" },
]
wheels = [
{ url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.26.0%2Bcu128-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:63e35234aed13b6edda37056f417b5c281249669db631e706811917af36b21d7", upload-time = "2026-04-09T23:21:35Z" },
@@ -7222,7 +7272,7 @@ name = "werkzeug"
version = "3.1.8"
source = { registry = "https://pypi.org/simple" }
dependencies = [
- { name = "markupsafe" },
+ { name = "markupsafe", marker = "sys_platform == 'linux'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/dd/b2/381be8cfdee792dd117872481b6e378f85c957dd7c5bca38897b08f765fd/werkzeug-3.1.8.tar.gz", hash = "sha256:9bad61a4268dac112f1c5cd4630a56ede601b6ed420300677a869083d70a4c44", size = 875852, upload-time = "2026-04-02T18:49:14.268Z" }
wheels = [