mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 04:59:44 +00:00
Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| cf927acb6e | |||
| 051b13573e | |||
| 7de2e4c1ef | |||
| 8db50611c2 | |||
| 92f96f33b3 | |||
| d4b3ca569c |
@@ -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), [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) |
|
||||
| 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) |
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -63,8 +63,6 @@
|
||||
title: π₀-FAST (Pi0Fast)
|
||||
- local: pi05
|
||||
title: π₀.₅ (Pi05)
|
||||
- local: pi052
|
||||
title: π₀.₅ with language supervision (Pi052)
|
||||
- local: molmoact2
|
||||
title: MolmoAct2
|
||||
- local: vla_jepa
|
||||
|
||||
@@ -189,162 +189,6 @@ def make_my_policy_pre_post_processors(
|
||||
|
||||
---
|
||||
|
||||
## Adding high- and low-level language control
|
||||
|
||||
The policy API above is sufficient for training and standard evaluation. To use a language-conditioned policy with interactive `lerobot-rollout`, also register a runtime adapter. The adapter keeps policy-specific prompting and tokenization out of the generic control loop.
|
||||
|
||||
The runtime supports two policy shapes:
|
||||
|
||||
| Policy shape | Behavior | Adapter |
|
||||
| ---------------- | ----------------------------------------------------------------------- | ---------------------------------------------- |
|
||||
| Low-level / flat | The operator's task or subtask directly conditions action prediction. | Reuse `DirectTaskPolicyAdapter`. |
|
||||
| High + low level | The policy generates subtasks or memory, then conditions actions on it. | Subclass `BaseLanguageAdapter`, as PI052 does. |
|
||||
|
||||
During a rollout, `RuntimeState` stores the high-level task and the active language context:
|
||||
|
||||
```text
|
||||
task ──> adapter.generate_text("subtask", ...) ──> state.language_context["subtask"]
|
||||
│
|
||||
observation ──> processors ──> adapter.select_action() ─┴─> action chunk ──> robot
|
||||
```
|
||||
|
||||
The generic runtime handles generation frequency, pause/resume, prompt replacement, action queues, and dispatch. The adapter only translates between that runtime contract and your policy.
|
||||
|
||||
### Low-level policies
|
||||
|
||||
If your policy already consumes the live task through its normal preprocessor and implements `predict_action_chunk`, register the shared direct adapter. PI0.5 and MolmoAct2 use this path:
|
||||
|
||||
```python
|
||||
# src/lerobot/runtime/registry.py
|
||||
_ADAPTERS = {
|
||||
# ...
|
||||
"my_policy": "lerobot.runtime.adapter:DirectTaskPolicyAdapter",
|
||||
}
|
||||
```
|
||||
|
||||
Run it with direct-subtask mode so the operator supplies the instruction used by the action policy:
|
||||
|
||||
```bash
|
||||
lerobot-rollout \
|
||||
--language \
|
||||
--policy.path=user/my_policy_checkpoint \
|
||||
--robot.type=so101_follower \
|
||||
--robot.port=/dev/ttyACM0 \
|
||||
--direct_subtask
|
||||
```
|
||||
|
||||
The rollout context builds the observation batch with the current instruction before `DirectTaskPolicyAdapter` calls `policy.predict_action_chunk(observation)`. No text-generation method is required.
|
||||
|
||||
### Hierarchical policies
|
||||
|
||||
For a policy that generates language and actions, subclass [`BaseLanguageAdapter`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/runtime/adapter.py) and implement two methods:
|
||||
|
||||
- `generate_text(kind, observation, state, user_text=None) -> str` generates a `subtask`, `memory`, or interjection response.
|
||||
- `select_action(observation, state)` builds the low-level prompt from the active context and returns an action chunk.
|
||||
|
||||
This abbreviated adapter follows [`PI052PolicyAdapter`](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/pi052/inference/pi052_adapter.py):
|
||||
|
||||
```python
|
||||
# inference/my_policy_adapter.py
|
||||
from typing import Any
|
||||
|
||||
from lerobot.runtime import RuntimeState
|
||||
from lerobot.runtime.adapter import BaseLanguageAdapter
|
||||
from lerobot.utils.constants import (
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
OBS_LANGUAGE_TOKENS,
|
||||
)
|
||||
|
||||
|
||||
class MyPolicyAdapter(BaseLanguageAdapter):
|
||||
def select_action(self, observation: dict[str, Any], state: RuntimeState):
|
||||
instruction = state.language_context.get("subtask") or state.task or ""
|
||||
tokens, attention_mask = tokenize_instruction(instruction)
|
||||
|
||||
batch = dict(observation)
|
||||
batch[OBS_LANGUAGE_TOKENS] = tokens
|
||||
batch[OBS_LANGUAGE_ATTENTION_MASK] = attention_mask
|
||||
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)
|
||||
batch, tokenizer = tokenize_messages(messages, observation)
|
||||
return self.policy.select_message(
|
||||
batch,
|
||||
tokenizer=tokenizer,
|
||||
min_new_tokens=self.gen.min_new_tokens,
|
||||
temperature=self.gen.temperature,
|
||||
top_p=self.gen.top_p,
|
||||
)
|
||||
|
||||
def build_messages(
|
||||
self, kind: str, state: RuntimeState, user_text: str | None
|
||||
) -> list[dict[str, str]]:
|
||||
if kind == "subtask":
|
||||
return [{"role": "user", "content": state.task or ""}]
|
||||
if kind == "memory":
|
||||
return [
|
||||
{"role": "user", "content": state.task or ""},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"Completed subtask: {state.extra.get('prior_subtask', '')}",
|
||||
},
|
||||
]
|
||||
if kind == "interjection":
|
||||
return [
|
||||
{"role": "user", "content": state.task or ""},
|
||||
{"role": "user", "content": user_text or ""},
|
||||
]
|
||||
raise ValueError(f"Unsupported text kind: {kind}")
|
||||
```
|
||||
|
||||
`tokenize_instruction` and `tokenize_messages` are policy-specific helpers. They must reproduce the prompt format used during training; PI052, for example, adds the discretized robot state to its low-level subtask prompt and uses the same PaliGemma formatting for `select_message`.
|
||||
|
||||
`BaseLanguageAdapter` provides the default hierarchy: regenerate a subtask at action-chunk boundaries, update memory when the subtask changes, and handle user interjections. Override `_regenerate_context` only if your policy uses a different hierarchy.
|
||||
|
||||
Register the adapter with a lazy import so importing LeRobot does not load the model or its optional dependencies:
|
||||
|
||||
```python
|
||||
# src/lerobot/runtime/registry.py
|
||||
_ADAPTERS = {
|
||||
# ...
|
||||
"my_policy": "lerobot.policies.my_policy.inference.my_policy_adapter:MyPolicyAdapter",
|
||||
}
|
||||
```
|
||||
|
||||
The key must match the policy's registered type. Once registered, the same checkpoint works through the shared entry point:
|
||||
|
||||
```bash
|
||||
lerobot-rollout \
|
||||
--language \
|
||||
--policy.path=user/my_hierarchical_checkpoint \
|
||||
--robot.type=so101_follower \
|
||||
--robot.port=/dev/ttyACM0 \
|
||||
--task="put the cup in the sink"
|
||||
```
|
||||
|
||||
For RoboCasa-compatible policies, replace the robot arguments with `--sim --sim.task=<task>`. Without `--direct_subtask`, the adapter generates the low-level subtask; with it, the operator bypasses high-level generation and supplies each subtask.
|
||||
|
||||
### Keep training and deployment aligned
|
||||
|
||||
The adapter is intentionally small, but its prompts are part of the model contract:
|
||||
|
||||
- Use the same tokenizer, role formatting, special tokens, image ordering, and state encoding as training.
|
||||
- Condition `select_action` on `state.language_context["subtask"]`, falling back to `state.task` for direct or not-yet-generated prompts.
|
||||
- Return a full action chunk from `select_action`; the runtime handles control-rate dispatch.
|
||||
- Keep optional model dependencies inside lazy imports.
|
||||
- Test adapter selection, generated-message routing, action-batch construction, and direct-subtask behavior with a lightweight fake policy.
|
||||
|
||||
PI052 is the complete in-tree reference: its [processor](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/pi052/processor_pi052.py) renders the training recipe, its [policy](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/pi052/modeling_pi052.py) exposes text and action generation, and its [adapter](https://github.com/huggingface/lerobot/blob/main/src/lerobot/policies/pi052/inference/pi052_adapter.py) reconstructs those same prompts at deployment.
|
||||
|
||||
---
|
||||
|
||||
## Path A: Out-of-tree plugin
|
||||
|
||||
The fastest way to ship a policy: package it as a standalone Python distribution and install it alongside LeRobot. No PR required, you own the release cycle, and you can publish to PyPI under your own namespace.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Policy Deployment (lerobot-rollout)
|
||||
|
||||
`lerobot-rollout` is the single CLI for deploying trained policies on real robots or in an interactive simulator. It supports multiple execution strategies and inference backends, from quick evaluation to continuous recording, language-driven control, and human-in-the-loop data collection.
|
||||
`lerobot-rollout` is the single CLI for deploying trained policies on real robots. It supports multiple execution strategies and inference backends, from quick evaluation to continuous recording and human-in-the-loop data collection.
|
||||
|
||||
## Quick Start
|
||||
|
||||
@@ -197,52 +197,6 @@ Teleop is optional — if omitted the robot holds its position during the reset
|
||||
|
||||
---
|
||||
|
||||
## Interactive language control
|
||||
|
||||
Language-conditioned policies can expose a high-level text head in addition to
|
||||
their action head. Add `--language` to open-prompt one of these policies on a
|
||||
real robot. Language-only flags such as `--direct_subtask` select this mode
|
||||
automatically.
|
||||
|
||||
MolmoAct2 has no high-level planner, so use direct-subtask mode and type each
|
||||
next low-level instruction yourself:
|
||||
|
||||
```bash
|
||||
lerobot-rollout \
|
||||
--policy.path=lerobot/MolmoAct2-SO100_101-LeRobot \
|
||||
--policy.device=cuda \
|
||||
--robot.type=so101_follower \
|
||||
--robot.port=/dev/ttyACM1 \
|
||||
--robot.cameras='{"cam0":{"type":"opencv","index_or_path":"/dev/video0","width":640,"height":480,"fps":30,"fourcc":"MJPG","backend":200},"cam1":{"type":"opencv","index_or_path":"/dev/video2","width":640,"height":480,"fps":30,"fourcc":"MJPG","backend":200}}' \
|
||||
--direct_subtask \
|
||||
--robot.max_relative_target='{"shoulder_pan":5,"shoulder_lift":5,"elbow_flex":5,"wrist_flex":5,"wrist_roll":5,"gripper":5}'
|
||||
```
|
||||
|
||||
The robot starts paused. Type a subtask, then use `/resume` and `/pause` to
|
||||
control action dispatch. Check the workspace and motion limits before resuming.
|
||||
Without `--direct_subtask`, a policy such as PI052 generates its active subtask
|
||||
from the high-level `--task` itself.
|
||||
|
||||
RoboCasa uses the same runtime and processor path. `--sim` selects it
|
||||
automatically, so no robot configuration is needed:
|
||||
|
||||
```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 simulator view. Add
|
||||
`--sim.direct_subtask` to bypass the language planner and make each typed prompt
|
||||
the action policy's current subtask.
|
||||
|
||||
---
|
||||
|
||||
## Inference Backends
|
||||
|
||||
Select a backend with `--inference.type=<name>`. All strategies work with both backends.
|
||||
|
||||
@@ -141,11 +141,6 @@ sample["target_message_indices"]
|
||||
|
||||
The renderer does not apply a tokenizer chat template. Policy processors decide how to serialize the messages for their backbone, which keeps the same dataset usable across SmolVLA, Pi0.5, and any future VLM that expects OpenAI-style chat messages.
|
||||
|
||||
## Blends
|
||||
|
||||
Blend recipes select one weighted sub-recipe deterministically from the sample index.
|
||||
`recipes/subtask_mem.yaml` trains the compact core blend — high-level subtask prediction, low-level execution, and memory. `recipes/subtask_mem_vqa_speech.yaml` is the fuller variant that also adds VQA and spoken interjection responses.
|
||||
|
||||
## Graceful absence
|
||||
|
||||
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
||||
|
||||
@@ -113,19 +113,6 @@ accelerate launch --num_processes=2 $(which lerobot-train) \
|
||||
--policy.type=act
|
||||
```
|
||||
|
||||
When the desired global batch is larger than the per-GPU batch that fits in memory, use gradient
|
||||
accumulation. `steps` continues to count optimizer updates:
|
||||
|
||||
```bash
|
||||
# 8 samples/GPU × 4 GPUs × 2 microbatches = effective global batch 64.
|
||||
accelerate launch --num_processes=4 $(which lerobot-train) \
|
||||
--batch_size=8 \
|
||||
--gradient_accumulation_steps=2 \
|
||||
--steps=80000 \
|
||||
--dataset.repo_id=lerobot/pusht \
|
||||
--policy.type=act
|
||||
```
|
||||
|
||||
## Training Large Models with FSDP
|
||||
|
||||
DDP replicates the full model on every GPU, so a model that doesn't fit on one GPU won't fit under
|
||||
|
||||
@@ -1,255 +0,0 @@
|
||||
# π₀.₅ 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_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 |
|
||||
|
||||
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.
|
||||
|
||||
<Tip>
|
||||
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.
|
||||
</Tip>
|
||||
|
||||
## 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 |
|
||||
|
||||
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.
|
||||
+9
-15
@@ -109,21 +109,15 @@ 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.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.
|
||||
| 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` |
|
||||
|
||||
## Inference
|
||||
|
||||
|
||||
+21
-81
@@ -44,18 +44,6 @@ docker build -f docker/Dockerfile.benchmark.robomme -t lerobot-robomme .
|
||||
|
||||
The `docker/Dockerfile.benchmark.robomme` image overrides `gymnasium==0.29.1` and `numpy==1.26.4` after lerobot's install. Both versions are runtime-safe for lerobot's actual API usage.
|
||||
|
||||
When evaluating a checkpoint saved on the host, run the container with the host UID. Checkpoint weights are intentionally private by default (`0600`), so the image's built-in user cannot otherwise read a bind-mounted `model.safetensors` file. Mount the Hugging Face cache at an accessible path as well:
|
||||
|
||||
```bash
|
||||
docker run --gpus all --rm --ipc=host \
|
||||
--user "$(id -u):$(id -g)" \
|
||||
-e HF_HOME=/tmp/hf-cache \
|
||||
-v "$HOME/.cache/huggingface:/tmp/hf-cache:ro" \
|
||||
-v "$PWD/outputs:/results" \
|
||||
lerobot-robomme \
|
||||
lerobot-eval --policy.path=/results/<checkpoint>/pretrained_model # ...
|
||||
```
|
||||
|
||||
## Running Evaluation
|
||||
|
||||
### Default (single task, single episode)
|
||||
@@ -73,7 +61,7 @@ lerobot-eval \
|
||||
|
||||
### Multi-task evaluation
|
||||
|
||||
Evaluate multiple tasks in one run by comma-separating task names. Use `task_ids` to select the dataset-backed episodes evaluated for every task. For the standard 50-episode test protocol, select IDs 0 through 49 and run each selected episode once.
|
||||
Evaluate multiple tasks in one run by comma-separating task names. Use `task_ids` to control which episodes are evaluated per task. Recommended: 50 episodes per task for the test split.
|
||||
|
||||
```bash
|
||||
lerobot-eval \
|
||||
@@ -81,22 +69,20 @@ lerobot-eval \
|
||||
--env.type=robomme \
|
||||
--env.task=PickXtimes,BinFill,StopCube,MoveCube,InsertPeg \
|
||||
--env.dataset_split=test \
|
||||
--env.task_ids=[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36,37,38,39,40,41,42,43,44,45,46,47,48,49] \
|
||||
--env.task_ids=[0,1,2,3,4,5,6,7,8,9] \
|
||||
--eval.batch_size=1 \
|
||||
--eval.n_episodes=1
|
||||
--eval.n_episodes=50
|
||||
```
|
||||
|
||||
### Key CLI options for `env.type=robomme`
|
||||
|
||||
| Option | Default | Description |
|
||||
| -------------------- | ------------- | ------------------------------------------------ |
|
||||
| `env.task` | `PickXtimes` | Any of the 16 task names above (comma-separated) |
|
||||
| `env.dataset_split` | `test` | `train`, `val`, or `test` |
|
||||
| `env.action_space` | `joint_angle` | `joint_angle` (8-D) or `ee_pose` (7-D) |
|
||||
| `env.episode_length` | `300` | Max steps per episode |
|
||||
| `env.task_ids` | `null` | Dataset episode indices (null = `[0]`) |
|
||||
|
||||
`eval.n_episodes` repeats each selected `task_id`; it does not advance to the next dataset episode. Keep it at `1` when enumerating the 50 distinct test episodes with `env.task_ids`.
|
||||
| Option | Default | Description |
|
||||
| -------------------- | ------------- | -------------------------------------------------- |
|
||||
| `env.task` | `PickXtimes` | Any of the 16 task names above (comma-separated) |
|
||||
| `env.dataset_split` | `test` | `train`, `val`, or `test` |
|
||||
| `env.action_space` | `joint_angle` | `joint_angle` (8-D) or `ee_pose` (7-D) |
|
||||
| `env.episode_length` | `300` | Max steps per episode |
|
||||
| `env.task_ids` | `null` | List of episode indices to evaluate (null = `[0]`) |
|
||||
|
||||
## Dataset
|
||||
|
||||
@@ -110,68 +96,22 @@ dataset = LeRobotDataset("lerobot/robomme")
|
||||
|
||||
### Dataset features
|
||||
|
||||
| Feature | Shape | Description |
|
||||
| ------------------ | ------------- | -------------------------------- |
|
||||
| `image` | (256, 256, 3) | Front camera RGB |
|
||||
| `wrist_image` | (256, 256, 3) | Wrist camera RGB |
|
||||
| `actions` | (8,) | Joint angles + gripper |
|
||||
| `state` | (8,) | Joint positions + gripper state |
|
||||
| `simple_subgoal` | str | High-level language annotation |
|
||||
| `grounded_subgoal` | str | Grounded language annotation |
|
||||
| `exec_start_idx` | scalar | First action-execution frame |
|
||||
| `is_demo` | bool | Video-demonstration prefix frame |
|
||||
| `episode_index` | int | Episode ID |
|
||||
| `frame_index` | int | Frame within episode |
|
||||
| Feature | Shape | Description |
|
||||
| ------------------ | ------------- | ------------------------------- |
|
||||
| `image` | (256, 256, 3) | Front camera RGB |
|
||||
| `wrist_image` | (256, 256, 3) | Wrist camera RGB |
|
||||
| `actions` | (8,) | Joint angles + gripper |
|
||||
| `state` | (8,) | Joint positions + gripper state |
|
||||
| `simple_subgoal` | str | High-level language annotation |
|
||||
| `grounded_subgoal` | str | Grounded language annotation |
|
||||
| `episode_index` | int | Episode ID |
|
||||
| `frame_index` | int | Frame within episode |
|
||||
|
||||
### Feature key alignment (training)
|
||||
|
||||
The env wrapper exposes `pixels/image` and `pixels/wrist_image` as observation keys. The `features_map` in `RoboMMEEnv` maps these to `observation.images.image` and `observation.images.wrist_image` for the policy. State is exposed as `agent_pos` and maps to `observation.state`.
|
||||
|
||||
The published dataset uses the raw keys shown above. Policies that expect canonical LeRobot keys should map them at training time. For example, the pretrained SmolVLA checkpoint expects three cameras, while RoboMME provides two:
|
||||
|
||||
```bash
|
||||
uv run lerobot-train \
|
||||
--policy.path=lerobot/smolvla_base \
|
||||
--policy.empty_cameras=1 \
|
||||
--dataset.repo_id=lerobot/robomme \
|
||||
--dataset.training_target_start_feature=exec_start_idx \
|
||||
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
||||
```
|
||||
|
||||
RoboMME episodes begin with video-demonstration/context frames that should remain available to a
|
||||
memory policy but should not be sampled as action-training targets. Setting
|
||||
`dataset.training_target_start_feature=exec_start_idx` starts target sampling at each episode's
|
||||
execution boundary while preserving earlier frames for temporal observation deltas. In the published
|
||||
training split this keeps 476,857 execution targets out of 768,897 total frames. One execution-target
|
||||
epoch therefore takes `ceil(476857 / effective_batch_size)` optimizer steps.
|
||||
|
||||
For a sample-matched SmolVLA visual-memory ablation, use
|
||||
`examples/robomme/smolvla_visual_memory_ablation.sh`. `TARGET_SAMPLES` counts examples across all
|
||||
GPUs, and `NUM_PROCESSES` is included when the script converts that target into optimizer steps. For
|
||||
example, the following reproduces 5.12 million example exposures (the exposure of RoboMME's
|
||||
80,000-step, global-batch-64 memory-policy recipe) with four GPUs, two accumulated microbatches,
|
||||
and an effective global batch of 64:
|
||||
|
||||
```bash
|
||||
TARGET_SAMPLES=5120000 \
|
||||
NUM_PROCESSES=4 \
|
||||
BATCH_SIZE=8 \
|
||||
GRADIENT_ACCUMULATION_STEPS=2 \
|
||||
VARIANT=visual-memory \
|
||||
RUN_EVAL=false \
|
||||
bash examples/robomme/smolvla_visual_memory_ablation.sh
|
||||
```
|
||||
|
||||
This becomes 80,000 optimizer steps, or about 10.74 execution-target epochs. Run the baseline with
|
||||
the same `TARGET_SAMPLES`, effective batch size, seed, and scheduler settings to isolate the visual
|
||||
memory change. RoboMME's released baseline uses a different global batch (128), so its nominal
|
||||
80,000-step recipe is not sample-matched to its global-batch-64 memory recipe.
|
||||
|
||||
At evaluation time the environment wrapper has already converted observations to canonical keys. Only the two camera suffixes need to be aligned with the checkpoint:
|
||||
|
||||
```bash
|
||||
'--rename_map={"observation.images.image":"observation.images.camera1","observation.images.wrist_image":"observation.images.camera2"}'
|
||||
```
|
||||
The dataset's `image` and `wrist_image` columns already align with the policy input keys, so no renaming is needed when fine-tuning.
|
||||
|
||||
## Action Spaces
|
||||
|
||||
|
||||
@@ -76,37 +76,6 @@ Fine-tuning is an art. For a complete overview of the options for finetuning, ru
|
||||
lerobot-train --help
|
||||
```
|
||||
|
||||
### Experimental MEM visual memory
|
||||
|
||||
SmolVLA can optionally fuse a short history of camera frames using the
|
||||
space-time separable vision encoder from [MEM](https://arxiv.org/abs/2603.03596).
|
||||
Every fourth SigLIP layer adds causal attention across time for matching image
|
||||
patches. Historical tokens are discarded inside the vision tower, so the VLM
|
||||
receives the same number of image tokens as the baseline policy.
|
||||
|
||||
The option is disabled by default. The following example uses six observations
|
||||
spaced one second apart for a 10 fps dataset:
|
||||
|
||||
```bash
|
||||
lerobot-train \
|
||||
--policy.path=lerobot/smolvla_base \
|
||||
--policy.use_visual_memory=true \
|
||||
--policy.visual_memory_frames=6 \
|
||||
--policy.visual_memory_stride=10 \
|
||||
--policy.visual_memory_temporal_attention_every=4 \
|
||||
--policy.freeze_vision_encoder=false \
|
||||
--policy.train_expert_only=false \
|
||||
--dataset.repo_id=${HF_USER}/mydataset \
|
||||
--output_dir=outputs/train/my_smolvla_mem
|
||||
```
|
||||
|
||||
`visual_memory_stride` is measured in dataset or environment steps. Keep all
|
||||
other settings and the random seed fixed when comparing against a baseline.
|
||||
Because the public SmolVLA checkpoint was not pretrained with this temporal
|
||||
attention pattern, this is a post-training-only MEM ablation; the MEM paper
|
||||
reports that memory-aware pretraining performs better than introducing visual
|
||||
memory only during task-specific fine-tuning.
|
||||
|
||||
<p align="center">
|
||||
<img
|
||||
src="https://cdn-uploads.huggingface.co/production/uploads/640e21ef3c82bd463ee5a76d/S-3vvVCulChREwHDkquoc.gif"
|
||||
|
||||
@@ -46,8 +46,11 @@ CMD = (
|
||||
"apt-get update -qq && apt-get install -y -qq git ffmpeg && "
|
||||
"pip install --no-deps "
|
||||
"'lerobot @ git+https://github.com/huggingface/lerobot.git@main' && "
|
||||
# Pins mirror pyproject.toml — unpinned installs pull av 18 / datasets 5 /
|
||||
# draccus 0.11, which break lerobot at import time.
|
||||
"pip install --upgrade-strategy only-if-needed "
|
||||
"datasets pyarrow av jsonlines draccus gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
||||
"'datasets>=4.7.0,<5.0.0' 'pyarrow>=21.0.0,<30.0.0' 'av>=15.0.0,<16.0.0' 'draccus==0.10.0' "
|
||||
"'pandas>=2.0.0,<3.0.0' jsonlines gymnasium torchcodec mergedeep pyyaml-include toml typing-inspect "
|
||||
"openai && "
|
||||
"export VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS=0 && "
|
||||
"export VLLM_VIDEO_BACKEND=pyav && "
|
||||
|
||||
@@ -1,131 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# Matched SmolVLA ablation for the 10 fps lerobot/robomme dataset.
|
||||
# Run from the LeRobot repository root on a CUDA machine.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
BATCH_SIZE="${BATCH_SIZE:-4}"
|
||||
NUM_PROCESSES="${NUM_PROCESSES:-1}"
|
||||
GRADIENT_ACCUMULATION_STEPS="${GRADIENT_ACCUMULATION_STEPS:-1}"
|
||||
# The published training split has 476,857 execution frames. By default, train on one
|
||||
# execution-frame epoch; set STEPS explicitly to use a different optimizer-step budget.
|
||||
TARGET_SAMPLES="${TARGET_SAMPLES:-476857}"
|
||||
EFFECTIVE_BATCH_SIZE=$((BATCH_SIZE * NUM_PROCESSES * GRADIENT_ACCUMULATION_STEPS))
|
||||
STEPS="${STEPS:-$(((TARGET_SAMPLES + EFFECTIVE_BATCH_SIZE - 1) / EFFECTIVE_BATCH_SIZE))}"
|
||||
SCHEDULER_WARMUP_STEPS="${SCHEDULER_WARMUP_STEPS:-$(((STEPS + 29) / 30))}"
|
||||
SCHEDULER_DECAY_STEPS="${SCHEDULER_DECAY_STEPS:-${STEPS}}"
|
||||
SEED="${SEED:-1000}"
|
||||
OUTPUT_ROOT="${OUTPUT_ROOT:-outputs/robomme-smolvla-mem-ablation}"
|
||||
WANDB_ENABLE="${WANDB_ENABLE:-false}"
|
||||
RUN_TRAIN="${RUN_TRAIN:-true}"
|
||||
RUN_EVAL="${RUN_EVAL:-true}"
|
||||
VARIANT="${VARIANT:-both}"
|
||||
TASKS="BinFill,PickXtimes,SwingXtimes,StopCube,VideoUnmask,VideoUnmaskSwap,ButtonUnmask,ButtonUnmaskSwap,PickHighlight,VideoRepick,VideoPlaceButton,VideoPlaceOrder,MoveCube,InsertPeg,PatternLock,RouteStick"
|
||||
TASK_IDS="[0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15,16,17,18,19,20,21,22,23,24,25,26,27,28,29,30,31,32,33,34,35,36,37,38,39,40,41,42,43,44,45,46,47,48,49]"
|
||||
|
||||
case "${VARIANT}" in
|
||||
baseline) VARIANTS=(baseline) ;;
|
||||
visual-memory) VARIANTS=(visual-memory) ;;
|
||||
both) VARIANTS=(baseline visual-memory) ;;
|
||||
*) echo "VARIANT must be baseline, visual-memory, or both" >&2; exit 2 ;;
|
||||
esac
|
||||
|
||||
TRAIN_COMMAND=()
|
||||
if [[ "${RUN_TRAIN}" == "true" ]]; then
|
||||
if ((NUM_PROCESSES > 1)); then
|
||||
TRAIN_ENTRYPOINT="$(uv run which lerobot-train)"
|
||||
TRAIN_COMMAND=(
|
||||
uv run accelerate launch
|
||||
--multi_gpu
|
||||
--num_processes="${NUM_PROCESSES}"
|
||||
--mixed_precision=bf16
|
||||
"${TRAIN_ENTRYPOINT}"
|
||||
)
|
||||
else
|
||||
TRAIN_COMMAND=(uv run lerobot-train)
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "Training ${VARIANT}: ${TARGET_SAMPLES} target samples, global batch ${EFFECTIVE_BATCH_SIZE}, ${STEPS} optimizer steps"
|
||||
|
||||
COMMON_TRAIN_ARGS=(
|
||||
--policy.path=lerobot/smolvla_base
|
||||
--policy.device=cuda
|
||||
--policy.push_to_hub=false
|
||||
--policy.empty_cameras=1
|
||||
--policy.freeze_vision_encoder=false
|
||||
--policy.train_expert_only=false
|
||||
--dataset.repo_id=lerobot/robomme
|
||||
--dataset.training_target_start_feature=exec_start_idx
|
||||
'--rename_map={"image":"observation.images.camera1","wrist_image":"observation.images.camera2","state":"observation.state","actions":"action"}'
|
||||
--batch_size="${BATCH_SIZE}"
|
||||
--gradient_accumulation_steps="${GRADIENT_ACCUMULATION_STEPS}"
|
||||
--steps="${STEPS}"
|
||||
--policy.scheduler_warmup_steps="${SCHEDULER_WARMUP_STEPS}"
|
||||
--policy.scheduler_decay_steps="${SCHEDULER_DECAY_STEPS}"
|
||||
--seed="${SEED}"
|
||||
--env_eval_freq=0
|
||||
--save_freq=5000
|
||||
--wandb.enable="${WANDB_ENABLE}"
|
||||
)
|
||||
|
||||
if [[ "${RUN_TRAIN}" == "true" ]]; then
|
||||
for variant in "${VARIANTS[@]}"; do
|
||||
MEMORY_ARGS=(--policy.use_visual_memory=false)
|
||||
if [[ "${variant}" == "visual-memory" ]]; then
|
||||
MEMORY_ARGS=(
|
||||
--policy.use_visual_memory=true
|
||||
--policy.visual_memory_frames=6
|
||||
--policy.visual_memory_stride=10
|
||||
--policy.visual_memory_temporal_attention_every=4
|
||||
)
|
||||
fi
|
||||
"${TRAIN_COMMAND[@]}" \
|
||||
"${COMMON_TRAIN_ARGS[@]}" \
|
||||
"${MEMORY_ARGS[@]}" \
|
||||
--output_dir="${OUTPUT_ROOT}/${variant}" \
|
||||
--job_name="robomme-smolvla-${variant}"
|
||||
done
|
||||
fi
|
||||
|
||||
if [[ "${RUN_EVAL}" == "true" ]]; then
|
||||
for variant in "${VARIANTS[@]}"; do
|
||||
uv run lerobot-eval \
|
||||
--policy.path="${OUTPUT_ROOT}/${variant}/checkpoints/last/pretrained_model" \
|
||||
--env.type=robomme \
|
||||
--env.task="${TASKS}" \
|
||||
--env.dataset_split=test \
|
||||
--env.task_ids="${TASK_IDS}" \
|
||||
'--rename_map={"observation.images.image":"observation.images.camera1","observation.images.wrist_image":"observation.images.camera2"}' \
|
||||
--eval.batch_size=1 \
|
||||
--eval.n_episodes=1 \
|
||||
--seed="${SEED}" \
|
||||
--output_dir="${OUTPUT_ROOT}/eval-${variant}"
|
||||
done
|
||||
fi
|
||||
|
||||
if [[ "${RUN_EVAL}" == "true" ]]; then
|
||||
uv run python - "${OUTPUT_ROOT}" <<'PY'
|
||||
import json
|
||||
import pathlib
|
||||
import sys
|
||||
|
||||
root = pathlib.Path(sys.argv[1])
|
||||
for variant in ("baseline", "visual-memory"):
|
||||
result_path = root / f"eval-{variant}" / "eval_info.json"
|
||||
if not result_path.exists():
|
||||
continue
|
||||
with result_path.open() as handle:
|
||||
info = json.load(handle)
|
||||
overall = info["overall"]
|
||||
print(
|
||||
f"{variant}: success={overall['pc_success']:.2f}% "
|
||||
f"avg_reward={overall['avg_sum_reward']:.4f}"
|
||||
)
|
||||
for task, metrics in info["per_group"].items():
|
||||
print(
|
||||
f" {task}: success={metrics['pc_success']:.2f}% "
|
||||
f"avg_reward={metrics['avg_sum_reward']:.4f}"
|
||||
)
|
||||
PY
|
||||
fi
|
||||
+1
-2
@@ -150,7 +150,6 @@ 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"]
|
||||
@@ -213,7 +212,7 @@ wallx = [
|
||||
"torchdiffeq>=0.2.4,<0.3.0",
|
||||
"lerobot[qwen-vl-utils-dep]",
|
||||
]
|
||||
pi = ["lerobot[transformers-dep]", "lerobot[scipy-dep]", "lerobot[sentencepiece-dep]"]
|
||||
pi = ["lerobot[transformers-dep]", "lerobot[scipy-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]"]
|
||||
|
||||
@@ -33,10 +33,6 @@ class DatasetConfig:
|
||||
# looked up under $HF_LEROBOT_HOME/repo_id and Hub downloads use a revision-safe cache under $HF_LEROBOT_HOME/hub.
|
||||
root: str | None = None
|
||||
episodes: list[int] | None = None
|
||||
# Optional per-frame feature containing the episode-relative index of the first frame that may be
|
||||
# sampled as a training target. Earlier frames remain in the dataset and can still be loaded through
|
||||
# temporal observation deltas. This is useful for datasets with demonstration/context prefixes.
|
||||
training_target_start_feature: str | None = None
|
||||
image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig)
|
||||
revision: str | None = None
|
||||
use_imagenet_stats: bool = True
|
||||
|
||||
@@ -147,7 +147,7 @@ class TrainingRecipe:
|
||||
return cls.from_dict(data)
|
||||
|
||||
def _validate_message_recipe(self) -> None:
|
||||
"""Validate bindings and require text or low-level action supervision."""
|
||||
"""Ensure every templated binding is known and at least one turn is a target."""
|
||||
assert self.messages is not None
|
||||
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
|
||||
|
||||
@@ -156,14 +156,8 @@ class TrainingRecipe:
|
||||
if missing:
|
||||
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
|
||||
|
||||
has_target = any(turn.target for turn in self.messages)
|
||||
has_low_level = any(turn.stream == "low_level" for turn in self.messages)
|
||||
if not (has_target or has_low_level):
|
||||
raise ValueError(
|
||||
"Message recipes must contain at least one supervised turn — "
|
||||
"either ``target: true`` (text CE) or ``stream: low_level`` "
|
||||
"(flow/action loss)."
|
||||
)
|
||||
if not any(turn.target for turn in self.messages):
|
||||
raise ValueError("Message recipes must contain at least one target turn.")
|
||||
|
||||
def _validate_blend_recipe(self) -> None:
|
||||
"""Ensure each blend component is a non-empty, weighted message recipe."""
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
# Predicts subtasks from tasks and trains subtask-conditioned action flow without memory or plans.
|
||||
# Requires `subtask` annotations; samples with missing `if_present` bindings do not render.
|
||||
|
||||
blend:
|
||||
|
||||
high_level_subtask:
|
||||
weight: 0.30
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||
|
||||
low_level_execution:
|
||||
weight: 0.70
|
||||
messages:
|
||||
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||
@@ -1,30 +0,0 @@
|
||||
# Trains subtask prediction, subtask-conditioned action flow, and memory updates without plans.
|
||||
# Requires `subtask` and `memory`; missing `if_present` bindings skip the affected sub-recipe.
|
||||
|
||||
blend:
|
||||
|
||||
high_level_subtask:
|
||||
weight: 0.25
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||
|
||||
low_level_execution:
|
||||
weight: 0.60
|
||||
messages:
|
||||
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||
|
||||
memory_update:
|
||||
# `active_at` densifies sparse boundaries while preserving the prior-memory/subtask mapping.
|
||||
# Inference controls update timing through `subtask_change` events.
|
||||
weight: 0.15
|
||||
bindings:
|
||||
prior_memory: "nth_prev(style=memory, offset=1)"
|
||||
current_memory: "active_at(t, style=memory)"
|
||||
completed_subtask: "nth_prev(style=subtask, offset=1)"
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: assistant, content: "Previous memory: ${prior_memory}", stream: high_level, if_present: prior_memory}
|
||||
- {role: user, content: "Completed subtask: ${completed_subtask}", stream: high_level, if_present: completed_subtask}
|
||||
- {role: assistant, content: "${current_memory}", stream: high_level, target: true, if_present: current_memory}
|
||||
@@ -1,70 +0,0 @@
|
||||
# Adds memory, spoken interjection responses, and camera-grounded VQA to subtask/action training.
|
||||
# Missing optional annotations skip only their sub-recipe; `say` tool calls tokenize as `<say>...</say>`.
|
||||
|
||||
blend:
|
||||
|
||||
high_level_subtask:
|
||||
weight: 0.25
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: assistant, content: "${subtask}", stream: high_level, target: true, if_present: subtask}
|
||||
|
||||
low_level_execution:
|
||||
weight: 0.40
|
||||
messages:
|
||||
# The low-level stream trains action flow on the generated or annotated subtask.
|
||||
- {role: user, content: "${subtask}", stream: low_level, if_present: subtask}
|
||||
|
||||
memory_update:
|
||||
# `active_at` densifies sparse boundaries while preserving the prior-memory/subtask mapping.
|
||||
# Inference controls update timing through `subtask_change` events.
|
||||
weight: 0.10
|
||||
bindings:
|
||||
prior_memory: "nth_prev(style=memory, offset=1)"
|
||||
current_memory: "active_at(t, style=memory)"
|
||||
completed_subtask: "nth_prev(style=subtask, offset=1)"
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: assistant, content: "Previous memory: ${prior_memory}", stream: high_level, if_present: prior_memory}
|
||||
- {role: user, content: "Completed subtask: ${completed_subtask}", stream: high_level, if_present: completed_subtask}
|
||||
- {role: assistant, content: "${current_memory}", stream: high_level, target: true, if_present: current_memory}
|
||||
|
||||
user_interjection_response:
|
||||
weight: 0.10
|
||||
bindings:
|
||||
interjection: "emitted_at(t, style=interjection)"
|
||||
speech: "emitted_at(t, role=assistant, tool_name=say)"
|
||||
messages:
|
||||
- {role: user, content: "${task}", stream: high_level}
|
||||
- {role: user, content: "${interjection}", stream: high_level, if_present: interjection}
|
||||
# The assistant target is a `say` tool call flattened to a `<say>...</say>` marker.
|
||||
- {role: assistant, stream: high_level, target: true, if_present: speech, tool_calls_from: speech}
|
||||
|
||||
# Each camera uses a separate VQA sub-recipe for view-specific binding.
|
||||
ask_vqa_top:
|
||||
weight: 0.075
|
||||
bindings:
|
||||
vqa_query: "emitted_at(t, style=vqa, role=user, camera=observation.images.front)"
|
||||
vqa: "emitted_at(t, style=vqa, role=assistant, camera=observation.images.front)"
|
||||
messages:
|
||||
- role: user
|
||||
stream: high_level
|
||||
if_present: vqa_query
|
||||
content:
|
||||
- {type: image, feature: observation.images.front}
|
||||
- {type: text, text: "${vqa_query}"}
|
||||
- {role: assistant, content: "${vqa}", stream: high_level, target: true, if_present: vqa}
|
||||
|
||||
ask_vqa_wrist:
|
||||
weight: 0.075
|
||||
bindings:
|
||||
vqa_query: "emitted_at(t, style=vqa, role=user, camera=observation.images.wrist)"
|
||||
vqa: "emitted_at(t, style=vqa, role=assistant, camera=observation.images.wrist)"
|
||||
messages:
|
||||
- role: user
|
||||
stream: high_level
|
||||
if_present: vqa_query
|
||||
content:
|
||||
- {type: image, feature: observation.images.wrist}
|
||||
- {type: text, text: "${vqa_query}"}
|
||||
- {role: assistant, content: "${vqa}", stream: high_level, target: true, if_present: vqa}
|
||||
@@ -99,9 +99,6 @@ class TrainPipelineConfig(HubMixin):
|
||||
# Number of workers for the dataloader.
|
||||
num_workers: int = 4
|
||||
batch_size: int = 8
|
||||
# Number of microbatches accumulated before each optimizer update. The effective global batch is
|
||||
# batch_size * accelerator.num_processes * gradient_accumulation_steps.
|
||||
gradient_accumulation_steps: int = 1
|
||||
prefetch_factor: int = 4
|
||||
persistent_workers: bool = True
|
||||
steps: int = 100_000
|
||||
@@ -224,8 +221,6 @@ class TrainPipelineConfig(HubMixin):
|
||||
)
|
||||
|
||||
active_cfg = self.trainable_config
|
||||
if self.gradient_accumulation_steps < 1:
|
||||
raise ValueError("gradient_accumulation_steps must be at least 1.")
|
||||
if self.rename_map and active_cfg.pretrained_path is None:
|
||||
raise ValueError(
|
||||
"`rename_map` requires a pretrained policy checkpoint. "
|
||||
|
||||
@@ -163,40 +163,10 @@ class DatasetReader:
|
||||
def _load_hf_dataset(self) -> datasets.Dataset:
|
||||
"""hf_dataset contains all the observations, states, actions, rewards, etc."""
|
||||
features = get_hf_features_from_features(self._meta.features)
|
||||
# Annotated datasets may have language columns absent from metadata.
|
||||
# Extend the schema before the strict Parquet cast.
|
||||
features = self._extend_features_with_language_columns(features)
|
||||
hf_dataset = load_nested_dataset(self.root / "data", features=features, episodes=self.episodes)
|
||||
hf_dataset.set_transform(hf_transform_to_torch)
|
||||
return hf_dataset
|
||||
|
||||
def _extend_features_with_language_columns(self, features: datasets.Features) -> datasets.Features:
|
||||
"""Register language columns found in Parquet but missing from metadata."""
|
||||
# Leave empty datasets to fail through the normal loading path.
|
||||
try:
|
||||
sample = next((self.root / "data").glob("*/*.parquet"))
|
||||
except StopIteration:
|
||||
return features
|
||||
|
||||
from pyarrow import parquet as _pq # noqa: PLC0415
|
||||
|
||||
schema_names = set(_pq.read_schema(sample).names)
|
||||
from .language import ( # noqa: PLC0415
|
||||
LANGUAGE_EVENTS,
|
||||
LANGUAGE_PERSISTENT,
|
||||
language_events_column_feature,
|
||||
language_persistent_column_feature,
|
||||
)
|
||||
|
||||
extra: dict[str, object] = {}
|
||||
if LANGUAGE_PERSISTENT in schema_names and LANGUAGE_PERSISTENT not in features:
|
||||
extra[LANGUAGE_PERSISTENT] = language_persistent_column_feature()
|
||||
if LANGUAGE_EVENTS in schema_names and LANGUAGE_EVENTS not in features:
|
||||
extra[LANGUAGE_EVENTS] = language_events_column_feature()
|
||||
if not extra:
|
||||
return features
|
||||
return datasets.Features({**features, **extra})
|
||||
|
||||
def _check_cached_episodes_sufficient(self) -> bool:
|
||||
"""Check if the cached dataset contains all requested episodes and their video files."""
|
||||
if self.hf_dataset is None or len(self.hf_dataset) == 0:
|
||||
|
||||
@@ -32,9 +32,7 @@ from .streaming_dataset import StreamingLeRobotDataset
|
||||
|
||||
|
||||
def resolve_delta_timestamps(
|
||||
cfg: PreTrainedConfig | RewardModelConfig,
|
||||
ds_meta: LeRobotDatasetMetadata,
|
||||
rename_map: dict[str, str] | None = None,
|
||||
cfg: PreTrainedConfig | RewardModelConfig, ds_meta: LeRobotDatasetMetadata
|
||||
) -> dict[str, list] | None:
|
||||
"""Resolves delta_timestamps by reading from the 'delta_indices' properties of the config.
|
||||
|
||||
@@ -55,12 +53,11 @@ def resolve_delta_timestamps(
|
||||
"""
|
||||
delta_timestamps = {}
|
||||
for key in ds_meta.features:
|
||||
policy_key = (rename_map or {}).get(key, key)
|
||||
if policy_key == REWARD and cfg.reward_delta_indices is not None:
|
||||
if key == REWARD and cfg.reward_delta_indices is not None:
|
||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.reward_delta_indices]
|
||||
if policy_key == ACTION and cfg.action_delta_indices is not None:
|
||||
if key == ACTION and cfg.action_delta_indices is not None:
|
||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.action_delta_indices]
|
||||
if policy_key.startswith(OBS_PREFIX) and cfg.observation_delta_indices is not None:
|
||||
if key.startswith(OBS_PREFIX) and cfg.observation_delta_indices is not None:
|
||||
delta_timestamps[key] = [i / ds_meta.fps for i in cfg.observation_delta_indices]
|
||||
|
||||
if len(delta_timestamps) == 0:
|
||||
@@ -89,7 +86,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
ds_meta = LeRobotDatasetMetadata(
|
||||
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
|
||||
)
|
||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta, cfg.rename_map)
|
||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
|
||||
if not cfg.dataset.streaming:
|
||||
dataset = LeRobotDataset(
|
||||
cfg.dataset.repo_id,
|
||||
@@ -133,9 +130,6 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
|
||||
for key in dataset.meta.camera_keys:
|
||||
if key in dataset.meta.depth_keys:
|
||||
continue # Exclude depth keys from ImageNet stats
|
||||
# Some video-only datasets omit camera statistics entirely. Visual
|
||||
# normalization can still use the requested ImageNet defaults.
|
||||
dataset.meta.stats.setdefault(key, {})
|
||||
for stats_type, stats in IMAGENET_STATS.items():
|
||||
dataset.meta.stats[key][stats_type] = torch.tensor(stats, dtype=torch.float32)
|
||||
|
||||
@@ -181,7 +175,7 @@ def make_train_eval_datasets(
|
||||
f"(eval_split={cfg.dataset.eval_split}, {len(task_to_episodes)} tasks)"
|
||||
)
|
||||
|
||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, full_dataset.meta, cfg.rename_map)
|
||||
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, full_dataset.meta)
|
||||
|
||||
train_image_transforms = (
|
||||
ImageTransforms(cfg.dataset.image_transforms) if cfg.dataset.image_transforms.enable else None
|
||||
|
||||
@@ -162,28 +162,14 @@ def render_sample(
|
||||
task: str | None = None,
|
||||
dataset_ctx: Any | None = None,
|
||||
) -> RenderedMessages | None:
|
||||
"""Resolve one sample's bindings and render its message recipe.
|
||||
"""Render the chat-style messages for a single dataset sample.
|
||||
|
||||
Returns ``None`` when no text or low-level action supervision applies.
|
||||
Resolves the recipe's bindings against ``persistent`` and ``events`` rows
|
||||
at frame timestamp ``t``, then expands the recipe's message templates.
|
||||
Returns ``None`` if the resolved sample contains no target message.
|
||||
"""
|
||||
persistent_rows = _normalize_rows(persistent or [])
|
||||
event_rows = _normalize_rows(events or [])
|
||||
|
||||
# Route sparse VQA frames to a matching view-specific component before weighted selection.
|
||||
# This avoids dropping annotated frames or selecting VQA without annotations.
|
||||
if recipe.blend is not None:
|
||||
vqa_rendered = _render_vqa_if_present(
|
||||
recipe,
|
||||
persistent=persistent_rows,
|
||||
events=event_rows,
|
||||
t=t,
|
||||
sample_idx=sample_idx,
|
||||
task=task,
|
||||
dataset_ctx=dataset_ctx,
|
||||
)
|
||||
if vqa_rendered is not None:
|
||||
return vqa_rendered
|
||||
|
||||
selected_recipe = _select_recipe(recipe, sample_idx)
|
||||
bindings = _resolve_bindings(
|
||||
selected_recipe,
|
||||
@@ -197,55 +183,6 @@ def render_sample(
|
||||
return _render_message_recipe(selected_recipe, bindings)
|
||||
|
||||
|
||||
def _render_vqa_if_present(
|
||||
recipe: TrainingRecipe,
|
||||
*,
|
||||
persistent: Sequence[LanguageRow],
|
||||
events: Sequence[LanguageRow],
|
||||
t: float,
|
||||
sample_idx: int,
|
||||
task: str | None,
|
||||
dataset_ctx: Any | None,
|
||||
) -> RenderedMessages | None:
|
||||
"""Render a matching VQA component, or return ``None`` for normal selection.
|
||||
|
||||
Multiple matching views are selected deterministically by relative weight.
|
||||
"""
|
||||
assert recipe.blend is not None
|
||||
renderable: list[tuple[float, RenderedMessages]] = []
|
||||
for name, component in recipe.blend.items():
|
||||
if not name.startswith("ask_vqa"):
|
||||
continue
|
||||
bindings = _resolve_bindings(
|
||||
component,
|
||||
persistent=persistent,
|
||||
events=events,
|
||||
t=t,
|
||||
sample_idx=sample_idx,
|
||||
task=task,
|
||||
dataset_ctx=dataset_ctx,
|
||||
)
|
||||
rendered = _render_message_recipe(component, bindings)
|
||||
if rendered is not None:
|
||||
renderable.append((float(component.weight or 0.0), rendered))
|
||||
|
||||
if not renderable:
|
||||
return None
|
||||
if len(renderable) == 1:
|
||||
return renderable[0][1]
|
||||
|
||||
# Choose among matching cameras by relative weight, or uniformly when all weights are zero.
|
||||
total = sum(w for w, _ in renderable) or float(len(renderable))
|
||||
digest = hashlib.blake2b(f"vqa:{sample_idx}".encode(), digest_size=8).digest()
|
||||
draw = int.from_bytes(digest, "big") / 2**64 * total
|
||||
cumulative = 0.0
|
||||
for w, rendered in renderable:
|
||||
cumulative += w or (total / len(renderable))
|
||||
if draw < cumulative:
|
||||
return rendered
|
||||
return renderable[-1][1]
|
||||
|
||||
|
||||
def _select_recipe(recipe: TrainingRecipe, sample_idx: int) -> TrainingRecipe:
|
||||
"""Pick a deterministic blend component for ``sample_idx`` (or return ``recipe``)."""
|
||||
if recipe.blend is None:
|
||||
@@ -409,9 +346,7 @@ def _render_message_recipe(
|
||||
if turn.target:
|
||||
target_indices.append(message_idx)
|
||||
|
||||
# Keep samples with either text targets or low-level action supervision.
|
||||
has_low_level = any(stream == "low_level" for stream in streams)
|
||||
if not target_indices and not has_low_level:
|
||||
if not target_indices:
|
||||
return None
|
||||
|
||||
rendered = {
|
||||
@@ -468,12 +403,14 @@ def _validate_rendered(rendered: RenderedMessages) -> None:
|
||||
|
||||
if len(streams) != len(messages):
|
||||
raise ValueError("message_streams must be aligned with messages.")
|
||||
# Require text or low-level action supervision.
|
||||
if not target_indices and not any(s == "low_level" for s in streams):
|
||||
raise ValueError("Rendered samples must contain a target message or a low_level-stream message.")
|
||||
if not target_indices:
|
||||
raise ValueError("Rendered samples must contain at least one target message.")
|
||||
for idx in target_indices:
|
||||
if idx < 0 or idx >= len(messages):
|
||||
raise ValueError(f"Target message index {idx} is out of bounds.")
|
||||
# ``stream`` is enforced non-None at MessageTurn construction time
|
||||
# (see ``MessageTurn.__post_init__``), so a missing stream here would
|
||||
# mean the dataclass invariant was bypassed; no need to re-check.
|
||||
|
||||
|
||||
def _nth_relative(
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
# limitations under the License.
|
||||
import logging
|
||||
import math
|
||||
from collections.abc import Iterator, Sequence
|
||||
from collections.abc import Iterator
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -49,7 +49,7 @@ class EpisodeAwareSampler:
|
||||
dataset_from_indices: list[int],
|
||||
dataset_to_indices: list[int],
|
||||
episode_indices_to_use: list | None = None,
|
||||
drop_n_first_frames: int | Sequence[int] = 0,
|
||||
drop_n_first_frames: int = 0,
|
||||
drop_n_last_frames: int = 0,
|
||||
shuffle: bool = False,
|
||||
seed: int = 0,
|
||||
@@ -60,12 +60,13 @@ class EpisodeAwareSampler:
|
||||
dataset_from_indices: Start index of each episode in the dataset.
|
||||
dataset_to_indices: End index of each episode in the dataset.
|
||||
episode_indices_to_use: Episode indices to use; None means all.
|
||||
drop_n_first_frames: Frames to drop from the start of each episode. An integer applies the
|
||||
same offset to every episode; a sequence supplies one offset per episode.
|
||||
drop_n_first_frames: Frames to drop from the start of each episode.
|
||||
drop_n_last_frames: Frames to drop from the end of each episode.
|
||||
shuffle: Whether to shuffle the indices.
|
||||
seed: Seed the permutation is derived from (together with the epoch).
|
||||
"""
|
||||
if drop_n_first_frames < 0:
|
||||
raise ValueError(f"drop_n_first_frames must be >= 0, got {drop_n_first_frames}")
|
||||
if drop_n_last_frames < 0:
|
||||
raise ValueError(f"drop_n_last_frames must be >= 0, got {drop_n_last_frames}")
|
||||
|
||||
@@ -77,27 +78,12 @@ class EpisodeAwareSampler:
|
||||
f"got {len(from_indices)} and {len(to_indices)}"
|
||||
)
|
||||
|
||||
if isinstance(drop_n_first_frames, int):
|
||||
first_frame_offsets = np.full(len(from_indices), drop_n_first_frames, dtype=np.int64)
|
||||
else:
|
||||
first_frame_offsets = np.asarray(drop_n_first_frames, dtype=np.int64)
|
||||
if first_frame_offsets.shape != from_indices.shape:
|
||||
raise ValueError(
|
||||
"drop_n_first_frames must be an integer or have one value per episode; "
|
||||
f"got {len(first_frame_offsets)} values for {len(from_indices)} episodes"
|
||||
)
|
||||
if np.any(first_frame_offsets < 0):
|
||||
raise ValueError(
|
||||
"drop_n_first_frames must be >= 0, got "
|
||||
f"{first_frame_offsets[first_frame_offsets < 0].tolist()}"
|
||||
)
|
||||
|
||||
used = np.ones(len(from_indices), dtype=bool)
|
||||
if episode_indices_to_use is not None:
|
||||
used = np.zeros(len(from_indices), dtype=bool)
|
||||
used[np.asarray(episode_indices_to_use, dtype=np.int64)] = True
|
||||
|
||||
starts = from_indices + first_frame_offsets
|
||||
starts = from_indices + drop_n_first_frames
|
||||
lengths = to_indices - drop_n_last_frames - starts
|
||||
for episode_idx in np.flatnonzero(used & (lengths <= 0)):
|
||||
logger.warning(
|
||||
@@ -105,7 +91,7 @@ class EpisodeAwareSampler:
|
||||
"drop_n_last_frames=%d removes all frames. Skipping.",
|
||||
episode_idx,
|
||||
to_indices[episode_idx] - from_indices[episode_idx],
|
||||
first_frame_offsets[episode_idx],
|
||||
drop_n_first_frames,
|
||||
drop_n_last_frames,
|
||||
)
|
||||
used &= lengths > 0
|
||||
|
||||
@@ -556,13 +556,7 @@ class RoboCasaEnv(EnvConfig):
|
||||
kwargs["split"] = self.split
|
||||
return kwargs
|
||||
|
||||
def create_envs(
|
||||
self,
|
||||
n_envs: int,
|
||||
use_async_envs: bool = False,
|
||||
terminate_on_success: bool = True,
|
||||
horizon: int | None = None,
|
||||
):
|
||||
def create_envs(self, n_envs: int, use_async_envs: bool = False):
|
||||
from .robocasa import create_robocasa_envs
|
||||
|
||||
if self.task is None:
|
||||
@@ -576,8 +570,6 @@ class RoboCasaEnv(EnvConfig):
|
||||
env_cls=env_cls,
|
||||
episode_length=self.episode_length,
|
||||
obj_registries=tuple(self.obj_registries),
|
||||
terminate_on_success=terminate_on_success,
|
||||
horizon=horizon,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -33,8 +33,8 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# Dimensions for the flat action/state vectors used by the LeRobot wrapper.
|
||||
# These correspond to the PandaOmron robot in RoboCasa365.
|
||||
OBS_STATE_DIM = 16 # ee_pos_rel(3) + ee_quat_rel(4) + base_pos(3) + base_quat(4) + gripper_qpos(2)
|
||||
ACTION_DIM = 12 # ee_pos(3) + ee_rot(3) + gripper(1) + base_motion(4) + control_mode(1)
|
||||
OBS_STATE_DIM = 16 # base_pos(3) + base_quat(4) + ee_pos_rel(3) + ee_quat_rel(4) + gripper_qpos(2)
|
||||
ACTION_DIM = 12 # base_motion(4) + control_mode(1) + ee_pos(3) + ee_rot(3) + gripper(1)
|
||||
ACTION_LOW = -1.0
|
||||
ACTION_HIGH = 1.0
|
||||
|
||||
@@ -101,15 +101,14 @@ def _resolve_tasks(task: str) -> tuple[list[str], str | None]:
|
||||
def convert_action(flat_action: np.ndarray) -> dict[str, Any]:
|
||||
"""Split a flat (12,) action vector into a RoboCasa action dict.
|
||||
|
||||
Layout (openpi / robocasa.utils.env_utils.convert_action order):
|
||||
ee_pos(3) + ee_rot(3) + gripper(1) + base_motion(4) + control_mode(1)
|
||||
Layout: base_motion(4) + control_mode(1) + ee_pos(3) + ee_rot(3) + gripper(1)
|
||||
"""
|
||||
return {
|
||||
"action.end_effector_position": flat_action[0:3],
|
||||
"action.end_effector_rotation": flat_action[3:6],
|
||||
"action.gripper_close": flat_action[6:7],
|
||||
"action.base_motion": flat_action[7:11],
|
||||
"action.control_mode": flat_action[11:12],
|
||||
"action.base_motion": flat_action[0:4],
|
||||
"action.control_mode": flat_action[4:5],
|
||||
"action.end_effector_position": flat_action[5:8],
|
||||
"action.end_effector_rotation": flat_action[8:11],
|
||||
"action.gripper_close": flat_action[11:12],
|
||||
}
|
||||
|
||||
|
||||
@@ -137,16 +136,9 @@ class RoboCasaEnv(gym.Env):
|
||||
episode_length: int | None = None,
|
||||
obj_registries: Sequence[str] = DEFAULT_OBJ_REGISTRIES,
|
||||
episode_index: int = 0,
|
||||
terminate_on_success: bool = True,
|
||||
horizon: int | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.task = task
|
||||
# When False, a task-success does NOT end/reset the episode — used by the
|
||||
# interactive sim so one kitchen persists across sequential prompts.
|
||||
self.terminate_on_success = terminate_on_success
|
||||
# Underlying robosuite horizon (steps before truncation). None -> default.
|
||||
self.horizon = horizon
|
||||
self.obs_type = obs_type
|
||||
self.render_mode = render_mode
|
||||
self.observation_width = observation_width
|
||||
@@ -218,16 +210,12 @@ class RoboCasaEnv(gym.Env):
|
||||
# (only None/"all"/"pretrain"/"target" are valid). Always pass a
|
||||
# valid value so we don't hit that default. Extra kwargs are
|
||||
# forwarded to the underlying kitchen env via create_env/robosuite.make.
|
||||
extra_kwargs: dict[str, Any] = {}
|
||||
if self.horizon is not None:
|
||||
extra_kwargs["horizon"] = int(self.horizon)
|
||||
self._env = RoboCasaGymEnv(
|
||||
env_name=self.task,
|
||||
camera_widths=self.observation_width,
|
||||
camera_heights=self.observation_height,
|
||||
split=self.split if self.split is not None else "all",
|
||||
obj_registries=self.obj_registries,
|
||||
**extra_kwargs,
|
||||
)
|
||||
|
||||
ep_meta = self._env.env.get_ep_meta()
|
||||
@@ -242,14 +230,12 @@ class RoboCasaEnv(gym.Env):
|
||||
return {"pixels": images}
|
||||
|
||||
# `state.*` keys come from PandaOmronKeyConverter inside the wrapper.
|
||||
# openpi state order: ee first, then base, then gripper (matches the
|
||||
# openpi robocasa pipeline / examples/robocasa/main.py state layout).
|
||||
agent_pos = np.concatenate(
|
||||
[
|
||||
raw_obs.get("state.end_effector_position_relative", np.zeros(3)),
|
||||
raw_obs.get("state.end_effector_rotation_relative", np.zeros(4)),
|
||||
raw_obs.get("state.base_position", np.zeros(3)),
|
||||
raw_obs.get("state.base_rotation", np.zeros(4)),
|
||||
raw_obs.get("state.end_effector_position_relative", np.zeros(3)),
|
||||
raw_obs.get("state.end_effector_rotation_relative", np.zeros(4)),
|
||||
raw_obs.get("state.gripper_qpos", np.zeros(2)),
|
||||
],
|
||||
axis=-1,
|
||||
@@ -294,7 +280,7 @@ class RoboCasaEnv(gym.Env):
|
||||
raw_obs, reward, done, truncated, info = self._env.step(action_dict)
|
||||
|
||||
is_success = bool(info.get("success", False))
|
||||
terminated = done or (is_success and self.terminate_on_success)
|
||||
terminated = done or is_success
|
||||
info.update({"task": self.task, "done": done, "is_success": is_success})
|
||||
|
||||
observation = self._format_raw_obs(raw_obs)
|
||||
@@ -327,8 +313,6 @@ def _make_env_fns(
|
||||
split: str | None,
|
||||
episode_length: int | None,
|
||||
obj_registries: Sequence[str],
|
||||
terminate_on_success: bool = True,
|
||||
horizon: int | None = None,
|
||||
) -> list[Callable[[], RoboCasaEnv]]:
|
||||
"""Build n_envs factory callables for a single task.
|
||||
|
||||
@@ -351,8 +335,6 @@ def _make_env_fns(
|
||||
episode_length=episode_length,
|
||||
obj_registries=obj_registries,
|
||||
episode_index=episode_index,
|
||||
terminate_on_success=terminate_on_success,
|
||||
horizon=horizon,
|
||||
)
|
||||
|
||||
return [partial(_make_env, i) for i in range(n_envs)]
|
||||
@@ -366,8 +348,6 @@ def create_robocasa_envs(
|
||||
env_cls: Callable[[Sequence[Callable[[], Any]]], Any] | None = None,
|
||||
episode_length: int | None = None,
|
||||
obj_registries: Sequence[str] = DEFAULT_OBJ_REGISTRIES,
|
||||
terminate_on_success: bool = True,
|
||||
horizon: int | None = None,
|
||||
) -> dict[str, dict[int, Any]]:
|
||||
"""Create vectorized RoboCasa365 environments with a consistent return shape.
|
||||
|
||||
@@ -429,8 +409,6 @@ def create_robocasa_envs(
|
||||
split=split,
|
||||
episode_length=episode_length,
|
||||
obj_registries=obj_registries,
|
||||
terminate_on_success=terminate_on_success,
|
||||
horizon=horizon,
|
||||
)
|
||||
|
||||
if is_async:
|
||||
|
||||
@@ -63,8 +63,6 @@ class RoboMMEGymEnv(gym.Env):
|
||||
from robomme.env_record_wrapper import BenchmarkEnvBuilder
|
||||
|
||||
self._task = task
|
||||
self.task = task
|
||||
self.task_description = task
|
||||
self._action_space_type = action_space_type
|
||||
self._dataset = dataset
|
||||
self._episode_idx = episode_idx
|
||||
@@ -107,12 +105,6 @@ class RoboMMEGymEnv(gym.Env):
|
||||
)
|
||||
obs, info = self._env.reset()
|
||||
self._last_raw_obs = obs
|
||||
task_goal = info.get("task_goal")
|
||||
# RoboMME returns [simple_subgoal, grounded_subgoal]. The published
|
||||
# LeRobot training dataset uses the simple subgoal as its `task` text.
|
||||
if isinstance(task_goal, (list, tuple)):
|
||||
task_goal = task_goal[0] if task_goal else ""
|
||||
self.task_description = str(task_goal or self._task)
|
||||
return self._convert_obs(obs), self._convert_info(info)
|
||||
|
||||
def step(self, action):
|
||||
|
||||
@@ -104,8 +104,6 @@ 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)
|
||||
|
||||
@@ -28,7 +28,6 @@ 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,7 +55,6 @@ __all__ = [
|
||||
"PI0Config",
|
||||
"PI0FastConfig",
|
||||
"PI05Config",
|
||||
"PI052Config",
|
||||
"SmolVLAConfig",
|
||||
"TDMPCConfig",
|
||||
"VQBeTConfig",
|
||||
|
||||
@@ -133,10 +133,6 @@ def get_policy_class(name: str) -> type[PreTrainedPolicy]:
|
||||
from .pi05.modeling_pi05 import PI05Policy
|
||||
|
||||
return PI05Policy
|
||||
elif name == "pi052":
|
||||
from .pi052.modeling_pi052 import PI052Policy
|
||||
|
||||
return PI052Policy
|
||||
elif name == "gaussian_actor":
|
||||
from .gaussian_actor.modeling_gaussian_actor import GaussianActorPolicy
|
||||
|
||||
@@ -197,8 +193,8 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
|
||||
Args:
|
||||
policy_type: The type of the policy. Supported types include "tdmpc",
|
||||
"multi_task_dit", "diffusion", "act", "vqbet", "pi0", "pi05", "pi052",
|
||||
"gaussian_actor", "smolvla", "wall_x", "molmoact2", "eo1", "evo1".
|
||||
"multi_task_dit", "diffusion", "act", "vqbet", "pi0", "pi05", "gaussian_actor",
|
||||
"smolvla", "wall_x", "molmoact2", "eo1", "evo1".
|
||||
**kwargs: Keyword arguments to be passed to the configuration class constructor.
|
||||
|
||||
Returns:
|
||||
@@ -221,10 +217,6 @@ def make_policy_config(policy_type: str, **kwargs) -> PreTrainedConfig:
|
||||
return PI0Config(**kwargs)
|
||||
elif policy_type == "pi05":
|
||||
return PI05Config(**kwargs)
|
||||
elif policy_type == "pi052":
|
||||
from .pi052.configuration_pi052 import PI052Config
|
||||
|
||||
return PI052Config(**kwargs)
|
||||
elif policy_type == "gaussian_actor":
|
||||
return GaussianActorConfig(**kwargs)
|
||||
elif policy_type == "smolvla":
|
||||
@@ -275,8 +267,6 @@ 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 repo used for optional processor fitting; omit it to use universal tokenizers.
|
||||
dataset_repo_id: str | None
|
||||
dataset_meta: Any | None
|
||||
|
||||
|
||||
@@ -311,22 +301,6 @@ def make_pre_post_processors(
|
||||
NotImplementedError: If a processor factory is not implemented for the given
|
||||
policy configuration type.
|
||||
"""
|
||||
if (
|
||||
pretrained_path
|
||||
and getattr(policy_cfg, "type", None) in {"pi0_fast", "pi052"}
|
||||
and getattr(policy_cfg, "auto_fit_fast_tokenizer", False)
|
||||
and kwargs.get("dataset_repo_id") is not None
|
||||
):
|
||||
from .pi052.fit_fast_tokenizer import resolve_fast_tokenizer
|
||||
|
||||
overrides = dict(kwargs.get("preprocessor_overrides") or {})
|
||||
action_tokenizer_override = {
|
||||
**overrides.get("action_tokenizer_processor", {}),
|
||||
"action_tokenizer_name": resolve_fast_tokenizer(policy_cfg, kwargs.get("dataset_repo_id")),
|
||||
}
|
||||
overrides["action_tokenizer_processor"] = action_tokenizer_override
|
||||
kwargs["preprocessor_overrides"] = overrides
|
||||
|
||||
if pretrained_path:
|
||||
if isinstance(policy_cfg, GrootConfig):
|
||||
from .groot.processor_groot import make_groot_pre_post_processors_from_pretrained
|
||||
@@ -428,26 +402,6 @@ def make_pre_post_processors(
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
)
|
||||
|
||||
elif policy_cfg.type == "pi0_fast":
|
||||
from .pi0_fast.processor_pi0_fast import make_pi0_fast_pre_post_processors
|
||||
|
||||
processors = make_pi0_fast_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
dataset_repo_id=kwargs.get("dataset_repo_id"),
|
||||
)
|
||||
|
||||
elif policy_cfg.type == "pi052":
|
||||
# PI052 must precede PI05 because its config subclasses PI05Config.
|
||||
from .pi052.processor_pi052 import make_pi052_pre_post_processors
|
||||
|
||||
processors = make_pi052_pre_post_processors(
|
||||
config=policy_cfg,
|
||||
dataset_stats=kwargs.get("dataset_stats"),
|
||||
# Without a dataset repo, FAST auto-fit falls back to the universal tokenizer.
|
||||
dataset_repo_id=kwargs.get("dataset_repo_id"),
|
||||
)
|
||||
|
||||
elif isinstance(policy_cfg, PI05Config):
|
||||
from .pi05.processor_pi05 import make_pi05_pre_post_processors
|
||||
|
||||
@@ -623,9 +577,6 @@ def make_policy(
|
||||
raise ValueError("env_cfg cannot be None when ds_meta is not provided")
|
||||
features = env_to_policy_features(env_cfg)
|
||||
|
||||
if rename_map:
|
||||
features = {rename_map.get(key, key): feature for key, feature in features.items()}
|
||||
|
||||
cfg.output_features = {key: ft for key, ft in features.items() if ft.type is FeatureType.ACTION}
|
||||
if not cfg.input_features:
|
||||
cfg.input_features = {key: ft for key, ft in features.items() if key not in cfg.output_features}
|
||||
|
||||
@@ -15,7 +15,6 @@
|
||||
# limitations under the License.
|
||||
|
||||
import builtins
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
from collections import deque
|
||||
@@ -24,7 +23,6 @@ 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
|
||||
@@ -34,7 +32,6 @@ 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,
|
||||
@@ -50,7 +47,6 @@ else:
|
||||
_gated_residual = None
|
||||
layernorm_forward = None
|
||||
PaliGemmaForConditionalGenerationWithPiGemma = None
|
||||
cached_file = None
|
||||
from lerobot.configs import PreTrainedConfig
|
||||
from lerobot.utils.constants import (
|
||||
ACTION,
|
||||
@@ -70,84 +66,6 @@ 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:
|
||||
@@ -645,12 +563,6 @@ 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
|
||||
@@ -694,11 +606,7 @@ 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
|
||||
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.paligemma.model.vision_tower.gradient_checkpointing = True
|
||||
self.paligemma_with_expert.gemma_expert.model.gradient_checkpointing = True
|
||||
logging.info("Enabled gradient checkpointing for PI05Pytorch model")
|
||||
|
||||
@@ -706,11 +614,7 @@ 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
|
||||
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.paligemma.model.vision_tower.gradient_checkpointing = False
|
||||
self.paligemma_with_expert.gemma_expert.model.gradient_checkpointing = False
|
||||
logging.info("Disabled gradient checkpointing for PI05Pytorch model")
|
||||
|
||||
@@ -725,13 +629,10 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
)
|
||||
return func(*args, **kwargs)
|
||||
|
||||
def _prepare_attention_masks_4d(self, att_2d_masks, dtype=None):
|
||||
def _prepare_attention_masks_4d(self, att_2d_masks):
|
||||
"""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
|
||||
return torch.where(att_2d_masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE)
|
||||
|
||||
def sample_noise(self, shape, device):
|
||||
return torch.normal(
|
||||
@@ -757,16 +658,13 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
pad_masks = []
|
||||
att_masks = []
|
||||
|
||||
if self.checkpoint_vision_embeddings:
|
||||
# Process images
|
||||
for img, img_mask in zip(images, img_masks, strict=True):
|
||||
|
||||
def embed_image(img):
|
||||
return self._apply_checkpoint(self.paligemma_with_expert.embed_image, img)
|
||||
def image_embed_func(img):
|
||||
return self.paligemma_with_expert.embed_image(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):
|
||||
img_emb = self._apply_checkpoint(image_embed_func, img)
|
||||
bsize, num_img_embs = img_emb.shape[:2]
|
||||
|
||||
embs.append(img_emb)
|
||||
@@ -836,14 +734,8 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
|
||||
embs = torch.cat(embs, dim=1)
|
||||
pad_masks = torch.cat(pad_masks, dim=1)
|
||||
if self.use_on_device_suffix_mask:
|
||||
n = len(att_masks)
|
||||
att_masks = torch.zeros(n, dtype=embs.dtype, device=embs.device)
|
||||
att_masks[0] = 1
|
||||
att_masks = att_masks[None, :].expand(bsize, n)
|
||||
else:
|
||||
att_masks = torch.tensor(att_masks, dtype=embs.dtype, device=embs.device)
|
||||
att_masks = att_masks[None, :].expand(bsize, len(att_masks))
|
||||
att_masks = torch.tensor(att_masks, dtype=embs.dtype, device=embs.device)
|
||||
att_masks = att_masks[None, :].expand(bsize, len(att_masks))
|
||||
|
||||
return embs, pad_masks, att_masks, adarms_cond
|
||||
|
||||
@@ -927,8 +819,7 @@ 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
|
||||
|
||||
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)
|
||||
prefix_att_2d_masks_4d = self._prepare_attention_masks_4d(prefix_att_2d_masks)
|
||||
self.paligemma_with_expert.paligemma.model.language_model.config._attn_implementation = "eager" # noqa: SLF001
|
||||
|
||||
_, past_key_values = self.paligemma_with_expert.forward(
|
||||
@@ -941,19 +832,10 @@ class PI05Pytorch(nn.Module): # see openpi `PI0Pytorch`
|
||||
|
||||
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)
|
||||
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
||||
|
||||
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
|
||||
return self.denoise_step(
|
||||
@@ -1031,9 +913,6 @@ class PI05Policy(PreTrainedPolicy):
|
||||
|
||||
config_class = PI05Config
|
||||
name = "pi05"
|
||||
model_class = PI05Pytorch
|
||||
eval_after_pretrained_load = False
|
||||
show_openpi_disclaimer = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -1051,7 +930,7 @@ class PI05Policy(PreTrainedPolicy):
|
||||
|
||||
# Initialize the core PI05 model
|
||||
self.init_rtc_processor()
|
||||
self.model = self.model_class(config, rtc_processor=self.rtc_processor)
|
||||
self.model = PI05Pytorch(config, rtc_processor=self.rtc_processor)
|
||||
|
||||
# Enable gradient checkpointing if requested
|
||||
if config.gradient_checkpointing:
|
||||
@@ -1077,16 +956,16 @@ class PI05Policy(PreTrainedPolicy):
|
||||
strict: bool = True,
|
||||
**kwargs,
|
||||
) -> T:
|
||||
"""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"
|
||||
)
|
||||
"""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"
|
||||
)
|
||||
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,
|
||||
@@ -1100,34 +979,84 @@ class PI05Policy(PreTrainedPolicy):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Initialize model without loading weights
|
||||
# Check if dataset_stats were provided in kwargs
|
||||
model = cls(config, **kwargs)
|
||||
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
|
||||
# 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}")
|
||||
|
||||
return model
|
||||
|
||||
def _fix_pytorch_state_dict_keys(
|
||||
self, state_dict, model_config
|
||||
@@ -1299,16 +1228,12 @@ class PI05Policy(PreTrainedPolicy):
|
||||
|
||||
# Action queue logic for n_action_steps > 1
|
||||
if len(self._action_queue) == 0:
|
||||
action_batch = self._prepare_action_batch(batch)
|
||||
actions = self.predict_action_chunk(action_batch)[:, : self.config.n_action_steps]
|
||||
actions = self.predict_action_chunk(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."""
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,165 +0,0 @@
|
||||
# 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 = 128
|
||||
"""Number of low-vocab tokens the FAST tokenizer skips to avoid
|
||||
collisions with PaliGemma's text vocabulary."""
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
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."""
|
||||
|
||||
# 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.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.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")
|
||||
@@ -1,256 +0,0 @@
|
||||
# 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 logging
|
||||
import os
|
||||
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_local_leader() -> bool:
|
||||
return int(os.environ.get("LOCAL_RANK", "0")) == 0
|
||||
|
||||
|
||||
def _dataset_signature(
|
||||
dataset_repo_id: str,
|
||||
base_tokenizer_name: str,
|
||||
n_samples: int,
|
||||
chunk_size: int,
|
||||
) -> str:
|
||||
"""Deterministic short hash for naming the cache directory.
|
||||
|
||||
Keys on (dataset, base tokenizer, sample count, chunk size) so any
|
||||
of those changing re-runs the fit. ``chunk_size`` matters because
|
||||
the tokenizer is fit on chunks of that length.
|
||||
"""
|
||||
h = hashlib.sha256()
|
||||
h.update(dataset_repo_id.encode("utf-8"))
|
||||
h.update(b"\0")
|
||||
h.update(base_tokenizer_name.encode("utf-8"))
|
||||
h.update(b"\0")
|
||||
h.update(str(n_samples).encode("utf-8"))
|
||||
h.update(b"\0")
|
||||
h.update(str(chunk_size).encode("utf-8"))
|
||||
return h.hexdigest()[:16]
|
||||
|
||||
|
||||
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,
|
||||
) -> 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)
|
||||
sig = _dataset_signature(dataset_repo_id, base_tokenizer_name, n_samples, chunk_size)
|
||||
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)
|
||||
|
||||
# Each node fits its node-local cache once; its other local ranks wait.
|
||||
is_leader = _is_local_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,
|
||||
)
|
||||
|
||||
from transformers import AutoProcessor # noqa: PLC0415
|
||||
|
||||
# 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
|
||||
from huggingface_hub import snapshot_download # noqa: PLC0415
|
||||
|
||||
snap = Path(snapshot_download(repo_id=dataset_repo_id, repo_type="dataset"))
|
||||
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}.")
|
||||
|
||||
# Load only episode indices and fixed-width actions across all shards.
|
||||
tables = [_pq.read_table(f, columns=["episode_index", "action"]) 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}.")
|
||||
|
||||
# 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]
|
||||
|
||||
samples_per_episode = max(1, n_samples // max(num_episodes, 1))
|
||||
collected = 0
|
||||
eps_visited = 0
|
||||
short_episodes = 0
|
||||
ep_indices = list(ep_to_slice.keys())
|
||||
for ep_idx in rng.permutation(ep_indices):
|
||||
if collected >= n_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])
|
||||
collected += 1
|
||||
if collected >= n_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)
|
||||
logger.info(
|
||||
"FAST fit: collected %d chunks of shape %s from %d episodes",
|
||||
actions.shape[0],
|
||||
actions.shape[1:],
|
||||
eps_visited,
|
||||
)
|
||||
|
||||
# Match training-time quantile normalization so FAST sees the same bounded action space.
|
||||
flat = actions.reshape(-1, actions.shape[-1])
|
||||
q01 = np.quantile(flat, 0.01, axis=0)
|
||||
q99 = np.quantile(flat, 0.99, axis=0)
|
||||
span = np.where((q99 - q01) > 1e-6, q99 - q01, 1.0)
|
||||
actions = np.clip((actions - q01) / span * 2.0 - 1.0, -1.0, 1.0).astype(np.float32)
|
||||
|
||||
base = AutoProcessor.from_pretrained(base_tokenizer_name, trust_remote_code=True)
|
||||
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."
|
||||
)
|
||||
|
||||
fitted = base.fit(actions)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
fitted.save_pretrained(str(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) -> 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
|
||||
|
||||
return fit_fast_tokenizer(
|
||||
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,
|
||||
)
|
||||
@@ -1,263 +0,0 @@
|
||||
# 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
|
||||
@@ -1,19 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,205 +0,0 @@
|
||||
# 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.
|
||||
content = subtask
|
||||
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
|
||||
content = f"{subtask}, State: {discretize_state_str(state_row)};"
|
||||
|
||||
text_batch = _build_text_batch(
|
||||
self.policy,
|
||||
[{"role": "user", "content": content}],
|
||||
add_generation_prompt=False,
|
||||
)
|
||||
batch = dict(observation)
|
||||
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)
|
||||
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 <loc>
|
||||
)
|
||||
|
||||
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:
|
||||
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 ""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,149 +0,0 @@
|
||||
# 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,
|
||||
) -> 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:
|
||||
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),
|
||||
max_action_tokens=config.max_action_tokens,
|
||||
fast_skip_tokens=config.fast_skip_tokens,
|
||||
paligemma_tokenizer_name="google/paligemma-3b-pt-224",
|
||||
)
|
||||
)
|
||||
|
||||
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)
|
||||
@@ -1,483 +0,0 @@
|
||||
# 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"<say>{t}</say>" 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 ``<locDDDD>`` strings as single tokens.
|
||||
|
||||
Without registration, the stock tokenizer splits each location into generic text pieces.
|
||||
"""
|
||||
if "<loc0000>" in getattr(tokenizer, "added_tokens_encoder", {}):
|
||||
return tokenizer
|
||||
tokenizer.add_tokens([f"<loc{i:04d}>" for i in range(1024)])
|
||||
return tokenizer
|
||||
|
||||
|
||||
def _loc_token(coord: float, scale: float = _VQA_COORD_SCALE) -> str:
|
||||
"""PaliGemma ``<locNNNN>`` for a coord on a ``[0, scale]`` axis."""
|
||||
idx = round(float(coord) / scale * 1023) if scale > 0 else 0
|
||||
return f"<loc{max(0, min(1023, idx)):04d}>"
|
||||
|
||||
|
||||
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 ``<loc>`` 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
|
||||
|
||||
|
||||
@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
|
||||
@@ -61,9 +61,6 @@ 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
|
||||
|
||||
@@ -101,7 +101,6 @@ 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,
|
||||
) -> tuple[
|
||||
PolicyProcessorPipeline[dict[str, Any], dict[str, Any]],
|
||||
PolicyProcessorPipeline[PolicyAction, PolicyAction],
|
||||
@@ -144,10 +143,6 @@ 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)
|
||||
|
||||
input_steps: list[ProcessorStep] = [
|
||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||
AddBatchDimensionProcessorStep(),
|
||||
@@ -165,7 +160,7 @@ def make_pi0_fast_pre_post_processors(
|
||||
padding="max_length",
|
||||
),
|
||||
ActionTokenizerProcessorStep(
|
||||
action_tokenizer_name=action_tokenizer_path,
|
||||
action_tokenizer_name=config.action_tokenizer_name,
|
||||
max_action_tokens=config.max_action_tokens,
|
||||
fast_skip_tokens=config.fast_skip_tokens,
|
||||
paligemma_tokenizer_name=config.text_tokenizer_name,
|
||||
|
||||
@@ -14,27 +14,18 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
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,
|
||||
@@ -58,8 +49,6 @@ else:
|
||||
GradientCheckpointingLayer = None
|
||||
BaseModelOutputWithPast = None
|
||||
create_causal_mask = None
|
||||
CONFIG_MAPPING = None
|
||||
modeling_gemma = None
|
||||
|
||||
|
||||
def _gated_residual(
|
||||
@@ -132,10 +121,7 @@ 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)
|
||||
# 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:
|
||||
if len(x.shape) == 3:
|
||||
modulation = modulation.unsqueeze(1)
|
||||
scale, shift, gate = modulation.chunk(3, dim=-1)
|
||||
normed = normed * (1 + scale.float()) + shift.float()
|
||||
@@ -289,8 +275,6 @@ 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)
|
||||
@@ -383,374 +367,3 @@ __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
|
||||
|
||||
@@ -23,8 +23,6 @@ from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, TypedDict, TypeVar, Unpack
|
||||
|
||||
import packaging
|
||||
import safetensors
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download, save_torch_state_dict
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
@@ -34,6 +32,7 @@ from torch import Tensor, nn
|
||||
from lerobot.__version__ import __version__
|
||||
from lerobot.configs import PreTrainedConfig
|
||||
from lerobot.configs.train import TrainPipelineConfig
|
||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||
from lerobot.utils.hub import HubMixin
|
||||
|
||||
from .utils import log_model_loading_keys
|
||||
@@ -221,26 +220,10 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
||||
|
||||
@classmethod
|
||||
def _load_as_safetensor(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:
|
||||
# Create base kwargs
|
||||
kwargs = {"strict": strict}
|
||||
|
||||
# Add device parameter for newer versions that support it
|
||||
if packaging.version.parse(safetensors.__version__) >= packaging.version.parse("0.4.3"):
|
||||
kwargs["device"] = map_location
|
||||
|
||||
# Load the model with appropriate kwargs
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(model, model_file, **kwargs)
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(
|
||||
model, model_file, strict=strict, device=resolve_safetensors_device(map_location)
|
||||
)
|
||||
log_model_loading_keys(missing_keys, unexpected_keys)
|
||||
|
||||
# For older versions, manually move to device if needed
|
||||
if "device" not in kwargs and map_location != "cpu":
|
||||
logging.warning(
|
||||
"Loading model weights on other devices than 'cpu' is not supported natively in your version of safetensors."
|
||||
" This means that the model is loaded on 'cpu' first and then copied to the device."
|
||||
" This leads to a slower loading time."
|
||||
" Please update safetensors to version 0.4.3 or above for improved performance."
|
||||
)
|
||||
model.to(map_location)
|
||||
return model
|
||||
|
||||
@abc.abstractmethod
|
||||
|
||||
@@ -44,14 +44,6 @@ class SmolVLAConfig(PreTrainedConfig):
|
||||
# Image preprocessing
|
||||
resize_imgs_with_padding: tuple[int, int] = (512, 512)
|
||||
|
||||
# MEM short-horizon visual memory (https://arxiv.org/abs/2603.03596).
|
||||
# Past frames are fused inside the vision tower and dropped before the VLM,
|
||||
# keeping the number of prefix tokens identical to single-frame SmolVLA.
|
||||
use_visual_memory: bool = False
|
||||
visual_memory_frames: int = 6
|
||||
visual_memory_stride: int = 10
|
||||
visual_memory_temporal_attention_every: int = 4
|
||||
|
||||
# Add empty images. Used by smolvla_aloha_sim which adds the empty
|
||||
# left and right wrist cameras in addition to the top camera.
|
||||
empty_cameras: int = 0
|
||||
@@ -127,12 +119,6 @@ class SmolVLAConfig(PreTrainedConfig):
|
||||
raise NotImplementedError(
|
||||
"`use_delta_joint_actions_aloha` is used by smolvla for aloha real models. It is not ported yet in LeRobot."
|
||||
)
|
||||
if self.visual_memory_frames < 1:
|
||||
raise ValueError("visual_memory_frames must be at least 1")
|
||||
if self.visual_memory_stride < 1:
|
||||
raise ValueError("visual_memory_stride must be at least 1")
|
||||
if self.visual_memory_temporal_attention_every < 1:
|
||||
raise ValueError("visual_memory_temporal_attention_every must be at least 1")
|
||||
|
||||
def validate_features(self) -> None:
|
||||
for i in range(self.empty_cameras):
|
||||
@@ -162,10 +148,7 @@ class SmolVLAConfig(PreTrainedConfig):
|
||||
|
||||
@property
|
||||
def observation_delta_indices(self) -> list:
|
||||
if not self.use_visual_memory:
|
||||
return [0]
|
||||
horizon = (self.visual_memory_frames - 1) * self.visual_memory_stride
|
||||
return list(range(-horizon, 1, self.visual_memory_stride))
|
||||
return [0]
|
||||
|
||||
@property
|
||||
def action_delta_indices(self) -> list:
|
||||
|
||||
@@ -71,7 +71,6 @@ from ..utils import (
|
||||
)
|
||||
from .configuration_smolvla import SmolVLAConfig
|
||||
from .smolvlm_with_expert import SmolVLMWithExpertModel
|
||||
from .visual_memory import sample_visual_history
|
||||
|
||||
|
||||
class ActionSelectKwargs(TypedDict, total=False):
|
||||
@@ -254,10 +253,6 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
self._queues = {
|
||||
ACTION: deque(maxlen=self.config.n_action_steps),
|
||||
}
|
||||
if self.config.use_visual_memory:
|
||||
history_length = (self.config.visual_memory_frames - 1) * self.config.visual_memory_stride + 1
|
||||
self._queues.update({key: deque(maxlen=history_length) for key in self.config.image_features})
|
||||
self._visual_memory_steps_seen = 0
|
||||
|
||||
def init_rtc_processor(self):
|
||||
"""Initialize RTC processor if RTC is enabled in config."""
|
||||
@@ -286,15 +281,9 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
# In the case of offline inference, we have the action in the batch
|
||||
# that why without the k != ACTION check, it will raise an error because we are trying to stack
|
||||
# on an empty container.
|
||||
for k in list(batch):
|
||||
for k in batch:
|
||||
if k in self._queues and k != ACTION:
|
||||
history = list(self._queues[k])
|
||||
batch[k], batch[f"{k}_is_pad"] = sample_visual_history(
|
||||
history,
|
||||
num_frames=self.config.visual_memory_frames,
|
||||
stride=self.config.visual_memory_stride,
|
||||
steps_seen=self._visual_memory_steps_seen,
|
||||
)
|
||||
batch[k] = torch.stack(list(self._queues[k]), dim=1)
|
||||
|
||||
images, img_masks = self.prepare_images(batch)
|
||||
state = self.prepare_state(batch)
|
||||
@@ -320,11 +309,6 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
|
||||
return batch
|
||||
|
||||
def _populate_observation_queues(self, batch: dict[str, Tensor]) -> None:
|
||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||
if self.config.use_visual_memory:
|
||||
self._visual_memory_steps_seen += 1
|
||||
|
||||
@torch.no_grad()
|
||||
def predict_action_chunk(
|
||||
self, batch: dict[str, Tensor], noise: Tensor | None = None, **kwargs: Unpack[ActionSelectKwargs]
|
||||
@@ -332,7 +316,7 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
self.eval()
|
||||
|
||||
batch = self._prepare_batch(batch)
|
||||
self._populate_observation_queues(batch)
|
||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||
|
||||
actions = self._get_action_chunk(batch, noise, **kwargs)
|
||||
return actions
|
||||
@@ -354,7 +338,7 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
|
||||
self.eval()
|
||||
batch = self._prepare_batch(batch)
|
||||
self._populate_observation_queues(batch)
|
||||
self._queues = populate_queues(self._queues, batch, exclude_keys=[ACTION])
|
||||
|
||||
if self._check_get_actions_condition():
|
||||
actions = self._get_action_chunk(batch, noise)
|
||||
@@ -443,30 +427,19 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
||||
)
|
||||
# Preprocess image features present in the batch
|
||||
for key in present_img_keys:
|
||||
img = batch[key]
|
||||
if img.ndim == 5 and not self.config.use_visual_memory:
|
||||
img = img[:, -1, :, :, :]
|
||||
img = batch[key][:, -1, :, :, :] if batch[key].ndim == 5 else batch[key]
|
||||
if self.config.resize_imgs_with_padding is not None:
|
||||
if img.ndim == 5:
|
||||
batch_size, num_frames = img.shape[:2]
|
||||
img = resize_with_pad(
|
||||
img.flatten(0, 1), *self.config.resize_imgs_with_padding, pad_value=0
|
||||
).unflatten(0, (batch_size, num_frames))
|
||||
else:
|
||||
img = resize_with_pad(img, *self.config.resize_imgs_with_padding, pad_value=0)
|
||||
img = resize_with_pad(img, *self.config.resize_imgs_with_padding, pad_value=0)
|
||||
|
||||
# Normalize from range [0,1] to [-1,1] as expacted by siglip
|
||||
img = img * 2.0 - 1.0
|
||||
|
||||
bsize = img.shape[0]
|
||||
device = img.device
|
||||
if f"{key}_is_pad" in batch:
|
||||
mask = ~batch[f"{key}_is_pad"].bool()
|
||||
elif f"{key}_padding_mask" in batch:
|
||||
if f"{key}_padding_mask" in batch:
|
||||
mask = batch[f"{key}_padding_mask"].bool()
|
||||
else:
|
||||
mask_shape = img.shape[:2] if img.ndim == 5 else (bsize,)
|
||||
mask = torch.ones(mask_shape, dtype=torch.bool, device=device)
|
||||
mask = torch.ones(bsize, dtype=torch.bool, device=device)
|
||||
images.append(img)
|
||||
img_masks.append(mask)
|
||||
|
||||
@@ -689,11 +662,7 @@ class VLAFlowMatching(nn.Module):
|
||||
embs.append(image_start_token)
|
||||
pad_masks.append(image_start_mask)
|
||||
|
||||
img_emb = self.vlm_with_expert.embed_image(
|
||||
img,
|
||||
frame_mask=img_mask if img.ndim == 5 else None,
|
||||
temporal_attention_every=self.config.visual_memory_temporal_attention_every,
|
||||
)
|
||||
img_emb = self.vlm_with_expert.embed_image(img)
|
||||
img_emb = img_emb
|
||||
|
||||
# Normalize image embeddings
|
||||
@@ -701,8 +670,6 @@ class VLAFlowMatching(nn.Module):
|
||||
img_emb = img_emb * torch.tensor(img_emb_dim**0.5, dtype=img_emb.dtype, device=img_emb.device)
|
||||
|
||||
bsize, num_img_embs = img_emb.shape[:2]
|
||||
if img_mask.ndim == 2:
|
||||
img_mask = img_mask[:, -1]
|
||||
img_mask = img_mask[:, None].expand(bsize, num_img_embs)
|
||||
|
||||
embs.append(img_emb)
|
||||
|
||||
@@ -20,8 +20,6 @@ from torch import nn
|
||||
|
||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||
|
||||
from .visual_memory import encode_video_with_mem
|
||||
|
||||
if TYPE_CHECKING or _transformers_available:
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
@@ -190,31 +188,17 @@ class SmolVLMWithExpertModel(nn.Module):
|
||||
if self.train_expert_only:
|
||||
self.vlm.eval()
|
||||
|
||||
def embed_image(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
*,
|
||||
frame_mask: torch.Tensor | None = None,
|
||||
temporal_attention_every: int = 4,
|
||||
):
|
||||
def embed_image(self, image: torch.Tensor):
|
||||
patch_attention_mask = None
|
||||
# Get sequence from the vision encoder
|
||||
vision_model = self.get_vlm_model().vision_model
|
||||
image = image.to(dtype=vision_model.dtype)
|
||||
if image.ndim == 5:
|
||||
if frame_mask is None:
|
||||
frame_mask = torch.ones(image.shape[:2], dtype=torch.bool, device=image.device)
|
||||
image_hidden_states = encode_video_with_mem(
|
||||
vision_model,
|
||||
image,
|
||||
frame_mask,
|
||||
temporal_attention_every=temporal_attention_every,
|
||||
)
|
||||
else:
|
||||
image_hidden_states = vision_model(
|
||||
pixel_values=image,
|
||||
image_hidden_states = (
|
||||
self.get_vlm_model()
|
||||
.vision_model(
|
||||
pixel_values=image.to(dtype=self.get_vlm_model().vision_model.dtype),
|
||||
patch_attention_mask=patch_attention_mask,
|
||||
).last_hidden_state
|
||||
)
|
||||
.last_hidden_state
|
||||
)
|
||||
# Modality projection & resampling
|
||||
image_hidden_states = self.get_vlm_model().connector(image_hidden_states)
|
||||
return image_hidden_states
|
||||
|
||||
@@ -1,161 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Short-horizon visual memory from MEM (arXiv:2603.03596).
|
||||
|
||||
The encoder keeps SigLIP's pretrained parameters and alternates its ordinary
|
||||
per-frame spatial attention with causal attention across time for matching
|
||||
spatial patches. Only the current frame's patch tokens leave the vision tower,
|
||||
so enabling memory does not increase the VLM prefix length.
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
def sample_visual_history(
|
||||
history: list[Tensor], *, num_frames: int, stride: int, steps_seen: int
|
||||
) -> tuple[Tensor, Tensor]:
|
||||
"""Subsample an inference queue and mark pre-episode padding frames."""
|
||||
sampled = history[::stride][-num_frames:]
|
||||
if len(sampled) != num_frames:
|
||||
raise ValueError(f"visual history has {len(sampled)} samples, expected {num_frames}")
|
||||
video = torch.stack(sampled, dim=1)
|
||||
required_ages = range((num_frames - 1) * stride, -1, -stride)
|
||||
valid = torch.tensor([steps_seen > age for age in required_ages], dtype=torch.bool, device=video.device)
|
||||
padding_mask = (~valid)[None, :].expand(video.shape[0], -1)
|
||||
return video, padding_mask
|
||||
|
||||
|
||||
def temporal_sinusoidal_embedding(
|
||||
num_frames: int, hidden_size: int, *, device: torch.device, dtype: torch.dtype
|
||||
) -> Tensor:
|
||||
"""Return fixed temporal embeddings with an exactly-zero current position."""
|
||||
if hidden_size % 2:
|
||||
raise ValueError(f"hidden_size must be even, got {hidden_size}")
|
||||
|
||||
# History is ordered oldest -> current, matching t in [-K, 0] in MEM.
|
||||
positions = torch.arange(1 - num_frames, 1, device=device, dtype=torch.float32)[:, None]
|
||||
frequencies = torch.exp(
|
||||
torch.arange(0, hidden_size, 2, device=device, dtype=torch.float32)
|
||||
* (-math.log(10_000.0) / hidden_size)
|
||||
)[None, :]
|
||||
angles = positions * frequencies
|
||||
embedding = torch.zeros(num_frames, hidden_size, device=device, dtype=torch.float32)
|
||||
embedding[:, 0::2] = torch.sin(angles)
|
||||
embedding[:, 1::2] = torch.cos(angles) - 1.0
|
||||
return embedding.to(dtype=dtype)
|
||||
|
||||
|
||||
def causal_temporal_mask(frame_mask: Tensor, *, dtype: torch.dtype, num_patches: int) -> Tensor:
|
||||
"""Build an additive causal/key-padding mask for per-patch temporal attention."""
|
||||
if frame_mask.ndim != 2:
|
||||
raise ValueError(f"frame_mask must have shape (batch, frames), got {tuple(frame_mask.shape)}")
|
||||
|
||||
batch_size, num_frames = frame_mask.shape
|
||||
allowed = torch.ones(num_frames, num_frames, dtype=torch.bool, device=frame_mask.device).tril()
|
||||
allowed = allowed[None, :, :] & frame_mask[:, None, :].bool()
|
||||
# The current frame is always present in normal use. Keeping the diagonal
|
||||
# valid also prevents NaNs for fully padded historical query rows.
|
||||
diagonal = torch.eye(num_frames, dtype=torch.bool, device=frame_mask.device)[None, :, :]
|
||||
allowed = allowed | diagonal
|
||||
mask = torch.zeros(batch_size, 1, num_frames, num_frames, dtype=dtype, device=frame_mask.device)
|
||||
mask.masked_fill_(~allowed[:, None, :, :], torch.finfo(dtype).min)
|
||||
return mask.repeat_interleave(num_patches, dim=0)
|
||||
|
||||
|
||||
def encode_video_with_mem(
|
||||
vision_model,
|
||||
pixel_values: Tensor,
|
||||
frame_mask: Tensor,
|
||||
*,
|
||||
temporal_attention_every: int,
|
||||
) -> Tensor:
|
||||
"""Encode ``(B,T,C,H,W)`` video with MEM's space-time separable attention.
|
||||
|
||||
No modules or learnable parameters are added. At temporal layers, the same
|
||||
layer-normalization and Q/K/V/out projections are reused for a second,
|
||||
causal attention operation along time. For a one-frame input the temporal
|
||||
branch is skipped, making this exactly equivalent to the original image
|
||||
encoder.
|
||||
"""
|
||||
if pixel_values.ndim != 5:
|
||||
raise ValueError(f"pixel_values must have shape (B,T,C,H,W), got {tuple(pixel_values.shape)}")
|
||||
if temporal_attention_every < 1:
|
||||
raise ValueError("temporal_attention_every must be at least 1")
|
||||
|
||||
batch_size, num_frames, channels, height, width = pixel_values.shape
|
||||
if frame_mask.shape != (batch_size, num_frames):
|
||||
raise ValueError(
|
||||
f"frame_mask must have shape {(batch_size, num_frames)}, got {tuple(frame_mask.shape)}"
|
||||
)
|
||||
|
||||
required = ("embeddings", "encoder", "post_layernorm")
|
||||
if any(not hasattr(vision_model, name) for name in required):
|
||||
raise TypeError("MEM visual memory currently requires a SigLIP-compatible vision tower")
|
||||
|
||||
flat_pixels = pixel_values.reshape(batch_size * num_frames, channels, height, width)
|
||||
if hasattr(vision_model, "patch_size"):
|
||||
patch_size = vision_model.patch_size
|
||||
patch_attention_mask = torch.ones(
|
||||
batch_size * num_frames,
|
||||
height // patch_size,
|
||||
width // patch_size,
|
||||
dtype=torch.bool,
|
||||
device=pixel_values.device,
|
||||
)
|
||||
hidden_states = vision_model.embeddings(
|
||||
pixel_values=flat_pixels,
|
||||
patch_attention_mask=patch_attention_mask,
|
||||
)
|
||||
else:
|
||||
hidden_states = vision_model.embeddings(flat_pixels)
|
||||
num_patches, hidden_size = hidden_states.shape[1:]
|
||||
hidden_states = hidden_states.reshape(batch_size, num_frames, num_patches, hidden_size)
|
||||
|
||||
temporal_positions = temporal_sinusoidal_embedding(
|
||||
num_frames, hidden_size, device=hidden_states.device, dtype=hidden_states.dtype
|
||||
)[None, :, None, :]
|
||||
temporal_mask = causal_temporal_mask(frame_mask, dtype=hidden_states.dtype, num_patches=num_patches)
|
||||
|
||||
for layer_index, layer in enumerate(vision_model.encoder.layers):
|
||||
spatial_input = hidden_states.reshape(batch_size * num_frames, num_patches, hidden_size)
|
||||
if num_frames == 1 or (layer_index + 1) % temporal_attention_every:
|
||||
hidden_states = layer(spatial_input, attention_mask=None).reshape(
|
||||
batch_size, num_frames, num_patches, hidden_size
|
||||
)
|
||||
continue
|
||||
|
||||
residual = hidden_states
|
||||
spatial_norm = layer.layer_norm1(spatial_input)
|
||||
spatial_output, _ = layer.self_attn(hidden_states=spatial_norm, attention_mask=None)
|
||||
spatial_output = spatial_output.reshape(batch_size, num_frames, num_patches, hidden_size)
|
||||
|
||||
temporal_input = (hidden_states + temporal_positions).permute(0, 2, 1, 3)
|
||||
temporal_input = temporal_input.reshape(batch_size * num_patches, num_frames, hidden_size)
|
||||
temporal_norm = layer.layer_norm1(temporal_input)
|
||||
temporal_output, _ = layer.self_attn(
|
||||
hidden_states=temporal_norm,
|
||||
attention_mask=temporal_mask,
|
||||
)
|
||||
temporal_output = temporal_output.reshape(batch_size, num_patches, num_frames, hidden_size)
|
||||
temporal_output = temporal_output.permute(0, 2, 1, 3)
|
||||
|
||||
hidden_states = residual + spatial_output + temporal_output
|
||||
hidden_states = hidden_states + layer.mlp(layer.layer_norm2(hidden_states))
|
||||
|
||||
# MEM drops all historical tokens before the VLA backbone.
|
||||
return vision_model.post_layernorm(hidden_states[:, -1])
|
||||
@@ -154,7 +154,7 @@ class XVLAModel(nn.Module):
|
||||
# Freeze or unfreeze policy transformer
|
||||
if not self.config.train_policy_transformer:
|
||||
for name, param in self.transformer.named_parameters():
|
||||
if "soft_prompts" not in name:
|
||||
if "soft_prompt" not in name:
|
||||
param.requires_grad = False
|
||||
|
||||
# Freeze or unfreeze soft prompts
|
||||
|
||||
@@ -175,6 +175,9 @@ class AddBatchDimensionComplementaryDataStep(ComplementaryDataProcessorStep):
|
||||
if isinstance(task_index_value, Tensor) and task_index_value.dim() == 0:
|
||||
complementary_data["task_index"] = task_index_value.unsqueeze(0)
|
||||
|
||||
complementary_data.pop("language_persistent", None)
|
||||
complementary_data.pop("language_events", None)
|
||||
|
||||
if "messages" in complementary_data:
|
||||
messages = complementary_data["messages"]
|
||||
if isinstance(messages, list) and (not messages or isinstance(messages[0], dict)):
|
||||
|
||||
@@ -21,23 +21,6 @@ from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
|
||||
from .pipeline import ObservationProcessorStep, ProcessorStepRegistry
|
||||
|
||||
_AUXILIARY_KEY_SUFFIXES = ("_is_pad", "_padding_mask")
|
||||
|
||||
|
||||
def rename_transition_key(key: str, rename_map: dict[str, str]) -> str:
|
||||
"""Rename a feature key and any sampling metadata derived from it."""
|
||||
if key in rename_map:
|
||||
return rename_map[key]
|
||||
for suffix in _AUXILIARY_KEY_SUFFIXES:
|
||||
if key.endswith(suffix) and key[: -len(suffix)] in rename_map:
|
||||
return f"{rename_map[key[: -len(suffix)]]}{suffix}"
|
||||
return key
|
||||
|
||||
|
||||
def rename_transition_keys(data: dict[str, Any], rename_map: dict[str, str]) -> dict[str, Any]:
|
||||
"""Rename all transition keys, including delta-sampling padding masks."""
|
||||
return {rename_transition_key(key, rename_map): value for key, value in data.items()}
|
||||
|
||||
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="rename_observations_processor")
|
||||
@@ -58,7 +41,14 @@ class RenameObservationsProcessorStep(ObservationProcessorStep):
|
||||
rename_map: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
def observation(self, observation):
|
||||
return rename_transition_keys(observation, self.rename_map)
|
||||
processed_obs = {}
|
||||
for key, value in observation.items():
|
||||
if key in self.rename_map:
|
||||
processed_obs[self.rename_map[key]] = value
|
||||
else:
|
||||
processed_obs[key] = value
|
||||
|
||||
return processed_obs
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"rename_map": self.rename_map}
|
||||
|
||||
@@ -16,7 +16,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||
@@ -32,18 +32,17 @@ from .pipeline import ProcessorStep, ProcessorStepRegistry
|
||||
@dataclass
|
||||
@ProcessorStepRegistry.register(name="render_messages_processor")
|
||||
class RenderMessagesStep(ProcessorStep):
|
||||
"""Render language columns into recipe-defined messages and supervision metadata."""
|
||||
"""Processor step that turns raw language columns into rendered chat messages.
|
||||
|
||||
Reads ``language_persistent`` and ``language_events`` from the transition's
|
||||
complementary data, renders them through ``recipe`` at the sample timestamp,
|
||||
and replaces the raw columns with the resulting ``messages`` /
|
||||
``message_streams`` / ``target_message_indices`` keys.
|
||||
"""
|
||||
|
||||
recipe: TrainingRecipe
|
||||
dataset_ctx: Any | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if isinstance(self.recipe, dict):
|
||||
self.recipe = TrainingRecipe.from_dict(self.recipe)
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
return {"recipe": asdict(self.recipe)}
|
||||
|
||||
def __call__(self, transition: EnvTransition) -> EnvTransition | None:
|
||||
"""Render messages for a single transition; return ``None`` to drop it."""
|
||||
complementary_data = transition.get(TransitionKey.COMPLEMENTARY_DATA) or {}
|
||||
@@ -51,17 +50,7 @@ class RenderMessagesStep(ProcessorStep):
|
||||
events = complementary_data.get(LANGUAGE_EVENTS) or []
|
||||
|
||||
if not persistent and not events:
|
||||
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||
if rendered is None:
|
||||
return transition
|
||||
new_transition = transition.copy()
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data.update(rendered)
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
if _is_batched_language(persistent) or _is_batched_language(events):
|
||||
return self._call_batch(transition, complementary_data, persistent, events)
|
||||
return transition
|
||||
|
||||
timestamp = complementary_data.get("timestamp")
|
||||
if timestamp is None:
|
||||
@@ -78,147 +67,18 @@ class RenderMessagesStep(ProcessorStep):
|
||||
dataset_ctx=self.dataset_ctx,
|
||||
)
|
||||
if rendered is None:
|
||||
rendered = _fallback_low_level_render(complementary_data.get("task"))
|
||||
if rendered is None:
|
||||
return None
|
||||
return None
|
||||
|
||||
new_transition = transition.copy()
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data = dict(complementary_data)
|
||||
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||
new_complementary_data.update(rendered)
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
def _call_batch(
|
||||
self,
|
||||
transition: EnvTransition,
|
||||
complementary_data: dict[str, Any],
|
||||
persistent_batch: list,
|
||||
events_batch: list,
|
||||
) -> EnvTransition | None:
|
||||
timestamp = complementary_data.get("timestamp")
|
||||
if timestamp is None:
|
||||
raise KeyError("RenderMessagesStep requires sample timestamp in complementary data.")
|
||||
|
||||
batch_size = max(len(persistent_batch), len(events_batch))
|
||||
messages: list[list[dict[str, Any]]] = []
|
||||
message_streams: list[list[str | None]] = []
|
||||
target_message_indices: list[list[int]] = []
|
||||
keep_indices: list[int] = []
|
||||
|
||||
for i in range(batch_size):
|
||||
rendered = render_sample(
|
||||
recipe=self.recipe,
|
||||
persistent=persistent_batch[i] if i < len(persistent_batch) else [],
|
||||
events=events_batch[i] if i < len(events_batch) else [],
|
||||
t=_batch_value(timestamp, i),
|
||||
sample_idx=int(_batch_value(complementary_data.get("index", 0), i)),
|
||||
task=_batch_value(complementary_data.get("task"), i),
|
||||
dataset_ctx=self.dataset_ctx,
|
||||
)
|
||||
if rendered is None:
|
||||
rendered = _fallback_low_level_render(_batch_value(complementary_data.get("task"), i))
|
||||
if rendered is None:
|
||||
continue
|
||||
keep_indices.append(i)
|
||||
messages.append(rendered["messages"])
|
||||
message_streams.append(rendered["message_streams"])
|
||||
target_message_indices.append(rendered["target_message_indices"])
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
new_transition = (
|
||||
_select_batch_indices(transition, keep_indices)
|
||||
if len(keep_indices) != batch_size
|
||||
else transition.copy()
|
||||
)
|
||||
new_complementary_data = dict(new_transition.get(TransitionKey.COMPLEMENTARY_DATA) or {})
|
||||
new_complementary_data.pop(LANGUAGE_PERSISTENT, None)
|
||||
new_complementary_data.pop(LANGUAGE_EVENTS, None)
|
||||
new_complementary_data["messages"] = messages
|
||||
new_complementary_data["message_streams"] = message_streams
|
||||
new_complementary_data["target_message_indices"] = target_message_indices
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = new_complementary_data
|
||||
return new_transition
|
||||
|
||||
def transform_features(
|
||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||
"""Pass features through unchanged; rendering only touches complementary data."""
|
||||
return features
|
||||
|
||||
|
||||
def _scalar(value: Any) -> float | int:
|
||||
"""Unwrap a tensor/array/single-element list into a Python scalar."""
|
||||
if hasattr(value, "item"):
|
||||
return value.item()
|
||||
if isinstance(value, list):
|
||||
if len(value) != 1:
|
||||
raise ValueError(f"Expected a scalar, got list of length {len(value)}: {value!r}")
|
||||
return _scalar(value[0])
|
||||
return value
|
||||
|
||||
|
||||
def _is_batched_language(value: Any) -> bool:
|
||||
return isinstance(value, list) and bool(value) and isinstance(value[0], list)
|
||||
|
||||
|
||||
def _batch_value(value: Any, index: int) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, list):
|
||||
return value[index]
|
||||
if hasattr(value, "ndim") and value.ndim > 0:
|
||||
return _scalar(value[index])
|
||||
return _scalar(value)
|
||||
|
||||
|
||||
def _select_batch_indices(transition: EnvTransition, indices: list[int]) -> EnvTransition:
|
||||
selected = transition.copy()
|
||||
for key in (TransitionKey.OBSERVATION, TransitionKey.COMPLEMENTARY_DATA):
|
||||
data = selected.get(key)
|
||||
if isinstance(data, dict):
|
||||
selected[key] = {k: _select_value(v, indices) for k, v in data.items()}
|
||||
action = selected.get(TransitionKey.ACTION)
|
||||
if action is not None:
|
||||
selected[TransitionKey.ACTION] = _select_value(action, indices)
|
||||
return selected
|
||||
|
||||
|
||||
def _select_value(value: Any, indices: list[int]) -> Any:
|
||||
if isinstance(value, list) and len(value) >= len(indices):
|
||||
return [value[i] for i in indices]
|
||||
if hasattr(value, "index_select") and hasattr(value, "new_tensor") and getattr(value, "ndim", 0) > 0:
|
||||
return value.index_select(0, value.new_tensor(indices).long())
|
||||
return value
|
||||
|
||||
|
||||
def _fallback_low_level_render(task: Any) -> dict[str, Any] | None:
|
||||
"""Keep action-only samples trainable when no recipe branch matches."""
|
||||
if hasattr(task, "item"):
|
||||
task = task.item()
|
||||
if isinstance(task, list):
|
||||
messages = []
|
||||
message_streams = []
|
||||
target_message_indices = []
|
||||
for t in task:
|
||||
rendered = _fallback_low_level_render(t)
|
||||
if rendered is None:
|
||||
return None
|
||||
messages.append(rendered["messages"])
|
||||
message_streams.append(rendered["message_streams"])
|
||||
target_message_indices.append(rendered["target_message_indices"])
|
||||
return {
|
||||
"messages": messages,
|
||||
"message_streams": message_streams,
|
||||
"target_message_indices": target_message_indices,
|
||||
}
|
||||
if not isinstance(task, str) or not task:
|
||||
return None
|
||||
return {
|
||||
"messages": [{"role": "user", "content": task}],
|
||||
"message_streams": ["low_level"],
|
||||
"target_message_indices": [],
|
||||
}
|
||||
|
||||
@@ -32,7 +32,6 @@ import torch
|
||||
from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.types import EnvTransition, RobotObservation, TransitionKey
|
||||
from lerobot.utils.constants import (
|
||||
ACTION_CODE_TOKEN_MASK,
|
||||
ACTION_TOKEN_MASK,
|
||||
ACTION_TOKENS,
|
||||
OBS_LANGUAGE_ATTENTION_MASK,
|
||||
@@ -413,15 +412,14 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
# During inference, no action is available, skip tokenization
|
||||
return new_transition
|
||||
|
||||
# Tokenize and get masks for the full formatted sequence and the discrete action codes.
|
||||
tokens, mask, code_mask = self._tokenize_action(action)
|
||||
# Tokenize and get both tokens and mask
|
||||
tokens, mask = self._tokenize_action(action)
|
||||
|
||||
# Store mask in complementary data
|
||||
complementary_data = new_transition.get(TransitionKey.COMPLEMENTARY_DATA, {})
|
||||
if complementary_data is None:
|
||||
complementary_data = {}
|
||||
complementary_data[ACTION_TOKEN_MASK] = mask
|
||||
complementary_data[ACTION_CODE_TOKEN_MASK] = code_mask
|
||||
complementary_data[ACTION_TOKENS] = tokens
|
||||
new_transition[TransitionKey.COMPLEMENTARY_DATA] = complementary_data
|
||||
return new_transition
|
||||
@@ -432,7 +430,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
"""
|
||||
return self._paligemma_tokenizer.vocab_size - 1 - self.fast_skip_tokens - tokens
|
||||
|
||||
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Tokenizes the action tensor and creates a mask.
|
||||
|
||||
@@ -461,7 +459,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
# The fast tokenizer expects action data and returns token IDs
|
||||
tokens_list = []
|
||||
masks_list = []
|
||||
code_masks_list = []
|
||||
|
||||
for i in range(batch_size):
|
||||
# Tokenize single action (move to CPU first as tokenizer uses scipy which requires numpy)
|
||||
@@ -479,26 +476,19 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
if tokens.dim() > 1:
|
||||
tokens = tokens.flatten()
|
||||
|
||||
action_code_tokens = self._act_tokens_to_paligemma_tokens(tokens)
|
||||
bos_id = self._paligemma_tokenizer.bos_token_id
|
||||
prompt_tokens = torch.tensor(
|
||||
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
||||
device=action.device,
|
||||
)
|
||||
end_tokens = torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device)
|
||||
|
||||
code_start = 1 + len(prompt_tokens)
|
||||
code_end = code_start + len(action_code_tokens)
|
||||
# add bos
|
||||
tokens = torch.cat(
|
||||
[
|
||||
torch.tensor([bos_id], device=action.device),
|
||||
prompt_tokens,
|
||||
action_code_tokens,
|
||||
end_tokens,
|
||||
torch.tensor(
|
||||
self._paligemma_tokenizer.encode("Action: ", add_special_tokens=False),
|
||||
device=action.device,
|
||||
),
|
||||
self._act_tokens_to_paligemma_tokens(tokens),
|
||||
torch.tensor(self._paligemma_tokenizer.encode("|"), device=action.device),
|
||||
]
|
||||
)
|
||||
code_mask = torch.zeros(len(tokens), dtype=torch.bool, device=action.device)
|
||||
code_mask[code_start:code_end] = True
|
||||
|
||||
# Truncate or pad to max_action_tokens
|
||||
if len(tokens) > self.max_action_tokens:
|
||||
@@ -507,49 +497,44 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
"Consider increasing the `max_action_tokens` in your model config if this happens frequently."
|
||||
)
|
||||
tokens = tokens[: self.max_action_tokens]
|
||||
code_mask = code_mask[: self.max_action_tokens]
|
||||
mask = torch.ones(self.max_action_tokens, dtype=torch.bool, device=action.device)
|
||||
else:
|
||||
pad_len = self.max_action_tokens - len(tokens)
|
||||
mask = torch.cat(
|
||||
[
|
||||
torch.ones(len(tokens), dtype=torch.bool, device=action.device),
|
||||
torch.zeros(pad_len, dtype=torch.bool, device=action.device),
|
||||
torch.zeros(
|
||||
self.max_action_tokens - len(tokens), dtype=torch.bool, device=action.device
|
||||
),
|
||||
]
|
||||
)
|
||||
code_mask = torch.nn.functional.pad(code_mask, (0, pad_len), value=False)
|
||||
# Pad tokens with zeros
|
||||
tokens = torch.nn.functional.pad(tokens, (0, pad_len), value=0)
|
||||
tokens = torch.nn.functional.pad(tokens, (0, self.max_action_tokens - len(tokens)), value=0)
|
||||
|
||||
tokens_list.append(tokens)
|
||||
masks_list.append(mask)
|
||||
code_masks_list.append(code_mask)
|
||||
|
||||
# Stack into batched tensors
|
||||
tokens_batch = torch.stack(tokens_list, dim=0) # (B, max_action_tokens)
|
||||
masks_batch = torch.stack(masks_list, dim=0) # (B, max_action_tokens)
|
||||
code_masks_batch = torch.stack(code_masks_list, dim=0) # (B, max_action_tokens)
|
||||
|
||||
# Remove batch dimension if input was single sample
|
||||
if single_sample:
|
||||
tokens_batch = tokens_batch.squeeze(0)
|
||||
masks_batch = masks_batch.squeeze(0)
|
||||
code_masks_batch = code_masks_batch.squeeze(0)
|
||||
|
||||
# Move to the same device as the input
|
||||
if device is not None:
|
||||
tokens_batch = tokens_batch.to(device)
|
||||
masks_batch = masks_batch.to(device)
|
||||
code_masks_batch = code_masks_batch.to(device)
|
||||
|
||||
return tokens_batch, masks_batch, code_masks_batch
|
||||
return tokens_batch, masks_batch
|
||||
|
||||
def action(self, action: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
This method is not used since we override __call__.
|
||||
Required by ActionProcessorStep ABC.
|
||||
"""
|
||||
tokens, _, _ = self._tokenize_action(action)
|
||||
tokens, _ = self._tokenize_action(action)
|
||||
return tokens
|
||||
|
||||
def get_config(self) -> dict[str, Any]:
|
||||
@@ -565,8 +550,6 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
|
||||
config = {
|
||||
"trust_remote_code": self.trust_remote_code,
|
||||
"max_action_tokens": self.max_action_tokens,
|
||||
"fast_skip_tokens": self.fast_skip_tokens,
|
||||
"paligemma_tokenizer_name": self.paligemma_tokenizer_name,
|
||||
}
|
||||
|
||||
# Only save tokenizer_name if it was used to create the tokenizer
|
||||
|
||||
@@ -21,8 +21,6 @@ from pathlib import Path
|
||||
from tempfile import TemporaryDirectory
|
||||
from typing import TYPE_CHECKING, Any, TypeVar
|
||||
|
||||
import packaging
|
||||
import safetensors
|
||||
from huggingface_hub import HfApi, ModelCard, ModelCardData, hf_hub_download
|
||||
from huggingface_hub.constants import SAFETENSORS_SINGLE_FILE
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
@@ -30,6 +28,7 @@ from safetensors.torch import load_model as load_model_as_safetensor, save_model
|
||||
from torch import Tensor, nn
|
||||
|
||||
from lerobot.configs.rewards import RewardModelConfig
|
||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||
from lerobot.utils.hub import HubMixin
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -129,29 +128,13 @@ class PreTrainedRewardModel(nn.Module, HubMixin, abc.ABC):
|
||||
|
||||
@classmethod
|
||||
def _load_as_safetensor(cls, model: T, model_file: str, map_location: str, strict: bool) -> T:
|
||||
# Create base kwargs
|
||||
kwargs = {"strict": strict}
|
||||
|
||||
# Add device parameter for newer versions that support it
|
||||
if packaging.version.parse(safetensors.__version__) >= packaging.version.parse("0.4.3"):
|
||||
kwargs["device"] = map_location
|
||||
|
||||
# Load the model with appropriate kwargs
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(model, model_file, **kwargs)
|
||||
missing_keys, unexpected_keys = load_model_as_safetensor(
|
||||
model, model_file, strict=strict, device=resolve_safetensors_device(map_location)
|
||||
)
|
||||
if missing_keys:
|
||||
logging.warning(f"Missing key(s) when loading model: {missing_keys}")
|
||||
if unexpected_keys:
|
||||
logging.warning(f"Unexpected key(s) when loading model: {unexpected_keys}")
|
||||
|
||||
# For older versions, manually move to device if needed
|
||||
if "device" not in kwargs and map_location != "cpu":
|
||||
logging.warning(
|
||||
"Loading model weights on other devices than 'cpu' is not supported natively in your version of safetensors."
|
||||
" This means that the model is loaded on 'cpu' first and then copied to the device."
|
||||
" This leads to a slower loading time."
|
||||
" Please update safetensors to version 0.4.3 or above for improved performance."
|
||||
)
|
||||
model.to(map_location)
|
||||
return model
|
||||
|
||||
def get_optim_params(self):
|
||||
|
||||
@@ -21,8 +21,6 @@ from lerobot.utils.import_utils import make_device_from_device_class
|
||||
from .config import RobotConfig
|
||||
from .robot import Robot
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def make_robot_from_config(config: RobotConfig) -> Robot:
|
||||
# TODO(Steven): Consider just using the make_device_from_device_class for all types
|
||||
@@ -120,7 +118,7 @@ def ensure_safe_goal_position(
|
||||
}
|
||||
|
||||
if warnings_dict:
|
||||
logger.warning(
|
||||
logging.warning(
|
||||
"Relative goal position magnitude had to be clamped to be safe.\n"
|
||||
f"{pformat(warnings_dict, indent=4)}"
|
||||
)
|
||||
|
||||
@@ -1,38 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Policy-agnostic runtime for language-conditioned policies.
|
||||
|
||||
Adapters registered in :mod:`lerobot.runtime.registry` are served by ``lerobot-rollout --language``.
|
||||
"""
|
||||
|
||||
from .adapter import BaseLanguageAdapter, GenerationConfig, LanguageDiagnostics
|
||||
from .language_runtime import (
|
||||
LanguageConditionedPolicyAdapter,
|
||||
LanguageConditionedRuntime,
|
||||
RuntimeState,
|
||||
Tick,
|
||||
TickClock,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"BaseLanguageAdapter",
|
||||
"GenerationConfig",
|
||||
"LanguageConditionedPolicyAdapter",
|
||||
"LanguageConditionedRuntime",
|
||||
"LanguageDiagnostics",
|
||||
"RuntimeState",
|
||||
"Tick",
|
||||
"TickClock",
|
||||
]
|
||||
@@ -1,165 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Policy adapters for the language runtime.
|
||||
|
||||
The base adapter owns generation control and diagnostics while subclasses provide policy-specific actions and text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from .language_runtime import RuntimeState
|
||||
|
||||
_SAY_RE = re.compile(r"<\s*say\s*>(.*?)<\s*/\s*say\s*>", re.IGNORECASE | re.DOTALL)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationConfig:
|
||||
"""Text-generation settings fixed for the adapter's lifetime."""
|
||||
|
||||
min_new_tokens: int = 0
|
||||
temperature: float = 0.0
|
||||
top_p: float = 1.0
|
||||
chunks_per_regen: int = 1 # regenerate the language context every N action chunks
|
||||
enable_memory: bool = True # generate a running memory note on subtask change
|
||||
enable_subtask: bool = True # generate the low-level subtask (off => use the given text directly)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LanguageDiagnostics:
|
||||
"""Runtime-panel generation counters keyed by text kind."""
|
||||
|
||||
last_raw: dict[str, str] = field(default_factory=dict)
|
||||
empty: dict[str, int] = field(default_factory=dict)
|
||||
repeat: int = 0
|
||||
|
||||
def _bump(self, table: dict[str, int], kind: str) -> int:
|
||||
table[kind] = table.get(kind, 0) + 1
|
||||
return table[kind]
|
||||
|
||||
|
||||
class BaseLanguageAdapter(ABC):
|
||||
"""Batteries-included adapter: generic high-level control, policy primitives abstract."""
|
||||
|
||||
def __init__(self, policy: Any, gen: GenerationConfig | None = None) -> None:
|
||||
self.policy = policy
|
||||
self.gen = gen or GenerationConfig()
|
||||
self.diag = LanguageDiagnostics()
|
||||
self._chunks_until_regen = 0
|
||||
|
||||
@abstractmethod
|
||||
def select_action(self, observation: dict[str, Any], state: RuntimeState) -> Any:
|
||||
"""Produce an action chunk from the observation + current language context."""
|
||||
|
||||
@abstractmethod
|
||||
def generate_text(
|
||||
self,
|
||||
kind: str,
|
||||
observation: dict[str, Any] | None,
|
||||
state: RuntimeState,
|
||||
user_text: str | None = None,
|
||||
) -> str:
|
||||
"""Generate one text stream (``kind``) and return the decoded string."""
|
||||
|
||||
def update_language_state(self, observation: dict[str, Any] | None, state: RuntimeState) -> None:
|
||||
"""Throttled regeneration of the language context (subtask / memory / ...)."""
|
||||
if self._chunks_until_regen > 0:
|
||||
self._chunks_until_regen -= 1
|
||||
return
|
||||
self._chunks_until_regen = max(1, self.gen.chunks_per_regen) - 1
|
||||
self._regenerate_context(observation, state)
|
||||
|
||||
def handle_interjection(
|
||||
self, user_text: str, observation: dict[str, Any] | None, state: RuntimeState
|
||||
) -> None:
|
||||
"""React to a mid-run user message by regenerating the plan."""
|
||||
out = self.generate_text("interjection", observation, state, user_text=user_text)
|
||||
plan = self.plan_from_text(out)
|
||||
if plan:
|
||||
state.set_context("plan", plan, label="plan")
|
||||
|
||||
def plan_from_text(self, text: str) -> str:
|
||||
"""Strip ``<say>`` speech markers from a generated plan."""
|
||||
plan, _speech = split_plan_and_say(text)
|
||||
return plan
|
||||
|
||||
def _regenerate_context(self, observation: dict[str, Any] | None, state: RuntimeState) -> None:
|
||||
"""Default hierarchy: regenerate the subtask, then memory when it changes.
|
||||
|
||||
Override for a policy with a different language hierarchy.
|
||||
"""
|
||||
if not self.gen.enable_subtask:
|
||||
# Preserve operator-provided subtasks in direct mode.
|
||||
return
|
||||
subtask = self._generate_filtered("subtask", observation, state)
|
||||
if subtask is None:
|
||||
return
|
||||
previous = state.language_context.get("subtask")
|
||||
if not state.set_context("subtask", subtask, label="subtask"):
|
||||
self.diag.repeat += 1
|
||||
return
|
||||
self.diag.repeat = 0
|
||||
if previous:
|
||||
state.extra["prior_subtask"] = previous
|
||||
if not self.gen.enable_memory:
|
||||
return
|
||||
memory = self._generate_filtered("memory", observation, state)
|
||||
if memory is not None:
|
||||
state.set_context("memory", memory, label="memory")
|
||||
|
||||
def _generate_filtered(
|
||||
self, kind: str, observation: dict[str, Any] | None, state: RuntimeState
|
||||
) -> str | None:
|
||||
"""Generate one ``kind``, record diagnostics, and drop empty output."""
|
||||
text = self.generate_text(kind, observation, state)
|
||||
self.diag.last_raw[kind] = text or ""
|
||||
if not text:
|
||||
count = self.diag._bump(self.diag.empty, kind)
|
||||
if count == 1 or count % 5 == 0:
|
||||
state.log(f" [info] {kind} gen returned empty (x{count})")
|
||||
return None
|
||||
return text
|
||||
|
||||
|
||||
class DirectTaskPolicyAdapter(BaseLanguageAdapter):
|
||||
"""Adapter for flat policies whose preprocessors condition actions on the operator's task."""
|
||||
|
||||
def select_action(self, observation: dict[str, Any], state: RuntimeState) -> Any:
|
||||
return self.policy.predict_action_chunk(observation)
|
||||
|
||||
def generate_text(
|
||||
self,
|
||||
kind: str,
|
||||
observation: dict[str, Any] | None,
|
||||
state: RuntimeState,
|
||||
user_text: str | None = None,
|
||||
) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def split_plan_and_say(text: str) -> tuple[str, str]:
|
||||
"""Split ``plan <say>speech</say>`` into ``(plan, speech)``."""
|
||||
if not text:
|
||||
return "", ""
|
||||
match = _SAY_RE.search(text)
|
||||
if not match:
|
||||
return text.strip(), ""
|
||||
speech = match.group(1).strip().strip('"').strip("'")
|
||||
plan = (text[: match.start()] + text[match.end() :]).strip()
|
||||
return plan, speech
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,349 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Small reusable runtime for language-conditioned robot policies."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RuntimeState:
|
||||
"""Explicit state shared by the runtime and policy adapter."""
|
||||
|
||||
task: str = ""
|
||||
language_context: dict[str, str] = field(default_factory=dict)
|
||||
action_queue: deque[Any] = field(default_factory=deque)
|
||||
events: set[str] = field(default_factory=set)
|
||||
log_lines: list[str] = field(default_factory=list)
|
||||
mode: str = "action"
|
||||
stop: bool = False
|
||||
tick: Tick | None = None
|
||||
actions_dispatched: int = 0
|
||||
action_deadline: float | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
revision: int = 0
|
||||
lock: Any = field(default_factory=threading.RLock, repr=False)
|
||||
|
||||
def emit(self, event_name: str) -> None:
|
||||
self.events.add(event_name)
|
||||
|
||||
def take_event(self, event_name: str) -> bool:
|
||||
if event_name not in self.events:
|
||||
return False
|
||||
self.events.remove(event_name)
|
||||
return True
|
||||
|
||||
def log(self, line: str) -> None:
|
||||
self.log_lines.append(line)
|
||||
|
||||
def set_context(self, key: str, value: str | None, *, label: str | None = None) -> bool:
|
||||
with self.lock:
|
||||
previous = self.language_context.get(key)
|
||||
if previous == value:
|
||||
return False
|
||||
if value is None:
|
||||
self.language_context.pop(key, None)
|
||||
else:
|
||||
self.language_context[key] = value
|
||||
self.revision += 1
|
||||
if label is not None and value:
|
||||
self.log(f" {label}: {value}")
|
||||
return True
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
try:
|
||||
return self[key]
|
||||
except KeyError:
|
||||
return default
|
||||
|
||||
def setdefault(self, key: str, default: Any = None) -> Any:
|
||||
current = self.get(key, None)
|
||||
if current is not None:
|
||||
return current
|
||||
self[key] = default
|
||||
return default
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
if hasattr(self, key):
|
||||
return getattr(self, key)
|
||||
if key in self.extra:
|
||||
return self.extra[key]
|
||||
raise KeyError(key)
|
||||
|
||||
def __setitem__(self, key: str, value: Any) -> None:
|
||||
with self.lock:
|
||||
if hasattr(self, key):
|
||||
if key == "mode" and self.mode != value:
|
||||
self.revision += 1
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
self.extra[key] = value
|
||||
|
||||
|
||||
class LanguageConditionedPolicyAdapter(Protocol):
|
||||
"""Runtime policy contract, implemented directly or through ``BaseLanguageAdapter``."""
|
||||
|
||||
def select_action(self, observation: dict[str, Any], state: RuntimeState) -> Any: ...
|
||||
|
||||
def update_language_state(self, observation: dict[str, Any] | None, state: RuntimeState) -> None: ...
|
||||
|
||||
def handle_interjection(
|
||||
self, user_text: str, observation: dict[str, Any] | None, state: RuntimeState
|
||||
) -> None: ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class Tick:
|
||||
index: int
|
||||
monotonic_seconds: float
|
||||
|
||||
|
||||
@dataclass
|
||||
class TickClock:
|
||||
max_rate_hz: float = 50.0
|
||||
_index: int = field(default=0, init=False)
|
||||
_last_seconds: float | None = field(default=None, init=False)
|
||||
|
||||
def advance(self) -> Tick:
|
||||
period = 1.0 / max(self.max_rate_hz, 0.1)
|
||||
now = time.monotonic()
|
||||
if self._last_seconds is not None:
|
||||
sleep_for = (self._last_seconds + period) - now
|
||||
if sleep_for > 0:
|
||||
time.sleep(sleep_for)
|
||||
now = time.monotonic()
|
||||
self._last_seconds = now
|
||||
self._index += 1
|
||||
return Tick(index=self._index, monotonic_seconds=now)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RateGate:
|
||||
hz: float
|
||||
_last_seconds: float | None = None
|
||||
|
||||
def due(self, tick: Tick, *, force: bool = False) -> bool:
|
||||
if force:
|
||||
self._last_seconds = tick.monotonic_seconds
|
||||
return True
|
||||
period = 1.0 / max(self.hz, 1e-6)
|
||||
if self._last_seconds is None or tick.monotonic_seconds - self._last_seconds >= period:
|
||||
self._last_seconds = tick.monotonic_seconds
|
||||
return True
|
||||
return False
|
||||
|
||||
def rearm(self) -> None:
|
||||
self._last_seconds = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class LanguageConditionedRuntime:
|
||||
"""Generic tick loop for language-conditioned robot policies."""
|
||||
|
||||
policy_adapter: LanguageConditionedPolicyAdapter
|
||||
observation_provider: Callable[[], dict[str, Any] | None] | None = None
|
||||
action_executor: Callable[[Any], None] | None = None
|
||||
event_collector: Callable[[RuntimeState], None] | None = None
|
||||
chunk_hz: float = 4.0
|
||||
ctrl_hz: float = 50.0
|
||||
high_level_hz: float = 1.0
|
||||
max_rate_hz: float = 50.0
|
||||
|
||||
state: RuntimeState = field(default_factory=RuntimeState)
|
||||
_chunk_gate: _RateGate = field(init=False)
|
||||
_ctrl_gate: _RateGate = field(init=False)
|
||||
_language_gate: _RateGate = field(init=False)
|
||||
_stop: bool = field(default=False, init=False)
|
||||
_last_dispatch_seconds: float | None = field(default=None, init=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._chunk_gate = _RateGate(self.chunk_hz)
|
||||
self._ctrl_gate = _RateGate(self.ctrl_hz)
|
||||
self._language_gate = _RateGate(self.high_level_hz)
|
||||
|
||||
@property
|
||||
def policy(self) -> Any:
|
||||
return getattr(self.policy_adapter, "policy", self.policy_adapter)
|
||||
|
||||
def set_task(self, task: str) -> None:
|
||||
with self.state.lock:
|
||||
if self.state.task != task:
|
||||
self.state.revision += 1
|
||||
self.state.task = task
|
||||
self.state.log(f"Task: {task}")
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop = True
|
||||
self.state.stop = True
|
||||
|
||||
def run(self, *, max_ticks: int | None = None) -> None:
|
||||
clock = TickClock(max_rate_hz=self.max_rate_hz)
|
||||
while not self._stop:
|
||||
tick = clock.advance()
|
||||
self._run_tick(tick)
|
||||
self._flush_logs()
|
||||
if self.state.stop:
|
||||
self._stop = True
|
||||
if max_ticks is not None and tick.index >= max_ticks:
|
||||
break
|
||||
self._on_shutdown()
|
||||
|
||||
def step_once(self) -> list[str]:
|
||||
previous = self.state.tick.index if self.state.tick is not None else 0
|
||||
tick = Tick(index=previous + 1, monotonic_seconds=time.monotonic())
|
||||
self._run_tick(tick, force_rates=True)
|
||||
return list(self.state.log_lines)
|
||||
|
||||
def _run_tick(self, tick: Tick, *, force_rates: bool = False) -> None:
|
||||
self.state.tick = tick
|
||||
self.state.log_lines = []
|
||||
if self.event_collector is not None:
|
||||
self.event_collector(self.state)
|
||||
self._handle_action_deadline()
|
||||
if self.state.stop:
|
||||
return
|
||||
self.maybe_update_language_state(force=force_rates)
|
||||
self.maybe_handle_user_events()
|
||||
self.maybe_enqueue_action_chunk(force=force_rates)
|
||||
self.dispatch_action(force=force_rates)
|
||||
self.state.events.clear()
|
||||
|
||||
def _current_observation(self) -> dict[str, Any] | None:
|
||||
if self.observation_provider is None:
|
||||
return None
|
||||
try:
|
||||
return self.observation_provider()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("observation_provider failed: %s", exc)
|
||||
return None
|
||||
|
||||
def maybe_update_language_state(self, *, force: bool = False) -> None:
|
||||
if self.state.mode != "action" or not self.state.task:
|
||||
return
|
||||
if self.state.action_queue:
|
||||
self._language_gate.rearm()
|
||||
return
|
||||
if self.state.tick is None or not self._language_gate.due(self.state.tick, force=force):
|
||||
return
|
||||
observation = self._current_observation()
|
||||
try:
|
||||
self.policy_adapter.update_language_state(observation, self.state)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("language update failed: %s", exc, exc_info=logger.isEnabledFor(logging.DEBUG))
|
||||
self.state.log(f" [warn] language update failed: {type(exc).__name__}: {exc}")
|
||||
|
||||
def maybe_handle_user_events(self) -> None:
|
||||
if self.state.take_event("user_interjection"):
|
||||
self._handle_user_interjection()
|
||||
|
||||
def _handle_user_interjection(self) -> None:
|
||||
text = str(self.state.extra.get("recent_interjection") or "")
|
||||
if not text:
|
||||
return
|
||||
observation = self._current_observation()
|
||||
self.policy_adapter.handle_interjection(text, observation, self.state)
|
||||
self.state.extra["recent_interjection"] = None
|
||||
|
||||
def maybe_enqueue_action_chunk(self, *, force: bool = False) -> None:
|
||||
with self.state.lock:
|
||||
if self.state.mode != "action" or not self.state.task:
|
||||
return
|
||||
if self.state.action_queue:
|
||||
return
|
||||
if self.state.tick is None or not self._chunk_gate.due(self.state.tick, force=force):
|
||||
return
|
||||
revision = self.state.revision
|
||||
observation = self._current_observation()
|
||||
if observation is None:
|
||||
return
|
||||
try:
|
||||
chunk = self.policy_adapter.select_action(observation, self.state)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("select_action failed: %s", exc, exc_info=logger.isEnabledFor(logging.DEBUG))
|
||||
self.state.log(f" [warn] select_action failed: {type(exc).__name__}: {exc}")
|
||||
return
|
||||
with self.state.lock:
|
||||
if (
|
||||
self.state.revision != revision
|
||||
or self.state.mode != "action"
|
||||
or self.state.stop
|
||||
or self._stop
|
||||
):
|
||||
logger.info("Discarded an action chunk invalidated during inference.")
|
||||
return
|
||||
self._enqueue_chunk(chunk)
|
||||
|
||||
def _enqueue_chunk(self, chunk: Any) -> None:
|
||||
if chunk is None:
|
||||
return
|
||||
chunk_iter = chunk[0] if getattr(chunk, "ndim", None) == 3 else chunk
|
||||
if getattr(chunk_iter, "ndim", None) == 1:
|
||||
chunk_iter = chunk_iter.unsqueeze(0)
|
||||
for step in chunk_iter:
|
||||
self.state.action_queue.append(step.unsqueeze(0) if hasattr(step, "unsqueeze") else step)
|
||||
try:
|
||||
self.state.extra["last_chunk_size"] = int(chunk_iter.shape[0])
|
||||
except Exception: # noqa: BLE001
|
||||
self.state.extra["last_chunk_size"] = len(self.state.action_queue)
|
||||
|
||||
def dispatch_action(self, *, force: bool = False) -> None:
|
||||
if self.state.mode != "action":
|
||||
self._last_dispatch_seconds = None
|
||||
return
|
||||
if self.state.tick is None or not self._ctrl_gate.due(self.state.tick, force=force):
|
||||
return
|
||||
queue = self.state.action_queue
|
||||
if not queue:
|
||||
self._last_dispatch_seconds = None
|
||||
return
|
||||
now = time.monotonic()
|
||||
if self._last_dispatch_seconds is None or self.ctrl_hz <= 0:
|
||||
n_to_pop = 1
|
||||
else:
|
||||
n_to_pop = max(1, min(len(queue), int(round((now - self._last_dispatch_seconds) * self.ctrl_hz))))
|
||||
self._last_dispatch_seconds = now
|
||||
latest = None
|
||||
for _ in range(n_to_pop):
|
||||
if not queue:
|
||||
break
|
||||
latest = queue.popleft()
|
||||
self.state.actions_dispatched += 1
|
||||
if latest is not None and self.action_executor is not None:
|
||||
self.action_executor(latest)
|
||||
|
||||
def _handle_action_deadline(self) -> None:
|
||||
deadline = self.state.action_deadline
|
||||
if self.state.mode == "action" and deadline is not None and time.monotonic() >= deadline:
|
||||
self.state.mode = "paused"
|
||||
self.state.action_deadline = None
|
||||
self.state.action_queue.clear()
|
||||
self.state.log("timed action elapsed — paused")
|
||||
|
||||
def _flush_logs(self) -> None:
|
||||
for line in self.state.log_lines:
|
||||
print(f"[runtime] {line}", flush=True)
|
||||
|
||||
def _on_shutdown(self) -> None:
|
||||
self.state.action_queue.clear()
|
||||
print("[runtime] stopped", flush=True)
|
||||
@@ -1,39 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Lazy mapping from policy types to language-runtime adapters."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
_ADAPTERS: dict[str, str] = {
|
||||
"pi052": "lerobot.policies.pi052.inference.pi052_adapter:PI052PolicyAdapter",
|
||||
"pi05": "lerobot.runtime.adapter:DirectTaskPolicyAdapter",
|
||||
"molmoact2": "lerobot.runtime.adapter:DirectTaskPolicyAdapter",
|
||||
}
|
||||
|
||||
|
||||
def get_language_adapter_factory(policy_type: str) -> Callable[..., Any]:
|
||||
"""Return the adapter class registered for ``policy_type``."""
|
||||
spec = _ADAPTERS.get(policy_type)
|
||||
if spec is None:
|
||||
raise ValueError(
|
||||
f"No language-runtime adapter registered for policy type {policy_type!r}. "
|
||||
f"Registered: {sorted(_ADAPTERS)}. Add an entry to lerobot.runtime.registry."
|
||||
)
|
||||
module_path, class_name = spec.split(":")
|
||||
return getattr(importlib.import_module(module_path), class_name)
|
||||
@@ -1,406 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""RoboCasa backend for interactive language-conditioned rollouts.
|
||||
|
||||
It reuses the eval observation/action pipeline while prompts control a persistent selected scene.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from lerobot.utils.io_utils import StreamingVideoWriter
|
||||
from lerobot.utils.video_annotation import annotate_frame
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _short_cam_name(cam: str) -> str:
|
||||
"""Human-friendly view label for a RoboCasa camera name."""
|
||||
c = cam.replace("robot0_", "")
|
||||
return {
|
||||
"agentview_left": "left",
|
||||
"agentview_right": "right",
|
||||
"eye_in_hand": "wrist",
|
||||
}.get(c, c)
|
||||
|
||||
|
||||
def _label_panel(img: np.ndarray, label: str) -> np.ndarray:
|
||||
"""Draw a small camera-view label in the bottom-left corner of a panel."""
|
||||
try:
|
||||
import cv2 # noqa: PLC0415
|
||||
except ImportError:
|
||||
return img
|
||||
y = img.shape[0] - 6
|
||||
cv2.putText(img, label, (5, y), cv2.FONT_HERSHEY_SIMPLEX, 0.45, (0, 0, 0), 3, cv2.LINE_AA)
|
||||
cv2.putText(img, label, (5, y), cv2.FONT_HERSHEY_SIMPLEX, 0.45, (255, 255, 0), 1, cv2.LINE_AA)
|
||||
return img
|
||||
|
||||
|
||||
# Two workers avoid broken single-worker EGL rendering; only env 0 is displayed.
|
||||
_SIM_N_ENVS = 2
|
||||
|
||||
|
||||
def create_sim_env(
|
||||
*,
|
||||
task: str,
|
||||
split: str | None,
|
||||
obj_registries: list[str],
|
||||
seed: int | None,
|
||||
render_size: int = 384,
|
||||
) -> tuple[Any, dict]:
|
||||
"""Create and reset the vectorized RoboCasa environment before CUDA initializes.
|
||||
|
||||
Two workers keep EGL stable, while only env 0 is driven and displayed.
|
||||
"""
|
||||
from lerobot.envs.configs import RoboCasaEnv as RoboCasaEnvConfig # noqa: PLC0415
|
||||
|
||||
# The policy resizes inputs, so render_size only affects display quality and cost.
|
||||
env_cfg = RoboCasaEnvConfig(
|
||||
task=task,
|
||||
split=split,
|
||||
obj_registries=list(obj_registries),
|
||||
observation_height=render_size,
|
||||
observation_width=render_size,
|
||||
)
|
||||
# Keep one kitchen alive across sequential prompts.
|
||||
envs = env_cfg.create_envs(
|
||||
n_envs=_SIM_N_ENVS,
|
||||
use_async_envs=True,
|
||||
terminate_on_success=False,
|
||||
horizon=100_000,
|
||||
)
|
||||
env = envs[next(iter(envs))][0]
|
||||
logger.info("[sim] resetting RoboCasa scene task=%r split=%r (n_envs=%d)", task, split, _SIM_N_ENVS)
|
||||
seeds = None if seed is None else [seed + i for i in range(_SIM_N_ENVS)]
|
||||
obs, _ = env.reset(seed=seeds)
|
||||
return env, obs
|
||||
|
||||
|
||||
def start_mjpeg_server(port: int, get_frame: Callable[[], np.ndarray | None]) -> Any:
|
||||
"""Start an MJPEG server that shows a placeholder until ``get_frame`` returns frames."""
|
||||
import io # noqa: PLC0415
|
||||
import threading # noqa: PLC0415
|
||||
import time # noqa: PLC0415
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer # noqa: PLC0415
|
||||
|
||||
from PIL import Image # noqa: PLC0415
|
||||
|
||||
_placeholder = Image.new("RGB", (256, 256), (17, 17, 17))
|
||||
|
||||
class _Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, *args): # silence per-request logging
|
||||
pass
|
||||
|
||||
def do_GET(self): # noqa: N802
|
||||
if self.path in ("/", "/index.html"):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/html")
|
||||
self.end_headers()
|
||||
self.wfile.write(
|
||||
b"<html><body style='margin:0;background:#111;text-align:center'>"
|
||||
b"<img src='/stream' style='max-width:100vw;max-height:100vh;"
|
||||
b"image-rendering:pixelated'></body></html>"
|
||||
)
|
||||
return
|
||||
if self.path != "/stream":
|
||||
self.send_response(404)
|
||||
self.end_headers()
|
||||
return
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "multipart/x-mixed-replace; boundary=frame")
|
||||
self.end_headers()
|
||||
try:
|
||||
while True:
|
||||
frame = get_frame()
|
||||
buf = io.BytesIO()
|
||||
img = Image.fromarray(frame) if frame is not None else _placeholder
|
||||
img.save(buf, format="JPEG", quality=80)
|
||||
data = buf.getvalue()
|
||||
self.wfile.write(
|
||||
b"--frame\r\nContent-Type: image/jpeg\r\nContent-Length: "
|
||||
+ str(len(data)).encode()
|
||||
+ b"\r\n\r\n"
|
||||
+ data
|
||||
+ b"\r\n"
|
||||
)
|
||||
time.sleep(0.05)
|
||||
except (BrokenPipeError, ConnectionResetError):
|
||||
pass
|
||||
|
||||
try:
|
||||
# Bind all interfaces intentionally so the viewer remains reachable
|
||||
# through the documented SSH port-forwarding workflow.
|
||||
server = ThreadingHTTPServer(("0.0.0.0", port), _Handler) # nosec B104
|
||||
except OSError as exc:
|
||||
logger.warning("[sim] could not start live stream on port %d: %s", port, exc)
|
||||
print(f"[runtime] WARNING: live stream port {port} unavailable ({exc})", flush=True)
|
||||
return None
|
||||
threading.Thread(target=server.serve_forever, daemon=True, name="sim-mjpeg").start()
|
||||
print(
|
||||
f"[runtime] live view: http://localhost:{port} "
|
||||
f"(over SSH: ssh -L {port}:localhost:{port} <host>) — loading until scene is ready",
|
||||
flush=True,
|
||||
)
|
||||
return server
|
||||
|
||||
|
||||
class RoboCasaSimBackend:
|
||||
"""Expose a RoboCasa environment through the runtime observation/action contract.
|
||||
|
||||
The environment must be created before the policy initializes CUDA.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
env: Any,
|
||||
last_obs: dict,
|
||||
task: str,
|
||||
seed: int | None,
|
||||
device: str,
|
||||
preprocessor: Any,
|
||||
postprocessor: Any,
|
||||
record: bool = True,
|
||||
output_dir: str | None = None,
|
||||
view_cams: list[str] | None = None,
|
||||
) -> None:
|
||||
self.env = env
|
||||
self._last_obs = last_obs
|
||||
self._scene_task = task
|
||||
self._view_cams = view_cams or [
|
||||
"robot0_agentview_left",
|
||||
"robot0_eye_in_hand",
|
||||
"robot0_agentview_right",
|
||||
]
|
||||
self.device = torch.device(device) if isinstance(device, str) else device
|
||||
self.preprocessor = preprocessor
|
||||
self.postprocessor = postprocessor
|
||||
self.seed = seed
|
||||
self.record = record
|
||||
self.output_dir = Path(output_dir) if output_dir else Path("outputs/runtime_sim")
|
||||
|
||||
self._video_writer: StreamingVideoWriter | None = None
|
||||
self._video_path: Path | None = None
|
||||
self._live_counter = 0
|
||||
self._latest_frame: np.ndarray | None = None
|
||||
self._stream_server: Any = None
|
||||
self._reset_count = 0
|
||||
# Bind these after runtime construction for live annotations.
|
||||
self._task_getter: Callable[[], str | None] | None = None
|
||||
self._subtask_getter: Callable[[], str | None] | None = None
|
||||
self._memory_getter: Callable[[], str | None] | None = None
|
||||
logger.info("[sim] scene ready — task_description=%r", self._scene_description())
|
||||
|
||||
def bind_runtime(self, runtime: Any) -> None:
|
||||
"""Wire live task/subtask/memory getters from the runtime state."""
|
||||
self._task_getter = lambda: runtime.state.get("task")
|
||||
self._subtask_getter = lambda: runtime.state.language_context.get("subtask")
|
||||
self._memory_getter = lambda: (runtime.state.get("language_context") or {}).get("memory")
|
||||
|
||||
def _scene_description(self) -> str:
|
||||
try:
|
||||
return str(self.env.get_attr("task_description")[0]) or self._scene_task
|
||||
except Exception: # noqa: BLE001
|
||||
return self._scene_task
|
||||
|
||||
def _current_task(self) -> str:
|
||||
task = self._task_getter() if self._task_getter else None
|
||||
return task or self._scene_description() or self._scene_task
|
||||
|
||||
def reset_scene(self) -> None:
|
||||
"""Re-roll the kitchen: reset the env to a fresh scene (new layout/style).
|
||||
|
||||
Uses a new seed each call so ``/reset`` explores different kitchens.
|
||||
"""
|
||||
self._reset_count += 1
|
||||
n = self.env.num_envs
|
||||
if self.seed is None:
|
||||
seeds = None
|
||||
else:
|
||||
base = self.seed + self._reset_count * 1000
|
||||
seeds = [base + i for i in range(n)]
|
||||
obs, _ = self.env.reset(seed=seeds)
|
||||
self._last_obs = obs
|
||||
logger.info("[sim] scene reset (#%d)", self._reset_count)
|
||||
|
||||
def _env0_obs(self) -> dict:
|
||||
"""Slice env 0 out of the batched vec-env observation (batch of 1)."""
|
||||
raw = self._last_obs or {}
|
||||
pixels = raw.get("pixels")
|
||||
out: dict[str, Any] = {}
|
||||
if isinstance(pixels, dict):
|
||||
out["pixels"] = {k: np.asarray(v)[0:1] for k, v in pixels.items()}
|
||||
agent_pos = raw.get("agent_pos")
|
||||
if agent_pos is not None:
|
||||
out["agent_pos"] = np.asarray(agent_pos)[0:1]
|
||||
return out
|
||||
|
||||
def observation_provider(self) -> dict | None:
|
||||
from lerobot.envs.utils import preprocess_observation # noqa: PLC0415
|
||||
|
||||
try:
|
||||
obs = preprocess_observation(self._env0_obs())
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("[sim] preprocess_observation failed: %s", exc)
|
||||
return None
|
||||
# The adapter later replaces this recipe input with its generated subtask.
|
||||
obs["task"] = [self._current_task()]
|
||||
if self.preprocessor is not None:
|
||||
try:
|
||||
obs = self.preprocessor(obs)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("[sim] preprocessor failed: %s", exc)
|
||||
return None
|
||||
return {
|
||||
k: (v.to(self.device) if isinstance(v, torch.Tensor) else v)
|
||||
for k, v in obs.items()
|
||||
if isinstance(k, str) and k.startswith("observation.")
|
||||
}
|
||||
|
||||
def action_executor(self, action: Any) -> None:
|
||||
try:
|
||||
if self.postprocessor is not None:
|
||||
action = self.postprocessor(action)
|
||||
if isinstance(action, torch.Tensor):
|
||||
if action.ndim > 1 and action.shape[0] == 1:
|
||||
action = action.squeeze(0)
|
||||
action = action.detach().to("cpu").numpy()
|
||||
# Tile env 0's action because the extra workers exist only for EGL stability.
|
||||
action_row = np.asarray(action, dtype=np.float32).reshape(-1)
|
||||
action_np = np.tile(action_row, (self.env.num_envs, 1))
|
||||
obs, _reward, terminated, truncated, _info = self.env.step(action_np)
|
||||
self._last_obs = obs
|
||||
self._capture_frame()
|
||||
# AsyncVectorEnv resets terminated sub-environments automatically.
|
||||
if bool(np.any(terminated)) or bool(np.any(truncated)):
|
||||
logger.info("[sim] episode ended — scene auto-reset")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.error("[sim] env.step failed: %s", exc, exc_info=True)
|
||||
|
||||
def _multiview_frame(self) -> np.ndarray | None:
|
||||
"""Label and compose env 0's existing observation views without extra rendering."""
|
||||
pixels = (self._last_obs or {}).get("pixels")
|
||||
if not isinstance(pixels, dict) or not pixels:
|
||||
return None
|
||||
panels: list[np.ndarray] = []
|
||||
for cam in self._view_cams:
|
||||
v = pixels.get(cam)
|
||||
if v is None:
|
||||
continue
|
||||
img = np.asarray(v)
|
||||
if img.ndim == 4: # (n_envs, H, W, C) -> env 0
|
||||
img = img[0]
|
||||
if img.ndim != 3 or img.shape[-1] != 3:
|
||||
continue
|
||||
panels.append(_label_panel(np.ascontiguousarray(img.astype(np.uint8)), _short_cam_name(cam)))
|
||||
if not panels:
|
||||
return None
|
||||
h = min(p.shape[0] for p in panels)
|
||||
panels = [p[:h] for p in panels]
|
||||
return np.concatenate(panels, axis=1)
|
||||
|
||||
def _capture_frame(self) -> None:
|
||||
frame = self._multiview_frame()
|
||||
if frame is None: # fallback to single env.render()
|
||||
try:
|
||||
rendered = self.env.call("render")[0]
|
||||
if isinstance(rendered, np.ndarray) and rendered.ndim == 3:
|
||||
frame = rendered
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("[sim] render failed: %s", exc)
|
||||
if frame is None:
|
||||
return
|
||||
subtask = self._subtask_getter() if self._subtask_getter else None
|
||||
memory = self._memory_getter() if self._memory_getter else None
|
||||
annotated = annotate_frame(
|
||||
frame,
|
||||
(("Task", self._current_task()), ("Subtask", subtask), ("Memory", memory)),
|
||||
)
|
||||
self._latest_frame = annotated # served by the live MJPEG stream
|
||||
self._write_live_frame(annotated)
|
||||
if self.record:
|
||||
self._write_recording_frame(annotated)
|
||||
|
||||
def _write_recording_frame(self, frame: np.ndarray) -> None:
|
||||
try:
|
||||
if self._video_writer is None:
|
||||
self.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
stamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
self._video_path = self.output_dir / f"sim_{stamp}.mp4"
|
||||
fps = int((getattr(self.env, "metadata", None) or {}).get("render_fps", 20))
|
||||
self._video_writer = StreamingVideoWriter(self._video_path, fps)
|
||||
self._video_writer.add_frame(frame)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("[sim] video encoding failed: %s", exc)
|
||||
self.record = False
|
||||
|
||||
def _write_live_frame(self, frame: np.ndarray) -> None:
|
||||
"""Write a rolling latest.png every few frames for live viewing over SSH.
|
||||
|
||||
Open ``{output_dir}/latest.png`` in an editor/viewer and refresh to watch
|
||||
the rollout in near-real-time without a GUI window. Written atomically
|
||||
(temp + replace) so a reader never sees a half-written file.
|
||||
"""
|
||||
self._live_counter += 1
|
||||
if self._live_counter % 3 != 0:
|
||||
return
|
||||
try:
|
||||
import os # noqa: PLC0415
|
||||
|
||||
from PIL import Image # noqa: PLC0415
|
||||
|
||||
self.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
tmp = self.output_dir / ".latest.tmp.png"
|
||||
Image.fromarray(frame).save(tmp)
|
||||
os.replace(tmp, self.output_dir / "latest.png")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("[sim] live frame write failed: %s", exc)
|
||||
|
||||
def _flush_video(self) -> None:
|
||||
if self._video_writer is None:
|
||||
return
|
||||
writer = self._video_writer
|
||||
self._video_writer = None
|
||||
try:
|
||||
writer.close()
|
||||
logger.info("[sim] wrote video (%d frames) to %s", writer.frames_written, self._video_path)
|
||||
print(f"[runtime] sim video saved to {self._video_path}", flush=True)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("[sim] video close failed: %s", exc)
|
||||
|
||||
def attach_stream_server(self, server: Any) -> None:
|
||||
"""Attach an already-running MJPEG server so disconnect() can stop it."""
|
||||
self._stream_server = server
|
||||
|
||||
def disconnect(self) -> None:
|
||||
"""Match the robot backend's cleanup contract."""
|
||||
if self._stream_server is not None:
|
||||
try:
|
||||
self._stream_server.shutdown()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("[sim] stream server shutdown raised %s", exc)
|
||||
self._flush_video()
|
||||
try:
|
||||
self.env.close()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.debug("[sim] env.close raised %s", exc)
|
||||
@@ -28,7 +28,12 @@ For distributed runs, see ``examples/annotations/run_hf_job.py``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
from lerobot.annotations.steerable_pipeline.config import AnnotationPipelineConfig
|
||||
from lerobot.annotations.steerable_pipeline.executor import Executor
|
||||
@@ -42,6 +47,12 @@ from lerobot.annotations.steerable_pipeline.validator import StagingValidator
|
||||
from lerobot.annotations.steerable_pipeline.vlm_client import make_vlm_client
|
||||
from lerobot.annotations.steerable_pipeline.writer import LanguageColumnsWriter
|
||||
from lerobot.configs import parser
|
||||
from lerobot.utils.import_utils import _datasets_available, require_package
|
||||
|
||||
if TYPE_CHECKING or _datasets_available:
|
||||
from lerobot.datasets.dataset_metadata import CODEBASE_VERSION
|
||||
from lerobot.datasets.io_utils import load_info
|
||||
from lerobot.datasets.utils import create_lerobot_dataset_card
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -50,8 +61,6 @@ def _resolve_root(cfg: AnnotationPipelineConfig) -> Path:
|
||||
if cfg.root is not None:
|
||||
return Path(cfg.root)
|
||||
if cfg.repo_id is not None:
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
return Path(snapshot_download(repo_id=cfg.repo_id, repo_type="dataset"))
|
||||
raise ValueError("Either --root or --repo_id must be provided.")
|
||||
|
||||
@@ -125,10 +134,7 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
|
||||
Pushes to ``cfg.new_repo_id`` when set, otherwise back to ``cfg.repo_id``.
|
||||
"""
|
||||
from huggingface_hub import HfApi # noqa: PLC0415
|
||||
|
||||
from lerobot.datasets.io_utils import load_info # noqa: PLC0415
|
||||
from lerobot.datasets.utils import create_lerobot_dataset_card # noqa: PLC0415
|
||||
require_package("datasets", "dataset")
|
||||
|
||||
repo_id = cfg.new_repo_id or cfg.repo_id
|
||||
commit_message = cfg.push_commit_message or "Add steerable annotations (lerobot-annotate)"
|
||||
@@ -163,8 +169,6 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
# ``RevisionNotFoundError``. Read the version straight from the
|
||||
# dataset's own ``meta/info.json`` so we tag whatever the writer
|
||||
# actually wrote (no accidental drift if the codebase floor moves).
|
||||
from lerobot.datasets.dataset_metadata import CODEBASE_VERSION # noqa: PLC0415
|
||||
|
||||
version_tag = (
|
||||
dataset_info.codebase_version if dataset_info.codebase_version.startswith("v") else CODEBASE_VERSION
|
||||
)
|
||||
@@ -178,10 +182,6 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
||||
tag_kwargs["revision"] = revision
|
||||
|
||||
try:
|
||||
from contextlib import suppress # noqa: PLC0415
|
||||
|
||||
from huggingface_hub.errors import RevisionNotFoundError # noqa: PLC0415
|
||||
|
||||
with suppress(RevisionNotFoundError):
|
||||
api.delete_tag(repo_id, tag=version_tag, repo_type="dataset")
|
||||
api.create_tag(**tag_kwargs)
|
||||
|
||||
@@ -94,19 +94,6 @@ from lerobot.utils.utils import (
|
||||
init_logging,
|
||||
inside_slurm,
|
||||
)
|
||||
from lerobot.utils.video_annotation import annotate_frame
|
||||
|
||||
|
||||
def _annotate_eval_frames(frames: np.ndarray, task: str | None, subtask: str | None) -> np.ndarray:
|
||||
"""Overlay the high-level task and predicted subtask onto rendered frames.
|
||||
|
||||
``frames`` is ``(n_envs, H, W, C)`` uint8. Best-effort: if OpenCV isn't
|
||||
available the frames are returned unchanged so eval never fails over a
|
||||
visualization concern.
|
||||
"""
|
||||
if frames.ndim != 4 or frames.shape[-1] != 3:
|
||||
return frames
|
||||
return np.stack([annotate_frame(frame, (("Task", task), ("Subtask", subtask))) for frame in frames])
|
||||
|
||||
|
||||
def _env_features_to_dataset_features(env_features: dict) -> dict:
|
||||
@@ -487,36 +474,11 @@ def eval_policy(
|
||||
return
|
||||
n_to_render_now = min(max_episodes_rendered - n_episodes_rendered, env.num_envs)
|
||||
if isinstance(env, gym.vector.SyncVectorEnv):
|
||||
frames = np.stack([env.envs[i].render() for i in range(n_to_render_now)]) # noqa: B023
|
||||
ep_frames.append(np.stack([env.envs[i].render() for i in range(n_to_render_now)])) # noqa: B023
|
||||
elif hasattr(env, "call"):
|
||||
# Here we must render all frames and discard any we don't need.
|
||||
# Covers AsyncVectorEnv and _LazyAsyncVectorEnv (which wraps one).
|
||||
frames = np.stack(env.call("render")[:n_to_render_now])
|
||||
else:
|
||||
return
|
||||
|
||||
# Overlay the high-level task and (for hierarchical policies like
|
||||
# pi052) the predicted low-level subtask onto each frame. Both are
|
||||
# best-effort: missing values just skip that line.
|
||||
try:
|
||||
tasks = list(env.call("task_description"))
|
||||
except (AttributeError, NotImplementedError):
|
||||
try:
|
||||
tasks = list(env.call("task"))
|
||||
except (AttributeError, NotImplementedError):
|
||||
tasks = None
|
||||
subtasks = getattr(policy, "last_subtasks", None)
|
||||
annotated = []
|
||||
for i in range(frames.shape[0]):
|
||||
subtask_i = subtasks[i] if subtasks is not None and i < len(subtasks) else None
|
||||
annotated.append(
|
||||
_annotate_eval_frames(
|
||||
frames[i : i + 1],
|
||||
tasks[i] if tasks is not None and i < len(tasks) else None,
|
||||
subtask_i,
|
||||
)[0]
|
||||
)
|
||||
ep_frames.append(np.stack(annotated))
|
||||
ep_frames.append(np.stack(env.call("render")[:n_to_render_now]))
|
||||
|
||||
if max_episodes_rendered > 0:
|
||||
video_paths: list[str] = []
|
||||
|
||||
@@ -151,7 +151,6 @@ Usage examples
|
||||
"""
|
||||
|
||||
import logging
|
||||
import sys
|
||||
|
||||
from lerobot.cameras.opencv import OpenCVCameraConfig # noqa: F401
|
||||
from lerobot.cameras.realsense import RealSenseCameraConfig # noqa: F401
|
||||
@@ -242,69 +241,10 @@ def rollout(cfg: RolloutConfig):
|
||||
logger.info("Rollout finished")
|
||||
|
||||
|
||||
_LANGUAGE_RUNTIME_FLAGS = {
|
||||
"--language",
|
||||
"--no_robot",
|
||||
"--sim",
|
||||
"--direct_subtask",
|
||||
"--sim.direct_subtask",
|
||||
"--disable_memory",
|
||||
"--fp8",
|
||||
}
|
||||
_LANGUAGE_RUNTIME_PREFIXES = (
|
||||
"--sim.",
|
||||
"--chunk_hz",
|
||||
"--ctrl_hz",
|
||||
"--high_level_hz",
|
||||
"--subtask_chunks_per_gen",
|
||||
"--text_min_new_tokens",
|
||||
"--text_temperature",
|
||||
"--text_top_p",
|
||||
)
|
||||
|
||||
|
||||
def _uses_language_runtime(argv: list[str]) -> bool:
|
||||
"""Return whether *argv* selects the interactive language runtime.
|
||||
|
||||
``--language`` is the explicit selector for real-robot runs whose other
|
||||
options overlap with the standard rollout CLI. Language-only options also
|
||||
select it automatically, which keeps the former language-runtime examples
|
||||
working after replacing their command name with ``lerobot-rollout``.
|
||||
"""
|
||||
return any(
|
||||
arg.split("=", 1)[0] in _LANGUAGE_RUNTIME_FLAGS or arg.startswith(_LANGUAGE_RUNTIME_PREFIXES)
|
||||
for arg in argv
|
||||
)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None):
|
||||
"""CLI entry point for ``lerobot-rollout``.
|
||||
|
||||
Standard policy deployment continues through :class:`RolloutConfig`.
|
||||
Interactive language-conditioned and RoboCasa runs share this entry point
|
||||
and are selected with ``--language`` or any language-runtime-only option.
|
||||
"""
|
||||
def main():
|
||||
"""CLI entry point for ``lerobot-rollout``."""
|
||||
register_third_party_plugins()
|
||||
cli_args = list(sys.argv[1:] if argv is None else argv)
|
||||
if _uses_language_runtime(cli_args):
|
||||
from lerobot.runtime.cli import run as run_language_runtime
|
||||
|
||||
# ``--language`` is a dispatcher flag, not part of the runtime's own
|
||||
# argparse surface. All other arguments pass through unchanged.
|
||||
runtime_args = [arg for arg in cli_args if arg != "--language"]
|
||||
return run_language_runtime(runtime_args, prog="lerobot-rollout")
|
||||
|
||||
if argv is None:
|
||||
return rollout()
|
||||
|
||||
# draccus reads sys.argv. Supporting an explicit argv keeps this entry
|
||||
# point easy to smoke-test and mirrors the language-runtime branch above.
|
||||
previous_argv = sys.argv
|
||||
try:
|
||||
sys.argv = [previous_argv[0], *cli_args]
|
||||
return rollout()
|
||||
finally:
|
||||
sys.argv = previous_argv
|
||||
rollout()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -20,11 +20,9 @@ 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
|
||||
|
||||
@@ -57,7 +55,6 @@ from lerobot.envs import close_envs, make_env, make_env_pre_post_processors
|
||||
from lerobot.jobs import submit_to_hf
|
||||
from lerobot.optim.factory import make_optimizer_and_scheduler
|
||||
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
||||
from lerobot.processor.rename_processor import rename_transition_keys
|
||||
from lerobot.rewards import make_reward_pre_post_processors
|
||||
from lerobot.utils.collate import lerobot_collate_fn
|
||||
from lerobot.utils.import_utils import register_third_party_plugins
|
||||
@@ -84,8 +81,6 @@ def update_policy(
|
||||
lr_scheduler=None,
|
||||
lock=None,
|
||||
sample_weighter=None,
|
||||
log_metrics: bool = True,
|
||||
track_update_time: bool = True,
|
||||
) -> tuple[MetricsTracker, dict | None]:
|
||||
"""
|
||||
Performs a single training step to update the policy's weights.
|
||||
@@ -103,7 +98,6 @@ 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:
|
||||
@@ -149,16 +143,13 @@ def update_policy(
|
||||
# Use accelerator's backward method
|
||||
accelerator.backward(loss)
|
||||
|
||||
# Accelerate suppresses gradient synchronization and optimizer updates on intermediate
|
||||
# microbatches. Clip and report the norm only when the accumulated update is complete.
|
||||
grad_norm = None
|
||||
if accelerator.sync_gradients:
|
||||
if grad_clip_norm > 0:
|
||||
grad_norm = accelerator.clip_grad_norm_(policy.parameters(), grad_clip_norm)
|
||||
else:
|
||||
grad_norm = torch.nn.utils.clip_grad_norm_(
|
||||
policy.parameters(), float("inf"), error_if_nonfinite=False
|
||||
)
|
||||
# Clip gradients if specified
|
||||
if grad_clip_norm > 0:
|
||||
grad_norm = accelerator.clip_grad_norm_(policy.parameters(), grad_clip_norm)
|
||||
else:
|
||||
grad_norm = torch.nn.utils.clip_grad_norm_(
|
||||
policy.parameters(), float("inf"), error_if_nonfinite=False
|
||||
)
|
||||
|
||||
# Optimizer step
|
||||
with lock if lock is not None else nullcontext():
|
||||
@@ -167,32 +158,22 @@ def update_policy(
|
||||
optimizer.zero_grad()
|
||||
|
||||
# Step through pytorch scheduler at every batch instead of epoch
|
||||
if lr_scheduler is not None and accelerator.sync_gradients:
|
||||
if lr_scheduler is not None:
|
||||
lr_scheduler.step()
|
||||
|
||||
# Update internal buffers if policy has update method
|
||||
if accelerator.sync_gradients and has_method(
|
||||
accelerator.unwrap_model(policy, keep_fp32_wrapper=True), "update"
|
||||
):
|
||||
if has_method(accelerator.unwrap_model(policy, keep_fp32_wrapper=True), "update"):
|
||||
accelerator.unwrap_model(policy, keep_fp32_wrapper=True).update()
|
||||
|
||||
if accelerator.sync_gradients:
|
||||
train_metrics.lr = optimizer.param_groups[0]["lr"]
|
||||
if torch.cuda.is_available() and accelerator.sync_gradients:
|
||||
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)
|
||||
if grad_norm is not None:
|
||||
train_metrics.accumulate_tensor("grad_norm", grad_norm)
|
||||
if track_update_time:
|
||||
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)
|
||||
return train_metrics, output_dict
|
||||
|
||||
|
||||
@@ -220,7 +201,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
|
||||
require_package("accelerate", extra="training")
|
||||
from accelerate import Accelerator
|
||||
from accelerate.utils import DistributedDataParallelKwargs, DistributedType, InitProcessGroupKwargs
|
||||
from accelerate.utils import DistributedDataParallelKwargs, DistributedType
|
||||
|
||||
cfg.validate()
|
||||
|
||||
@@ -229,16 +210,7 @@ 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:
|
||||
# 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))
|
||||
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
|
||||
# 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"
|
||||
@@ -247,9 +219,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
mixed_precision = {"bfloat16": "bf16", "float16": "fp16", "float32": "no"}.get(policy_dtype)
|
||||
accelerator = Accelerator(
|
||||
step_scheduler_with_optimizer=False,
|
||||
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
|
||||
mixed_precision=mixed_precision,
|
||||
kwargs_handlers=[ddp_kwargs, ipg_kwargs],
|
||||
kwargs_handlers=[ddp_kwargs],
|
||||
cpu=force_cpu,
|
||||
)
|
||||
|
||||
@@ -345,31 +316,19 @@ 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
|
||||
|
||||
dataset_stats = {cfg.rename_map.get(key, key): value for key, value in dataset.meta.stats.items()}
|
||||
processor_kwargs = {}
|
||||
if (processor_pretrained_path and not cfg.resume) or not processor_pretrained_path:
|
||||
processor_kwargs["dataset_stats"] = dataset_stats
|
||||
processor_kwargs["dataset_stats"] = dataset.meta.stats
|
||||
|
||||
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
|
||||
|
||||
if not cfg.is_reward_model_training and processor_pretrained_path is not None:
|
||||
preprocessor_overrides = {
|
||||
"device_processor": {"device": device.type},
|
||||
"normalizer_processor": {
|
||||
"stats": dataset_stats,
|
||||
"stats": dataset.meta.stats,
|
||||
"features": {**policy.config.input_features, **policy.config.output_features},
|
||||
"norm_map": policy.config.normalization_mapping,
|
||||
},
|
||||
@@ -377,7 +336,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
}
|
||||
postprocessor_overrides = {
|
||||
"unnormalizer_processor": {
|
||||
"stats": dataset_stats,
|
||||
"stats": dataset.meta.stats,
|
||||
"features": policy.config.output_features,
|
||||
"norm_map": policy.config.normalization_mapping,
|
||||
},
|
||||
@@ -449,11 +408,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
logging.info(f"{dataset.num_frames=} ({format_big_number(dataset.num_frames)})")
|
||||
logging.info(f"{dataset.num_episodes=}")
|
||||
num_processes = accelerator.num_processes
|
||||
effective_bs = cfg.batch_size * num_processes * cfg.gradient_accumulation_steps
|
||||
logging.info(
|
||||
"Effective batch size: "
|
||||
f"{cfg.batch_size} x {num_processes} x {cfg.gradient_accumulation_steps} = {effective_bs}"
|
||||
)
|
||||
effective_bs = cfg.batch_size * num_processes
|
||||
logging.info(f"Effective batch size: {cfg.batch_size} x {num_processes} = {effective_bs}")
|
||||
logging.info(f"{num_learnable_params=} ({format_big_number(num_learnable_params)})")
|
||||
logging.info(f"{num_total_params=} ({format_big_number(num_total_params)})")
|
||||
|
||||
@@ -464,53 +420,13 @@ 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
|
||||
|
||||
drop_n_first_frames: int | list[int] = 0
|
||||
target_start_feature = cfg.dataset.training_target_start_feature
|
||||
if target_start_feature is not None:
|
||||
if not hasattr(dataset, "hf_dataset"):
|
||||
raise ValueError(
|
||||
"dataset.training_target_start_feature is only supported for a single map-style "
|
||||
"LeRobotDataset."
|
||||
)
|
||||
if target_start_feature not in dataset.hf_dataset.column_names:
|
||||
raise ValueError(
|
||||
f"Training target start feature {target_start_feature!r} is not present in the dataset. "
|
||||
f"Available features: {dataset.hf_dataset.column_names}"
|
||||
)
|
||||
start_column = dataset.hf_dataset.data.column(target_start_feature)
|
||||
absolute_to_relative_idx = dataset.absolute_to_relative_idx
|
||||
episode_indices = dataset.episodes if dataset.episodes is not None else range(len(from_indices))
|
||||
drop_n_first_frames = [0] * len(from_indices)
|
||||
for episode_index in episode_indices:
|
||||
absolute_index = int(from_indices[episode_index])
|
||||
relative_index = (
|
||||
absolute_to_relative_idx[absolute_index]
|
||||
if absolute_to_relative_idx is not None
|
||||
else absolute_index
|
||||
)
|
||||
drop_n_first_frames[episode_index] = int(start_column[relative_index].as_py())
|
||||
logging.info(
|
||||
"Restricting training targets with %s: keeping %d of %d frames",
|
||||
target_start_feature,
|
||||
sum(
|
||||
int(to_indices[index]) - int(from_indices[index]) - drop_n_first_frames[index]
|
||||
for index in episode_indices
|
||||
),
|
||||
sum(int(to_indices[index]) - int(from_indices[index]) for index in episode_indices),
|
||||
)
|
||||
|
||||
sampler = EpisodeAwareSampler(
|
||||
from_indices,
|
||||
to_indices,
|
||||
dataset.meta.episodes["dataset_from_index"],
|
||||
dataset.meta.episodes["dataset_to_index"],
|
||||
episode_indices_to_use=dataset.episodes,
|
||||
drop_n_first_frames=drop_n_first_frames,
|
||||
drop_n_last_frames=getattr(active_cfg, "drop_n_last_frames", 0),
|
||||
shuffle=True,
|
||||
seed=seed,
|
||||
seed=cfg.seed if cfg.seed is not None else 0,
|
||||
absolute_to_relative_idx=dataset.absolute_to_relative_idx,
|
||||
)
|
||||
if cfg.resume and step > 0:
|
||||
@@ -533,12 +449,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
f"batch_size={saved_batch_size}. The data order resumes at the right epoch/offset, "
|
||||
"but per-rank sample-exactness requires the same batch size."
|
||||
)
|
||||
sampler_state = compute_sampler_state(
|
||||
step,
|
||||
len(sampler),
|
||||
ckpt_batch_size * cfg.gradient_accumulation_steps,
|
||||
ckpt_num_processes,
|
||||
)
|
||||
sampler_state = compute_sampler_state(step, len(sampler), ckpt_batch_size, ckpt_num_processes)
|
||||
sampler.load_state_dict(sampler_state)
|
||||
if is_main_process:
|
||||
logging.info(
|
||||
@@ -553,8 +464,6 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
# declares language columns; otherwise stay on PyTorch's default
|
||||
# collate so non-language training runs are unaffected.
|
||||
collate_fn = lerobot_collate_fn if dataset.meta.has_language_columns else None
|
||||
# Allow spawn/forkserver workers where forking large rank processes exhausts memory.
|
||||
mp_context = os.environ.get("LEROBOT_DATALOADER_MP_CONTEXT") or None
|
||||
dataloader = torch.utils.data.DataLoader(
|
||||
dataset,
|
||||
num_workers=cfg.num_workers,
|
||||
@@ -566,7 +475,6 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
collate_fn=collate_fn,
|
||||
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None,
|
||||
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
|
||||
multiprocessing_context=mp_context if cfg.num_workers > 0 else None,
|
||||
)
|
||||
|
||||
# Build eval dataloader if a held-out split exists
|
||||
@@ -635,9 +543,9 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
train_metrics["gpu_mem_gb"] = AverageMeter("mem_gb", ":.2f", reduction="max")
|
||||
|
||||
# Keep global batch size for logging; MetricsTracker handles world size internally.
|
||||
effective_batch_size = cfg.batch_size * accelerator.num_processes * cfg.gradient_accumulation_steps
|
||||
effective_batch_size = cfg.batch_size * accelerator.num_processes
|
||||
train_tracker = MetricsTracker(
|
||||
cfg.batch_size * cfg.gradient_accumulation_steps,
|
||||
cfg.batch_size,
|
||||
dataset.num_frames,
|
||||
dataset.num_episodes,
|
||||
train_metrics,
|
||||
@@ -659,42 +567,24 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
)
|
||||
|
||||
for _ in range(step, cfg.steps):
|
||||
update_start_time = time.perf_counter()
|
||||
dataloading_s = 0.0
|
||||
output_dict = None
|
||||
for microbatch_idx in range(cfg.gradient_accumulation_steps):
|
||||
start_time = time.perf_counter()
|
||||
batch = next(dl_iter)
|
||||
for cam_key in dataset.meta.camera_keys:
|
||||
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
||||
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
||||
if cfg.rename_map:
|
||||
batch = rename_transition_keys(batch, cfg.rename_map)
|
||||
batch = preprocessor(batch)
|
||||
dataloading_s += time.perf_counter() - start_time
|
||||
start_time = time.perf_counter()
|
||||
batch = next(dl_iter)
|
||||
for cam_key in dataset.meta.camera_keys:
|
||||
if cam_key in batch and batch[cam_key].dtype == torch.uint8:
|
||||
batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0
|
||||
batch = preprocessor(batch)
|
||||
train_tracker.dataloading_s = time.perf_counter() - start_time
|
||||
|
||||
# Synchronize GPU metrics only on the final microbatch of logged optimizer updates.
|
||||
log_metrics = (
|
||||
cfg.log_freq > 0
|
||||
and (step + 1) % cfg.log_freq == 0
|
||||
and microbatch_idx == cfg.gradient_accumulation_steps - 1
|
||||
)
|
||||
|
||||
with accelerator.accumulate(policy):
|
||||
train_tracker, output_dict = update_policy(
|
||||
train_tracker,
|
||||
policy,
|
||||
batch,
|
||||
optimizer,
|
||||
cfg.optimizer.grad_clip_norm,
|
||||
accelerator=accelerator,
|
||||
lr_scheduler=lr_scheduler,
|
||||
sample_weighter=sample_weighter,
|
||||
log_metrics=log_metrics,
|
||||
track_update_time=False,
|
||||
)
|
||||
train_tracker.dataloading_s = dataloading_s
|
||||
train_tracker.update_s = time.perf_counter() - update_start_time - dataloading_s
|
||||
train_tracker, _ = update_policy(
|
||||
train_tracker,
|
||||
policy,
|
||||
batch,
|
||||
optimizer,
|
||||
cfg.optimizer.grad_clip_norm,
|
||||
accelerator=accelerator,
|
||||
lr_scheduler=lr_scheduler,
|
||||
sample_weighter=sample_weighter,
|
||||
)
|
||||
|
||||
# Note: eval and checkpoint happens *after* the `step`th training update has completed, so we
|
||||
# increment `step` here.
|
||||
@@ -718,9 +608,10 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
||||
train_tracker.samples_per_s = effective_batch_size / step_time
|
||||
logging.info(train_tracker)
|
||||
if wandb_logger:
|
||||
# Policy sub-losses (latent_loss, action_loss, ...) are aggregated into the
|
||||
# tracker by update_policy, so to_dict() already carries their windowed,
|
||||
# rank-reduced averages — no per-step output_dict passthrough needed.
|
||||
wandb_log_dict = train_tracker.to_dict()
|
||||
if output_dict:
|
||||
wandb_log_dict.update(output_dict)
|
||||
# Log sample weighting statistics if enabled
|
||||
if sample_weighter is not None:
|
||||
weighter_stats = sample_weighter.get_stats()
|
||||
@@ -793,11 +684,10 @@ 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=eval_target_policy,
|
||||
policy=accelerator.unwrap_model(policy),
|
||||
env_preprocessor=env_preprocessor,
|
||||
env_postprocessor=env_postprocessor,
|
||||
preprocessor=preprocessor,
|
||||
|
||||
@@ -22,7 +22,7 @@ from torch.utils.data._utils.collate import default_collate
|
||||
|
||||
from lerobot.datasets.language import LANGUAGE_COLUMNS
|
||||
|
||||
_PYTHON_LIST_KEYS = {"messages", "message_streams", "target_message_indices", *LANGUAGE_COLUMNS}
|
||||
_PYTHON_LIST_KEYS = {"messages", "message_streams", "target_message_indices"}
|
||||
|
||||
|
||||
def lerobot_collate_fn(batch: list[dict[str, Any] | None]) -> dict[str, Any] | None:
|
||||
|
||||
@@ -34,7 +34,6 @@ ACTION = "action"
|
||||
ACTION_PREFIX = ACTION + "."
|
||||
ACTION_TOKENS = ACTION + ".tokens"
|
||||
ACTION_TOKEN_MASK = ACTION + ".token_mask"
|
||||
ACTION_CODE_TOKEN_MASK = ACTION + ".code_token_mask"
|
||||
REWARD = "next.reward"
|
||||
TRUNCATED = "next.truncated"
|
||||
DONE = "next.done"
|
||||
|
||||
@@ -59,6 +59,20 @@ def get_safe_torch_device(try_device: str, log: bool = False) -> torch.device:
|
||||
return device
|
||||
|
||||
|
||||
def resolve_safetensors_device(map_location: str | torch.device) -> str:
|
||||
"""Resolve a device string for a safetensors load, working around a device-mapping quirk.
|
||||
|
||||
safetensors' load maps the bare string "cuda" to cuda:0 regardless of the current device
|
||||
(unlike torch's .to("cuda"), which honors torch.cuda.current_device()). Under multi-GPU
|
||||
accelerate/FSDP every rank would then load its weights onto GPU 0, OOMing it before sharding.
|
||||
Resolve "cuda" to the concrete current-device index so each rank loads onto its own GPU.
|
||||
"""
|
||||
map_location = str(map_location)
|
||||
if map_location == "cuda" and torch.cuda.is_available():
|
||||
return f"cuda:{torch.cuda.current_device()}"
|
||||
return map_location
|
||||
|
||||
|
||||
def get_safe_dtype(dtype: torch.dtype, device: str | torch.device):
|
||||
"""
|
||||
mps is currently not compatible with float64
|
||||
|
||||
@@ -23,46 +23,6 @@ logger = logging.getLogger(__name__)
|
||||
JsonLike = str | int | float | bool | None | list["JsonLike"] | dict[str, "JsonLike"] | tuple["JsonLike", ...]
|
||||
|
||||
|
||||
class StreamingVideoWriter:
|
||||
"""Incrementally encode RGB frames to an MP4 without retaining them in memory."""
|
||||
|
||||
def __init__(self, video_path: str | Path, fps: int) -> None:
|
||||
from .import_utils import require_package
|
||||
|
||||
require_package("av", extra="av-dep")
|
||||
import av
|
||||
|
||||
self._av = av
|
||||
self._container = av.open(str(video_path), mode="w")
|
||||
self._stream = self._container.add_stream("libx264", rate=fps)
|
||||
self._shape: tuple[int, int] | None = None
|
||||
self.frames_written = 0
|
||||
|
||||
def add_frame(self, frame_array) -> None:
|
||||
orig_height, orig_width = frame_array.shape[:2]
|
||||
height = orig_height - orig_height % 2
|
||||
width = orig_width - orig_width % 2
|
||||
if self._shape is None:
|
||||
self._shape = (height, width)
|
||||
self._stream.width = width
|
||||
self._stream.height = height
|
||||
self._stream.pix_fmt = "yuv420p"
|
||||
elif self._shape != (height, width):
|
||||
raise ValueError(f"Video frame shape changed from {self._shape} to {(height, width)}")
|
||||
frame = self._av.VideoFrame.from_ndarray(frame_array[:height, :width], format="rgb24")
|
||||
for packet in self._stream.encode(frame):
|
||||
self._container.mux(packet)
|
||||
self.frames_written += 1
|
||||
|
||||
def close(self) -> None:
|
||||
if self._container is None:
|
||||
return
|
||||
for packet in self._stream.encode():
|
||||
self._container.mux(packet)
|
||||
self._container.close()
|
||||
self._container = None
|
||||
|
||||
|
||||
def load_json(fpath: Path) -> Any:
|
||||
"""Load data from a JSON file.
|
||||
|
||||
@@ -98,12 +58,36 @@ def write_video(video_path: str | Path, stacked_frames: list, fps: int) -> None:
|
||||
stacked_frames: List of HWC uint8 numpy arrays (RGB).
|
||||
fps: Frames per second for the output video.
|
||||
"""
|
||||
writer = StreamingVideoWriter(video_path, fps)
|
||||
try:
|
||||
from .import_utils import require_package
|
||||
|
||||
require_package("av", extra="av-dep")
|
||||
import av
|
||||
|
||||
with av.open(str(video_path), mode="w") as container:
|
||||
orig_height, orig_width = stacked_frames[0].shape[:2]
|
||||
# yuv420p requires even dimensions; crop by one pixel if needed
|
||||
height = orig_height if orig_height % 2 == 0 else orig_height - 1
|
||||
width = orig_width if orig_width % 2 == 0 else orig_width - 1
|
||||
if height != orig_height or width != orig_width:
|
||||
logger.warning(
|
||||
"Frame dimensions %dx%d are not even; cropping to %dx%d for yuv420p compatibility.",
|
||||
orig_width,
|
||||
orig_height,
|
||||
width,
|
||||
height,
|
||||
)
|
||||
stream = container.add_stream("libx264", rate=fps)
|
||||
stream.width = width
|
||||
stream.height = height
|
||||
stream.pix_fmt = "yuv420p"
|
||||
for frame_array in stacked_frames:
|
||||
writer.add_frame(frame_array)
|
||||
finally:
|
||||
writer.close()
|
||||
if height != orig_height or width != orig_width:
|
||||
frame_array = frame_array[:height, :width]
|
||||
frame = av.VideoFrame.from_ndarray(frame_array, format="rgb24")
|
||||
for packet in stream.encode(frame):
|
||||
container.mux(packet)
|
||||
for packet in stream.encode():
|
||||
container.mux(packet)
|
||||
|
||||
|
||||
def deserialize_json_into_object[T: JsonLike](fpath: Path, obj: T) -> T:
|
||||
|
||||
@@ -104,8 +104,7 @@ class MetricsTracker:
|
||||
"episodes",
|
||||
"epochs",
|
||||
"accelerator",
|
||||
"_tensor_sums",
|
||||
"_tensor_counts",
|
||||
"_caller_metrics",
|
||||
]
|
||||
|
||||
def __init__(
|
||||
@@ -131,8 +130,9 @@ class MetricsTracker:
|
||||
self.episodes = self.samples / self._avg_samples_per_ep
|
||||
self.epochs = self.samples / self._num_frames
|
||||
self.accelerator = accelerator
|
||||
self._tensor_sums: dict[str, torch.Tensor] = {}
|
||||
self._tensor_counts: dict[str, int] = {}
|
||||
# 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)
|
||||
|
||||
def __getattr__(self, name: str) -> int | dict[str, AverageMeter] | AverageMeter | Any:
|
||||
if name in self.__dict__:
|
||||
@@ -160,21 +160,20 @@ 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 update_metrics(self, values: dict[str, Any]) -> None:
|
||||
"""Accumulate a dict of scalar metrics, auto-registering a meter for each new key.
|
||||
|
||||
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()
|
||||
Non-numeric values and bools are ignored.
|
||||
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)):
|
||||
continue
|
||||
if name in self._caller_metrics:
|
||||
continue
|
||||
if name not in self.metrics:
|
||||
self.metrics[name] = AverageMeter(name, ":.3f", reduction="mean")
|
||||
self.metrics[name].update(float(value))
|
||||
|
||||
def reduce_across_ranks(self) -> None:
|
||||
"""
|
||||
@@ -196,21 +195,10 @@ class MetricsTracker:
|
||||
if not buckets:
|
||||
return
|
||||
|
||||
# NB: don't use ``accelerator.reduce(..., reduction="max")`` — accelerate only implements
|
||||
# "sum"/"mean" (it always all-reduces with SUM and divides for "mean"), so "max" silently
|
||||
# returns the SUM across ranks, inflating every "max" metric by ``num_processes`` (e.g. a
|
||||
# 3.5s step reported as 28s on 8 GPUs). Gather per-rank values and reduce them explicitly.
|
||||
device = self.accelerator.device
|
||||
num_processes = self.accelerator.num_processes
|
||||
for reduction, names in buckets.items():
|
||||
local = torch.tensor([self.metrics[n].avg for n in names], dtype=torch.float32, device=device)
|
||||
gathered = self.accelerator.gather(local).view(num_processes, len(names))
|
||||
if reduction == "max":
|
||||
reduced = gathered.amax(dim=0)
|
||||
elif reduction == "sum":
|
||||
reduced = gathered.sum(dim=0)
|
||||
else: # "mean"
|
||||
reduced = gathered.mean(dim=0)
|
||||
tensor = torch.tensor([self.metrics[n].avg for n in names], dtype=torch.float32, device=device)
|
||||
reduced = self.accelerator.reduce(tensor, reduction=reduction)
|
||||
for name, value in zip(names, reduced.tolist(), strict=True):
|
||||
meter = self.metrics[name]
|
||||
# Preserve avg == sum / count so a later .update() on this meter accumulates
|
||||
@@ -247,5 +235,3 @@ class MetricsTracker:
|
||||
"""Resets average meters."""
|
||||
for m in self.metrics.values():
|
||||
m.reset()
|
||||
self._tensor_sums.clear()
|
||||
self._tensor_counts.clear()
|
||||
|
||||
@@ -38,10 +38,7 @@ def _is_scalar(x):
|
||||
|
||||
|
||||
def init_rerun(
|
||||
session_name: str = "lerobot_control_loop",
|
||||
ip: str | None = None,
|
||||
port: int | None = None,
|
||||
web_port: int | None = None,
|
||||
session_name: str = "lerobot_control_loop", ip: str | None = None, port: int | None = None
|
||||
) -> None:
|
||||
"""
|
||||
Initializes the Rerun SDK for visualizing the control loop.
|
||||
@@ -50,7 +47,6 @@ def init_rerun(
|
||||
session_name: Name of the Rerun session.
|
||||
ip: Optional IP for connecting to a Rerun server.
|
||||
port: Optional port for connecting to a Rerun server.
|
||||
web_port: Serve a headless web viewer on this port, using ``port`` for gRPC.
|
||||
"""
|
||||
|
||||
require_package("rerun-sdk", extra="viz", import_name="rerun")
|
||||
@@ -64,10 +60,6 @@ def init_rerun(
|
||||
memory_limit = os.getenv("LEROBOT_RERUN_MEMORY_LIMIT", "10%")
|
||||
if ip and port:
|
||||
rr.connect_grpc(url=f"rerun+http://{ip}:{port}/proxy")
|
||||
elif web_port is not None:
|
||||
grpc_port = port or 9876
|
||||
url = rr.serve_grpc(grpc_port=grpc_port)
|
||||
rr.serve_web_viewer(web_port=web_port, open_browser=False, connect_to=url)
|
||||
else:
|
||||
rr.spawn(memory_limit=memory_limit)
|
||||
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
# 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.
|
||||
|
||||
"""Best-effort text overlays shared by evaluation and interactive rollouts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def annotate_frame(frame: np.ndarray, fields: Iterable[tuple[str, str | None]]) -> np.ndarray:
|
||||
"""Return an RGB frame annotated with the non-empty labeled ``fields``."""
|
||||
if frame.ndim != 3 or frame.shape[-1] != 3:
|
||||
return frame
|
||||
try:
|
||||
import cv2 # noqa: PLC0415
|
||||
except ImportError:
|
||||
return frame
|
||||
|
||||
text_rows = [f"{label}: {value}" for label, value in fields if value]
|
||||
if not text_rows:
|
||||
return frame
|
||||
|
||||
image = np.ascontiguousarray(frame).copy()
|
||||
font, scale, thickness, margin = cv2.FONT_HERSHEY_SIMPLEX, 0.5, 1, 6
|
||||
max_width = image.shape[1] - 2 * margin
|
||||
lines: list[str] = []
|
||||
for text in text_rows:
|
||||
current = ""
|
||||
for word in text.split():
|
||||
candidate = f"{current} {word}".strip()
|
||||
width = cv2.getTextSize(candidate, font, scale, thickness)[0][0]
|
||||
if width > max_width and current:
|
||||
lines.append(current)
|
||||
current = word
|
||||
else:
|
||||
current = candidate
|
||||
if current:
|
||||
lines.append(current)
|
||||
|
||||
line_height = 20
|
||||
header_height = min(image.shape[0], len(lines) * line_height + 6)
|
||||
backdrop = image.copy()
|
||||
cv2.rectangle(backdrop, (0, 0), (image.shape[1], header_height), (0, 0, 0), -1)
|
||||
cv2.addWeighted(backdrop, 0.55, image, 0.45, 0, dst=image)
|
||||
|
||||
for index, line in enumerate(lines):
|
||||
cv2.putText(
|
||||
image,
|
||||
line,
|
||||
(margin, 18 + index * line_height),
|
||||
font,
|
||||
scale,
|
||||
(255, 255, 255),
|
||||
thickness,
|
||||
cv2.LINE_AA,
|
||||
)
|
||||
return image
|
||||
@@ -29,13 +29,6 @@ def test_message_recipe_validates_unknown_binding():
|
||||
)
|
||||
|
||||
|
||||
def test_canonical_recipe_loads():
|
||||
"""The canonical PI052 blend YAML loads + validates."""
|
||||
recipe = TrainingRecipe.from_yaml(Path("src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml"))
|
||||
assert recipe.blend is not None
|
||||
assert sum(c.weight for c in recipe.blend.values()) == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_message_turn_requires_a_stream():
|
||||
"""Every turn must declare a stream — None is rejected at construction.
|
||||
|
||||
|
||||
@@ -343,84 +343,6 @@ def test_resolve_task_explicit_override_beats_rephrasings():
|
||||
assert rendered["messages"][0]["content"] == "explicit override wins"
|
||||
|
||||
|
||||
def test_flow_only_low_level_recipe_renders_without_target():
|
||||
"""Regression: a flow-only ``low_level`` recipe has no ``target`` turn —
|
||||
its supervision is the action-expert flow loss, not text-CE. It must
|
||||
still render (not ``None``), otherwise every blend draw of it is dropped
|
||||
and the action expert never receives a flow loss."""
|
||||
recipe = TrainingRecipe(
|
||||
messages=[
|
||||
MessageTurn(
|
||||
role="user",
|
||||
content="${subtask}",
|
||||
stream="low_level",
|
||||
if_present="subtask",
|
||||
),
|
||||
],
|
||||
bindings={"subtask": "active_at(t, style=subtask)"},
|
||||
)
|
||||
|
||||
rendered = render_sample(
|
||||
recipe=recipe,
|
||||
persistent=PERSISTENT,
|
||||
events=[],
|
||||
t=0.5,
|
||||
sample_idx=0,
|
||||
task="clean kitchen",
|
||||
)
|
||||
|
||||
assert rendered is not None
|
||||
assert rendered["messages"] == [{"role": "user", "content": "subtask 0"}]
|
||||
assert rendered["message_streams"] == ["low_level"]
|
||||
assert rendered["target_message_indices"] == []
|
||||
|
||||
|
||||
def test_vqa_frame_is_consumed_over_the_weighted_blend():
|
||||
"""A frame carrying a VQA annotation renders the ``ask_vqa*`` sub-recipe
|
||||
even when its blend weight is tiny — VQA annotations are sparse and must
|
||||
never be wasted on a subtask/action draw."""
|
||||
recipe = TrainingRecipe(
|
||||
blend={
|
||||
"high_level_subtask": TrainingRecipe(
|
||||
weight=0.99,
|
||||
messages=[
|
||||
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||
MessageTurn(role="assistant", content="a subtask", stream="high_level", target=True),
|
||||
],
|
||||
),
|
||||
"ask_vqa_top": TrainingRecipe(
|
||||
weight=0.01,
|
||||
bindings={
|
||||
"vqa_query": "emitted_at(t, style=vqa, role=user, camera=observation.images.top)",
|
||||
"vqa": "emitted_at(t, style=vqa, role=assistant, camera=observation.images.top)",
|
||||
},
|
||||
messages=[
|
||||
MessageTurn(
|
||||
role="user", content="${vqa_query}", stream="high_level", if_present="vqa_query"
|
||||
),
|
||||
MessageTurn(
|
||||
role="assistant",
|
||||
content="${vqa}",
|
||||
stream="high_level",
|
||||
target=True,
|
||||
if_present="vqa",
|
||||
),
|
||||
],
|
||||
),
|
||||
}
|
||||
)
|
||||
# A frame WITH a vqa event renders VQA on every sample_idx, despite the
|
||||
# ask_vqa weight being only 0.01.
|
||||
for sample_idx in range(20):
|
||||
rendered = render_sample(
|
||||
recipe=recipe, persistent=PERSISTENT, events=EVENTS_AT_1, t=1.0, sample_idx=sample_idx, task="x"
|
||||
)
|
||||
assert rendered["messages"][-1]["content"] == '{"count": 2}', sample_idx
|
||||
# A frame WITHOUT a vqa event falls back to the normal weighted blend.
|
||||
rendered = render_sample(recipe=recipe, persistent=PERSISTENT, events=[], t=1.0, sample_idx=0, task="x")
|
||||
assert rendered["messages"][-1]["content"] == "a subtask"
|
||||
|
||||
|
||||
def test_emitted_at_persistent_tolerates_small_timestamp_drift():
|
||||
"""Persistent ``emitted_at`` should match within EMITTED_AT_TOLERANCE_S
|
||||
so callers that derive ``t`` arithmetically (``frame_idx / fps``) still
|
||||
|
||||
@@ -25,7 +25,7 @@ from datasets import Dataset # noqa: E402
|
||||
from lerobot.datasets.io_utils import (
|
||||
hf_transform_to_torch,
|
||||
)
|
||||
from lerobot.datasets.sampler import EpisodeAwareSampler, compute_sampler_state
|
||||
from lerobot.datasets.sampler import EpisodeAwareSampler
|
||||
|
||||
|
||||
def calculate_episode_data_index(hf_dataset: Dataset) -> dict[str, torch.Tensor]:
|
||||
@@ -60,26 +60,6 @@ def test_drop_n_first_frames():
|
||||
assert list(sampler) == [1, 4, 5]
|
||||
|
||||
|
||||
def test_drop_different_number_of_first_frames_per_episode():
|
||||
sampler = EpisodeAwareSampler(
|
||||
dataset_from_indices=[0, 3, 5],
|
||||
dataset_to_indices=[3, 5, 9],
|
||||
drop_n_first_frames=[1, 0, 3],
|
||||
)
|
||||
assert sampler.indices == [1, 2, 3, 4, 8]
|
||||
assert len(sampler) == 5
|
||||
|
||||
|
||||
def test_drop_first_frames_sequence_must_match_episode_count():
|
||||
with pytest.raises(ValueError, match="one value per episode"):
|
||||
EpisodeAwareSampler([0, 3], [3, 6], drop_n_first_frames=[1])
|
||||
|
||||
|
||||
def test_drop_first_frames_sequence_must_be_non_negative():
|
||||
with pytest.raises(ValueError, match="must be >= 0"):
|
||||
EpisodeAwareSampler([0, 3], [3, 6], drop_n_first_frames=[1, -1])
|
||||
|
||||
|
||||
def test_drop_n_last_frames():
|
||||
dataset = Dataset.from_dict(
|
||||
{
|
||||
@@ -174,6 +154,8 @@ def test_partial_episode_drop_warns(caplog):
|
||||
|
||||
# --- seeded (seed, epoch) shuffling, resume, and state ---
|
||||
|
||||
from lerobot.datasets.sampler import compute_sampler_state # noqa: E402
|
||||
|
||||
EPISODE_BOUNDS = ([0, 2, 3], [2, 3, 6]) # episodes of 2, 1 and 3 frames
|
||||
|
||||
|
||||
|
||||
@@ -1,151 +0,0 @@
|
||||
#!/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
|
||||
@@ -1,146 +0,0 @@
|
||||
#!/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
|
||||
@@ -1,106 +0,0 @@
|
||||
#!/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."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.nn import functional as F # noqa: N812
|
||||
|
||||
pytest.importorskip("transformers")
|
||||
pytest.importorskip("liger_kernel")
|
||||
|
||||
from lerobot.policies.pi052.modeling_pi052 import _fast_lin_ce # 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
|
||||
@@ -1,65 +0,0 @@
|
||||
#!/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
|
||||
@@ -1,27 +0,0 @@
|
||||
# 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)
|
||||
@@ -1,85 +0,0 @@
|
||||
# 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. <say>heading to the sink</say>"
|
||||
|
||||
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]
|
||||
@@ -1,147 +0,0 @@
|
||||
#!/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)
|
||||
@@ -1,223 +0,0 @@
|
||||
#!/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
|
||||
``<say>...</say>`` 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.\n<say>On it!</say>"
|
||||
|
||||
|
||||
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"] == "<say>hi</say>"
|
||||
|
||||
|
||||
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\n<say>hello there</say>"
|
||||
|
||||
|
||||
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>")
|
||||
# EOS is appended to the supervised target (assistant) turn only.
|
||||
assert prompt == "User: pick cube\nAssistant: move to cube<eos>\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<eos>"
|
||||
|
||||
|
||||
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()
|
||||
@@ -1,141 +0,0 @@
|
||||
#!/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)
|
||||
@@ -1,98 +0,0 @@
|
||||
#!/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,
|
||||
}
|
||||
|
||||
|
||||
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_each_node_uses_its_local_rank_zero_as_fit_leader(monkeypatch):
|
||||
monkeypatch.setenv("RANK", "8")
|
||||
monkeypatch.setenv("LOCAL_RANK", "0")
|
||||
assert fit_module._is_local_leader()
|
||||
|
||||
monkeypatch.setenv("LOCAL_RANK", "1")
|
||||
assert not fit_module._is_local_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: "/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"}
|
||||
}
|
||||
@@ -16,12 +16,8 @@
|
||||
|
||||
"""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")
|
||||
|
||||
@@ -35,26 +31,6 @@ 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():
|
||||
|
||||
@@ -1,210 +0,0 @@
|
||||
# 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
|
||||
|
||||
from lerobot.datasets.factory import resolve_delta_timestamps
|
||||
from lerobot.policies.smolvla.configuration_smolvla import SmolVLAConfig
|
||||
from lerobot.policies.smolvla.visual_memory import (
|
||||
causal_temporal_mask,
|
||||
encode_video_with_mem,
|
||||
sample_visual_history,
|
||||
temporal_sinusoidal_embedding,
|
||||
)
|
||||
|
||||
|
||||
def test_visual_memory_observation_delta_indices():
|
||||
baseline = SmolVLAConfig()
|
||||
memory = SmolVLAConfig(use_visual_memory=True, visual_memory_frames=6, visual_memory_stride=10)
|
||||
|
||||
assert baseline.observation_delta_indices == [0]
|
||||
assert memory.observation_delta_indices == [-50, -40, -30, -20, -10, 0]
|
||||
|
||||
|
||||
def test_delta_timestamps_respect_raw_dataset_rename_map():
|
||||
class RawMetadata:
|
||||
fps = 10
|
||||
features = {"image": {}, "state": {}, "actions": {}}
|
||||
|
||||
config = SmolVLAConfig(use_visual_memory=True, visual_memory_frames=3, visual_memory_stride=5)
|
||||
delta_timestamps = resolve_delta_timestamps(
|
||||
config,
|
||||
RawMetadata(),
|
||||
{
|
||||
"image": "observation.images.camera1",
|
||||
"state": "observation.state",
|
||||
"actions": "action",
|
||||
},
|
||||
)
|
||||
|
||||
assert delta_timestamps["image"] == [-1.0, -0.5, 0.0]
|
||||
assert delta_timestamps["state"] == [-1.0, -0.5, 0.0]
|
||||
assert delta_timestamps["actions"] == [index / 10 for index in range(50)]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("visual_memory_frames", 0),
|
||||
("visual_memory_stride", 0),
|
||||
("visual_memory_temporal_attention_every", 0),
|
||||
],
|
||||
)
|
||||
def test_visual_memory_config_rejects_non_positive_values(field, value):
|
||||
with pytest.raises(ValueError, match=field):
|
||||
SmolVLAConfig(**{field: value})
|
||||
|
||||
|
||||
def test_current_temporal_position_is_exactly_zero():
|
||||
embedding = temporal_sinusoidal_embedding(4, 16, device=torch.device("cpu"), dtype=torch.float32)
|
||||
|
||||
torch.testing.assert_close(embedding[-1], torch.zeros(16))
|
||||
assert torch.count_nonzero(embedding[:-1]) > 0
|
||||
|
||||
|
||||
def test_causal_temporal_mask_combines_causality_and_padding():
|
||||
frame_mask = torch.tensor([[False, True, True]])
|
||||
mask = causal_temporal_mask(frame_mask, dtype=torch.float32, num_patches=2)
|
||||
|
||||
assert mask.shape == (2, 1, 3, 3)
|
||||
assert mask[0, 0, 1, 0] < -1e30
|
||||
assert mask[0, 0, 1, 1] == 0
|
||||
assert mask[0, 0, 1, 2] < -1e30
|
||||
assert mask[0, 0, 2, 1] == 0
|
||||
|
||||
|
||||
def test_inference_history_matches_training_order_and_padding():
|
||||
history = [torch.full((2, 1), value) for value in range(11)]
|
||||
|
||||
initial_video, initial_padding = sample_visual_history(history, num_frames=3, stride=5, steps_seen=1)
|
||||
full_video, full_padding = sample_visual_history(history, num_frames=3, stride=5, steps_seen=11)
|
||||
|
||||
assert initial_video[:, :, 0].tolist() == [[0, 5, 10], [0, 5, 10]]
|
||||
assert initial_padding.tolist() == [[True, True, False], [True, True, False]]
|
||||
torch.testing.assert_close(full_video[:, :, 0], torch.tensor([[0, 5, 10], [0, 5, 10]]))
|
||||
assert not full_padding.any()
|
||||
|
||||
|
||||
def test_single_frame_mem_matches_original_siglip_encoder():
|
||||
transformers = pytest.importorskip("transformers")
|
||||
config = transformers.SiglipVisionConfig(
|
||||
image_size=16,
|
||||
patch_size=8,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=4,
|
||||
vision_use_head=False,
|
||||
)
|
||||
from transformers.models.siglip.modeling_siglip import SiglipVisionTransformer
|
||||
|
||||
model = SiglipVisionTransformer(config).eval()
|
||||
image = torch.randn(2, 3, 16, 16)
|
||||
|
||||
expected = model(image).last_hidden_state
|
||||
actual = encode_video_with_mem(
|
||||
model,
|
||||
image[:, None],
|
||||
torch.ones(2, 1, dtype=torch.bool),
|
||||
temporal_attention_every=4,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(actual, expected)
|
||||
|
||||
|
||||
def test_mem_video_encoder_compresses_time_without_new_parameters():
|
||||
transformers = pytest.importorskip("transformers")
|
||||
config = transformers.SiglipVisionConfig(
|
||||
image_size=16,
|
||||
patch_size=8,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=4,
|
||||
vision_use_head=False,
|
||||
)
|
||||
from transformers.models.siglip.modeling_siglip import SiglipVisionTransformer
|
||||
|
||||
model = SiglipVisionTransformer(config)
|
||||
parameter_ids = {id(parameter) for parameter in model.parameters()}
|
||||
video = torch.randn(2, 3, 3, 16, 16, requires_grad=True)
|
||||
|
||||
output = encode_video_with_mem(
|
||||
model,
|
||||
video,
|
||||
torch.ones(2, 3, dtype=torch.bool),
|
||||
temporal_attention_every=4,
|
||||
)
|
||||
output.sum().backward()
|
||||
|
||||
assert output.shape == (2, 4, 16)
|
||||
assert {id(parameter) for parameter in model.parameters()} == parameter_ids
|
||||
assert video.grad is not None
|
||||
|
||||
|
||||
def test_mem_video_encoder_supports_smolvlm_vision_tower():
|
||||
transformers = pytest.importorskip("transformers")
|
||||
config = transformers.SmolVLMVisionConfig(
|
||||
image_size=16,
|
||||
patch_size=8,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=4,
|
||||
)
|
||||
from transformers.models.smolvlm.modeling_smolvlm import SmolVLMVisionTransformer
|
||||
|
||||
model = SmolVLMVisionTransformer(config).eval()
|
||||
video = torch.randn(2, 3, 3, 16, 16)
|
||||
|
||||
output = encode_video_with_mem(
|
||||
model,
|
||||
video,
|
||||
torch.ones(2, 3, dtype=torch.bool),
|
||||
temporal_attention_every=4,
|
||||
)
|
||||
single_frame = encode_video_with_mem(
|
||||
model,
|
||||
video[:, -1:],
|
||||
torch.ones(2, 1, dtype=torch.bool),
|
||||
temporal_attention_every=4,
|
||||
)
|
||||
|
||||
assert output.shape == (2, 4, 16)
|
||||
torch.testing.assert_close(single_frame, model(video[:, -1]).last_hidden_state)
|
||||
|
||||
|
||||
def test_masked_history_cannot_change_current_embedding():
|
||||
transformers = pytest.importorskip("transformers")
|
||||
config = transformers.SmolVLMVisionConfig(
|
||||
image_size=16,
|
||||
patch_size=8,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=4,
|
||||
)
|
||||
from transformers.models.smolvlm.modeling_smolvlm import SmolVLMVisionTransformer
|
||||
|
||||
model = SmolVLMVisionTransformer(config).eval()
|
||||
first_video = torch.randn(1, 3, 3, 16, 16)
|
||||
second_video = first_video.clone()
|
||||
second_video[:, :2] = torch.randn_like(second_video[:, :2]) * 100
|
||||
frame_mask = torch.tensor([[False, False, True]])
|
||||
|
||||
first_output = encode_video_with_mem(model, first_video, frame_mask, temporal_attention_every=4)
|
||||
second_output = encode_video_with_mem(model, second_video, frame_mask, temporal_attention_every=4)
|
||||
|
||||
torch.testing.assert_close(first_output, second_output)
|
||||
@@ -27,7 +27,7 @@ from lerobot.processor import (
|
||||
TransitionKey,
|
||||
)
|
||||
from lerobot.processor.converters import create_transition, identity_transition
|
||||
from lerobot.processor.rename_processor import rename_stats, rename_transition_keys
|
||||
from lerobot.processor.rename_processor import rename_stats
|
||||
from lerobot.utils.constants import ACTION, OBS_IMAGE, OBS_IMAGES, OBS_STATE
|
||||
from tests.conftest import assert_contract_is_typed
|
||||
|
||||
@@ -64,22 +64,6 @@ def test_basic_renaming():
|
||||
assert processed_obs["unchanged_key"] == "keep_me"
|
||||
|
||||
|
||||
def test_renaming_preserves_feature_suffixes_for_sampling_metadata():
|
||||
data = {
|
||||
"image": torch.zeros(1),
|
||||
"image_is_pad": torch.ones(1, dtype=torch.bool),
|
||||
"image_padding_mask": torch.ones(1, dtype=torch.bool),
|
||||
}
|
||||
|
||||
result = rename_transition_keys(data, {"image": "observation.images.camera1"})
|
||||
|
||||
assert set(result) == {
|
||||
"observation.images.camera1",
|
||||
"observation.images.camera1_is_pad",
|
||||
"observation.images.camera1_padding_mask",
|
||||
}
|
||||
|
||||
|
||||
def test_empty_rename_map():
|
||||
"""Test processor with empty rename map (should pass through unchanged)."""
|
||||
processor = RenameObservationsProcessorStep(rename_map={})
|
||||
|
||||
@@ -12,9 +12,7 @@ from lerobot.processor.render_messages_processor import RenderMessagesStep # no
|
||||
from lerobot.types import TransitionKey # noqa: E402
|
||||
|
||||
|
||||
def test_render_messages_step_renders_task_fallback_without_language_columns():
|
||||
"""No language columns + a task string → low-level task fallback render,
|
||||
matching what the policy sees at eval time on unannotated observations."""
|
||||
def test_render_messages_step_noops_without_language_columns():
|
||||
recipe = TrainingRecipe(
|
||||
messages=[
|
||||
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||
@@ -23,24 +21,6 @@ def test_render_messages_step_renders_task_fallback_without_language_columns():
|
||||
)
|
||||
transition = create_transition(complementary_data={"task": "do it"})
|
||||
|
||||
out = RenderMessagesStep(recipe)(transition)
|
||||
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||
|
||||
assert data["messages"] == [{"role": "user", "content": "do it"}]
|
||||
assert data["message_streams"] == ["low_level"]
|
||||
assert data["target_message_indices"] == []
|
||||
assert data["task"] == "do it"
|
||||
|
||||
|
||||
def test_render_messages_step_noops_without_language_columns_or_task():
|
||||
recipe = TrainingRecipe(
|
||||
messages=[
|
||||
MessageTurn(role="user", content="${task}", stream="high_level"),
|
||||
MessageTurn(role="assistant", content="${subtask}", stream="low_level", target=True),
|
||||
]
|
||||
)
|
||||
transition = create_transition(complementary_data={})
|
||||
|
||||
assert RenderMessagesStep(recipe)(transition) == transition
|
||||
|
||||
|
||||
@@ -78,70 +58,3 @@ def test_render_messages_step_renders_and_drops_raw_language():
|
||||
assert data["messages"][-1]["content"] == "reach carefully"
|
||||
assert data["message_streams"] == ["high_level", "low_level"]
|
||||
assert data["target_message_indices"] == [1]
|
||||
|
||||
|
||||
def test_render_messages_step_falls_back_to_low_level_task_when_recipe_misses():
|
||||
recipe = TrainingRecipe(
|
||||
messages=[
|
||||
MessageTurn(
|
||||
role="assistant",
|
||||
content="${subtask}",
|
||||
stream="high_level",
|
||||
target=True,
|
||||
if_present="subtask",
|
||||
),
|
||||
]
|
||||
)
|
||||
transition = create_transition(
|
||||
complementary_data={
|
||||
"task": "pick the cube",
|
||||
"timestamp": torch.tensor(0.0),
|
||||
"index": torch.tensor(7),
|
||||
"language_persistent": [],
|
||||
"language_events": [{"style": "unmatched", "timestamp": 0.0}],
|
||||
}
|
||||
)
|
||||
|
||||
out = RenderMessagesStep(recipe)(transition)
|
||||
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||
|
||||
assert data["messages"] == [{"role": "user", "content": "pick the cube"}]
|
||||
assert data["message_streams"] == ["low_level"]
|
||||
assert data["target_message_indices"] == []
|
||||
|
||||
|
||||
def test_render_messages_step_falls_back_per_sample_in_batched_language():
|
||||
recipe = TrainingRecipe(
|
||||
messages=[
|
||||
MessageTurn(
|
||||
role="assistant",
|
||||
content="${subtask}",
|
||||
stream="high_level",
|
||||
target=True,
|
||||
if_present="subtask",
|
||||
),
|
||||
]
|
||||
)
|
||||
transition = create_transition(
|
||||
action=torch.arange(4).reshape(2, 2),
|
||||
complementary_data={
|
||||
"task": ["pick the cube", "open the drawer"],
|
||||
"timestamp": torch.tensor([0.0, 1.0]),
|
||||
"index": torch.tensor([7, 8]),
|
||||
"language_persistent": [[], []],
|
||||
"language_events": [
|
||||
[{"style": "unmatched", "timestamp": 0.0}],
|
||||
[{"style": "unmatched", "timestamp": 1.0}],
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
out = RenderMessagesStep(recipe)(transition)
|
||||
data = out[TransitionKey.COMPLEMENTARY_DATA]
|
||||
|
||||
assert data["messages"] == [
|
||||
[{"role": "user", "content": "pick the cube"}],
|
||||
[{"role": "user", "content": "open the drawer"}],
|
||||
]
|
||||
assert data["message_streams"] == [["low_level"], ["low_level"]]
|
||||
assert data["target_message_indices"] == [[], []]
|
||||
|
||||
@@ -25,7 +25,7 @@ import pytest
|
||||
import torch
|
||||
|
||||
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
|
||||
from lerobot.processor import ActionTokenizerProcessorStep, DataProcessorPipeline, TokenizerProcessorStep
|
||||
from lerobot.processor import DataProcessorPipeline, TokenizerProcessorStep
|
||||
from lerobot.processor.converters import create_transition, identity_transition
|
||||
from lerobot.types import TransitionKey
|
||||
from lerobot.utils.constants import (
|
||||
@@ -88,24 +88,6 @@ class MockTokenizer:
|
||||
return result
|
||||
|
||||
|
||||
def test_action_tokenizer_config_preserves_token_mapping():
|
||||
processor = object.__new__(ActionTokenizerProcessorStep)
|
||||
processor.trust_remote_code = True
|
||||
processor.max_action_tokens = 384
|
||||
processor.fast_skip_tokens = 64
|
||||
processor.paligemma_tokenizer_name = "custom/paligemma"
|
||||
processor.action_tokenizer_name = "custom/fast"
|
||||
processor.action_tokenizer_input_object = None
|
||||
|
||||
assert processor.get_config() == {
|
||||
"trust_remote_code": True,
|
||||
"max_action_tokens": 384,
|
||||
"fast_skip_tokens": 64,
|
||||
"paligemma_tokenizer_name": "custom/paligemma",
|
||||
"action_tokenizer_name": "custom/fast",
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_tokenizer():
|
||||
"""Provide a mock tokenizer for testing."""
|
||||
|
||||
@@ -1,105 +0,0 @@
|
||||
# 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 lerobot.runtime import RuntimeState
|
||||
from lerobot.runtime.adapter import (
|
||||
BaseLanguageAdapter,
|
||||
DirectTaskPolicyAdapter,
|
||||
GenerationConfig,
|
||||
)
|
||||
|
||||
|
||||
class ScriptedAdapter(BaseLanguageAdapter):
|
||||
"""Base adapter whose text generation returns queued strings per kind."""
|
||||
|
||||
def __init__(self, scripts, gen=None):
|
||||
super().__init__(policy=object(), gen=gen)
|
||||
self.scripts = {k: list(v) for k, v in scripts.items()}
|
||||
self.calls = []
|
||||
|
||||
def select_action(self, observation, state):
|
||||
return None
|
||||
|
||||
def generate_text(self, kind, observation, state, user_text=None):
|
||||
self.calls.append(kind)
|
||||
queue = self.scripts.get(kind, [])
|
||||
return queue.pop(0) if queue else ""
|
||||
|
||||
|
||||
def test_cascade_sets_subtask_then_memory():
|
||||
adapter = ScriptedAdapter({"subtask": ["pick the red cup"], "memory": ["the cup is grasped"]})
|
||||
state = RuntimeState(task="clean")
|
||||
|
||||
adapter.update_language_state(None, state)
|
||||
|
||||
assert state.language_context["subtask"] == "pick the red cup"
|
||||
assert state.language_context["memory"] == "the cup is grasped"
|
||||
assert adapter.calls == ["subtask", "memory"]
|
||||
|
||||
|
||||
def test_nonempty_generation_is_used_verbatim():
|
||||
adapter = ScriptedAdapter({"subtask": [":::: ::"], "memory": ["memory"]})
|
||||
state = RuntimeState(task="clean")
|
||||
|
||||
adapter.update_language_state(None, state)
|
||||
|
||||
assert state.language_context["subtask"] == ":::: ::"
|
||||
assert state.language_context["memory"] == "memory"
|
||||
assert adapter.calls == ["subtask", "memory"]
|
||||
|
||||
|
||||
def test_throttle_regenerates_every_n_chunks():
|
||||
adapter = ScriptedAdapter(
|
||||
{
|
||||
"subtask": ["pick the first cup", "pick the second cup"],
|
||||
"memory": ["memory one two three", "memory four five six"],
|
||||
},
|
||||
gen=GenerationConfig(chunks_per_regen=2),
|
||||
)
|
||||
state = RuntimeState(task="clean")
|
||||
|
||||
adapter.update_language_state(None, state) # generates
|
||||
assert state.language_context["subtask"] == "pick the first cup"
|
||||
adapter.update_language_state(None, state) # throttled — no generation
|
||||
assert state.language_context["subtask"] == "pick the first cup"
|
||||
adapter.update_language_state(None, state) # generates again
|
||||
assert state.language_context["subtask"] == "pick the second cup"
|
||||
|
||||
|
||||
def test_handle_interjection_sets_plan_and_strips_say():
|
||||
adapter = ScriptedAdapter({"interjection": ["turn to the left now <say>heading left</say>"]})
|
||||
state = RuntimeState(task="clean")
|
||||
|
||||
adapter.handle_interjection("turn", None, state)
|
||||
|
||||
assert state.language_context["plan"] == "turn to the left now"
|
||||
|
||||
|
||||
def test_direct_task_adapter_delegates_action_chunk():
|
||||
class Policy:
|
||||
def predict_action_chunk(self, observation):
|
||||
return ("chunk", observation)
|
||||
|
||||
observation = {"task": "pick up the cube"}
|
||||
adapter = DirectTaskPolicyAdapter(Policy())
|
||||
|
||||
assert adapter.select_action(observation, RuntimeState()) == ("chunk", observation)
|
||||
assert adapter.generate_text("subtask", observation, RuntimeState()) == ""
|
||||
|
||||
|
||||
def test_flat_policy_registry_reuses_direct_task_adapter():
|
||||
from lerobot.runtime.registry import get_language_adapter_factory
|
||||
|
||||
assert get_language_adapter_factory("pi05") is DirectTaskPolicyAdapter
|
||||
assert get_language_adapter_factory("molmoact2") is DirectTaskPolicyAdapter
|
||||
@@ -1,75 +0,0 @@
|
||||
# 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 unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from lerobot.runtime.cli import _build_rollout_runtime_io, _parse_args
|
||||
|
||||
|
||||
def test_parse_args_preserves_rollout_robot_overrides():
|
||||
args = _parse_args(
|
||||
[
|
||||
"--policy.path=checkpoint",
|
||||
"--robot.type=so101_follower",
|
||||
"--robot.calibration_dir=/tmp/calibration",
|
||||
]
|
||||
)
|
||||
|
||||
assert args.robot_type == "so101_follower"
|
||||
assert "--robot.calibration_dir=/tmp/calibration" in args.raw_argv
|
||||
|
||||
|
||||
def test_parse_args_rejects_removed_dataset_replay_flags():
|
||||
with pytest.raises(SystemExit):
|
||||
_parse_args(["--policy.path=checkpoint", "--dataset.repo_id=dataset"])
|
||||
|
||||
|
||||
def test_rollout_runtime_io_uses_context_processors():
|
||||
robot = MagicMock()
|
||||
robot.robot_type = "mock_robot"
|
||||
robot.cameras = {}
|
||||
robot.get_observation.return_value = {"joint.pos": 1.5}
|
||||
ctx = SimpleNamespace(
|
||||
hardware=SimpleNamespace(robot_wrapper=robot),
|
||||
runtime=SimpleNamespace(cfg=SimpleNamespace(device="cpu")),
|
||||
processors=SimpleNamespace(
|
||||
robot_observation_processor=lambda observation: observation,
|
||||
robot_action_processor=lambda pair: pair[0],
|
||||
),
|
||||
policy=SimpleNamespace(
|
||||
preprocessor=lambda observation: observation,
|
||||
postprocessor=lambda action: action,
|
||||
),
|
||||
data=SimpleNamespace(
|
||||
dataset_features={
|
||||
"observation.state": {
|
||||
"dtype": "float32",
|
||||
"shape": (1,),
|
||||
"names": ["joint.pos"],
|
||||
},
|
||||
"action": {"dtype": "float32", "shape": (1,), "names": ["joint.pos"]},
|
||||
}
|
||||
),
|
||||
)
|
||||
provider, executor = _build_rollout_runtime_io(ctx, rerun_log=False, get_task=lambda: "move")
|
||||
|
||||
observation = provider()
|
||||
executor(torch.tensor([[2.0]]))
|
||||
|
||||
assert observation["observation.state"].shape == (1, 1)
|
||||
robot.send_action.assert_called_once_with({"joint.pos": 2.0})
|
||||
@@ -1,100 +0,0 @@
|
||||
# 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 threading
|
||||
import time
|
||||
|
||||
from lerobot.runtime import LanguageConditionedRuntime, Tick
|
||||
|
||||
|
||||
class FakeAdapter:
|
||||
def __init__(self):
|
||||
self.updated = False
|
||||
self.interjections = []
|
||||
|
||||
def select_action(self, observation, state):
|
||||
assert observation == {"observation.state": 1}
|
||||
assert state.task == "clean"
|
||||
return ["a0", "a1"]
|
||||
|
||||
def update_language_state(self, observation, state):
|
||||
self.updated = True
|
||||
state.set_context("subtask", "pick cup", label="subtask")
|
||||
|
||||
def handle_interjection(self, user_text, observation, state):
|
||||
self.interjections.append(user_text)
|
||||
state.set_context("plan", "new plan", label="plan")
|
||||
|
||||
|
||||
def test_runtime_tick_updates_language_enqueues_and_dispatches_action():
|
||||
adapter = FakeAdapter()
|
||||
executed = []
|
||||
runtime = LanguageConditionedRuntime(
|
||||
policy_adapter=adapter,
|
||||
observation_provider=lambda: {"observation.state": 1},
|
||||
action_executor=executed.append,
|
||||
)
|
||||
runtime.set_task("clean")
|
||||
|
||||
logs = runtime.step_once()
|
||||
|
||||
assert adapter.updated
|
||||
assert runtime.state.language_context["subtask"] == "pick cup"
|
||||
assert executed == ["a0"]
|
||||
assert list(runtime.state.action_queue) == ["a1"]
|
||||
assert " subtask: pick cup" in logs
|
||||
|
||||
|
||||
def test_runtime_handles_user_interjection():
|
||||
adapter = FakeAdapter()
|
||||
runtime = LanguageConditionedRuntime(
|
||||
policy_adapter=adapter,
|
||||
observation_provider=lambda: {"observation.state": 1},
|
||||
)
|
||||
runtime.set_task("clean")
|
||||
runtime.state.extra["recent_interjection"] = "please say ok"
|
||||
runtime.state.emit("user_interjection")
|
||||
|
||||
runtime.step_once()
|
||||
|
||||
assert "please say ok" in adapter.interjections
|
||||
assert runtime.state.language_context["plan"] == "new plan"
|
||||
|
||||
|
||||
def test_prompt_change_discards_in_flight_action_chunk():
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
class BlockingAdapter(FakeAdapter):
|
||||
def select_action(self, observation, state):
|
||||
started.set()
|
||||
assert release.wait(timeout=2)
|
||||
return ["stale"]
|
||||
|
||||
runtime = LanguageConditionedRuntime(
|
||||
policy_adapter=BlockingAdapter(),
|
||||
observation_provider=lambda: {"observation.state": 1},
|
||||
)
|
||||
runtime.set_task("old task")
|
||||
runtime.state.tick = Tick(index=1, monotonic_seconds=time.monotonic())
|
||||
inference = threading.Thread(target=runtime.maybe_enqueue_action_chunk, kwargs={"force": True})
|
||||
inference.start()
|
||||
assert started.wait(timeout=2)
|
||||
|
||||
runtime.set_task("new task")
|
||||
release.set()
|
||||
inference.join(timeout=2)
|
||||
|
||||
assert not inference.is_alive()
|
||||
assert list(runtime.state.action_queue) == []
|
||||
@@ -1,82 +0,0 @@
|
||||
# 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 sys
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
|
||||
from lerobot.runtime.sim_robocasa import RoboCasaSimBackend
|
||||
from lerobot.utils.video_annotation import annotate_frame
|
||||
|
||||
|
||||
def test_overlay_draws_each_label_once(monkeypatch):
|
||||
put_text_calls = []
|
||||
rectangle_calls = []
|
||||
|
||||
def put_text(image, text, origin, font, scale, color, thickness, line_type):
|
||||
put_text_calls.append((text, color, thickness))
|
||||
return image
|
||||
|
||||
def rectangle(image, start, end, color, thickness):
|
||||
rectangle_calls.append((start, end, color, thickness))
|
||||
return image
|
||||
|
||||
def add_weighted(src1, alpha, src2, beta, gamma, *, dst):
|
||||
dst[:] = src1 * alpha + src2 * beta + gamma
|
||||
return dst
|
||||
|
||||
fake_cv2 = SimpleNamespace(
|
||||
FONT_HERSHEY_SIMPLEX=0,
|
||||
LINE_AA=16,
|
||||
getTextSize=lambda text, font, scale, thickness: ((len(text) * 7, 10), 0),
|
||||
putText=put_text,
|
||||
rectangle=rectangle,
|
||||
addWeighted=add_weighted,
|
||||
)
|
||||
monkeypatch.setitem(sys.modules, "cv2", fake_cv2)
|
||||
|
||||
frame = np.full((120, 480, 3), 200, dtype=np.uint8)
|
||||
annotated = annotate_frame(
|
||||
frame,
|
||||
(("Task", "close the fridge"), ("Subtask", "reach for the handle"), ("Memory", None)),
|
||||
)
|
||||
|
||||
assert [call[0] for call in put_text_calls] == [
|
||||
"Task: close the fridge",
|
||||
"Subtask: reach for the handle",
|
||||
]
|
||||
assert all(color == (255, 255, 255) and thickness == 1 for _, color, thickness in put_text_calls)
|
||||
assert len(rectangle_calls) == 1
|
||||
assert not np.shares_memory(annotated, frame)
|
||||
|
||||
|
||||
def test_capture_updates_live_frame_when_recording_is_disabled(monkeypatch):
|
||||
backend = object.__new__(RoboCasaSimBackend)
|
||||
frame = np.full((8, 8, 3), 42, dtype=np.uint8)
|
||||
written = []
|
||||
backend.record = False
|
||||
backend.runtime_state = None
|
||||
backend._multiview_frame = lambda: frame
|
||||
backend._current_task = lambda: "task"
|
||||
backend._subtask_getter = None
|
||||
backend._memory_getter = None
|
||||
backend._latest_frame = None
|
||||
backend._write_live_frame = written.append
|
||||
monkeypatch.setattr("lerobot.runtime.sim_robocasa.annotate_frame", lambda image, labels: image)
|
||||
|
||||
backend._capture_frame()
|
||||
|
||||
assert backend._latest_frame is frame
|
||||
assert written == [frame]
|
||||
@@ -18,6 +18,8 @@ import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
# ``lerobot.scripts.lerobot_annotate`` (and the ``_push_to_hub`` path it
|
||||
# exercises) imports ``lerobot.datasets``, which only ships under the
|
||||
@@ -26,7 +28,7 @@ pytest.importorskip("datasets", reason="datasets is required (install lerobot[da
|
||||
|
||||
|
||||
def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
from lerobot.scripts.lerobot_annotate import _push_to_hub
|
||||
from lerobot.scripts import lerobot_annotate
|
||||
|
||||
root = tmp_path / "dataset"
|
||||
(root / "meta").mkdir(parents=True)
|
||||
@@ -45,9 +47,6 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
return SimpleNamespace(oid="abc123")
|
||||
|
||||
def delete_tag(self, repo_id, **kwargs):
|
||||
import requests
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
calls["delete_tag"] = {"repo_id": repo_id, **kwargs}
|
||||
# Simulate the common case: no stale tag to delete.
|
||||
raise RevisionNotFoundError("no such tag", response=requests.Response())
|
||||
@@ -55,7 +54,7 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
def create_tag(self, **kwargs):
|
||||
calls["create_tag"] = kwargs
|
||||
|
||||
monkeypatch.setattr("huggingface_hub.HfApi", FakeHfApi)
|
||||
monkeypatch.setattr(lerobot_annotate, "HfApi", FakeHfApi)
|
||||
|
||||
def fake_card_push(self, **kwargs):
|
||||
calls["card_push"] = {"content": str(self), **kwargs}
|
||||
@@ -69,7 +68,7 @@ def test_push_to_hub_tags_uploaded_dataset_revision(tmp_path, monkeypatch):
|
||||
push_commit_message=None,
|
||||
)
|
||||
|
||||
_push_to_hub(root, cfg)
|
||||
lerobot_annotate._push_to_hub(root, cfg)
|
||||
|
||||
assert calls["create_repo"] == {
|
||||
"repo_id": "annotated/dataset",
|
||||
|
||||
@@ -44,10 +44,7 @@ def _install_robomme_stub():
|
||||
"joint_state_list": [np.zeros(7, dtype=np.float32)],
|
||||
"gripper_state_list": [np.zeros(2, dtype=np.float32)],
|
||||
}
|
||||
env.reset.return_value = (
|
||||
obs,
|
||||
{"status": "ongoing", "task_goal": ["pick the cube", "pick the blue cube"]},
|
||||
)
|
||||
env.reset.return_value = (obs, {"status": "ongoing", "task_goal": "pick the cube"})
|
||||
env.step.return_value = (obs, 0.0, False, False, {"status": "ongoing", "task_goal": ""})
|
||||
return env
|
||||
|
||||
@@ -119,21 +116,6 @@ def test_robomme_features_action_dim_ee_pose():
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_reset_exposes_episode_task_description():
|
||||
"""VLA evaluation receives the episode-specific language instruction."""
|
||||
_install_robomme_stub()
|
||||
try:
|
||||
from lerobot.envs.robomme import RoboMMEGymEnv
|
||||
|
||||
env = RoboMMEGymEnv(task="PickXtimes")
|
||||
env.reset()
|
||||
|
||||
assert env.task == "PickXtimes"
|
||||
assert env.task_description == "pick the cube"
|
||||
finally:
|
||||
_uninstall_robomme_stub()
|
||||
|
||||
|
||||
def test_convert_obs_list_format():
|
||||
"""_convert_obs takes the last element from list-format obs fields and
|
||||
emits a nested ``pixels`` dict (image, wrist_image) plus ``agent_pos``.
|
||||
|
||||
@@ -1,84 +0,0 @@
|
||||
#!/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
|
||||
from accelerate import Accelerator
|
||||
from torch import nn
|
||||
|
||||
from lerobot.scripts.lerobot_train import update_policy
|
||||
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
|
||||
|
||||
|
||||
class TinyPolicy(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.projection = nn.Linear(2, 1, bias=False)
|
||||
|
||||
def forward(self, batch):
|
||||
loss = self.projection(batch["x"]).square().mean()
|
||||
return loss, {}
|
||||
|
||||
|
||||
def test_gradient_accumulation_steps_optimizer_and_scheduler_once():
|
||||
accelerator = Accelerator(
|
||||
cpu=True,
|
||||
gradient_accumulation_steps=2,
|
||||
step_scheduler_with_optimizer=False,
|
||||
)
|
||||
policy = TinyPolicy()
|
||||
optimizer = torch.optim.SGD(policy.parameters(), lr=0.1)
|
||||
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=0.5)
|
||||
policy, optimizer, scheduler = accelerator.prepare(policy, optimizer, scheduler)
|
||||
metrics = {
|
||||
"loss": AverageMeter("loss"),
|
||||
"grad_norm": AverageMeter("grad_norm"),
|
||||
"lr": AverageMeter("lr"),
|
||||
"update_s": AverageMeter("update_s"),
|
||||
}
|
||||
tracker = MetricsTracker(1, 2, 1, metrics, accelerator=accelerator)
|
||||
batch = {"x": torch.ones(1, 2)}
|
||||
before = policy.projection.weight.detach().clone()
|
||||
|
||||
with accelerator.accumulate(policy):
|
||||
update_policy(
|
||||
tracker,
|
||||
policy,
|
||||
batch,
|
||||
optimizer,
|
||||
grad_clip_norm=0,
|
||||
accelerator=accelerator,
|
||||
lr_scheduler=scheduler,
|
||||
log_metrics=False,
|
||||
)
|
||||
after_first_microbatch = policy.projection.weight.detach().clone()
|
||||
|
||||
with accelerator.accumulate(policy):
|
||||
update_policy(
|
||||
tracker,
|
||||
policy,
|
||||
batch,
|
||||
optimizer,
|
||||
grad_clip_norm=0,
|
||||
accelerator=accelerator,
|
||||
lr_scheduler=scheduler,
|
||||
log_metrics=False,
|
||||
)
|
||||
after_optimizer_step = policy.projection.weight.detach().clone()
|
||||
|
||||
torch.testing.assert_close(after_first_microbatch, before)
|
||||
assert not torch.equal(after_optimizer_step, after_first_microbatch)
|
||||
assert optimizer.param_groups[0]["lr"] == pytest.approx(0.05)
|
||||
@@ -37,12 +37,6 @@ 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")
|
||||
@@ -174,18 +168,6 @@ 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")
|
||||
@@ -251,3 +233,37 @@ def test_metrics_tracker_reduce_across_ranks_invokes_reduce():
|
||||
# accumulate against the cluster view rather than the stale per-rank sum.
|
||||
meter = tracker.update_s
|
||||
assert meter.sum / meter.count == pytest.approx(meter.avg)
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_registers_and_averages():
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics={})
|
||||
tracker.update_metrics({"latent_loss": 0.2, "action_loss": 0.4})
|
||||
tracker.update_metrics({"latent_loss": 0.4, "action_loss": 0.6})
|
||||
|
||||
# New keys are auto-registered as mean-reduced meters and averaged over the window.
|
||||
assert tracker.metrics["latent_loss"].reduction == "mean"
|
||||
assert tracker.metrics["latent_loss"].avg == pytest.approx(0.3)
|
||||
assert tracker.metrics["action_loss"].avg == pytest.approx(0.5)
|
||||
assert tracker.to_dict()["latent_loss"] == pytest.approx(0.3)
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_skips_non_numeric():
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics={})
|
||||
tracker.update_metrics({"loss": 0.5, "head_mode": "sparse", "enabled": True})
|
||||
|
||||
# strings and bools ignored
|
||||
assert "loss" in tracker.metrics
|
||||
assert "head_mode" not in tracker.metrics
|
||||
assert "enabled" not in tracker.metrics
|
||||
|
||||
|
||||
def test_metrics_tracker_update_metrics_does_not_override_caller_meter():
|
||||
# A policy that echoes "loss" in its output dict must not overwrite the caller-owned,
|
||||
# already-aggregated loss meter.
|
||||
metrics = {"loss": AverageMeter("loss", ":.3f", reduction="mean")}
|
||||
tracker = MetricsTracker(batch_size=32, num_frames=1000, num_episodes=50, metrics=metrics)
|
||||
tracker.loss = 1.0 # caller-set optimized loss
|
||||
tracker.update_metrics({"loss": 99.0, "latent_loss": 0.2})
|
||||
|
||||
assert tracker.metrics["loss"].avg == pytest.approx(1.0) # snapshot ignored
|
||||
assert tracker.metrics["latent_loss"].avg == pytest.approx(0.2)
|
||||
|
||||
@@ -1129,7 +1129,7 @@ name = "decord"
|
||||
version = "0.6.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "numpy", marker = "(platform_machine != 'arm64' and platform_machine != 's390x' and sys_platform == 'darwin') or (platform_machine == 'AMD64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (platform_machine != 's390x' and sys_platform != 'darwin' and sys_platform != 'linux')" },
|
||||
{ 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" },
|
||||
@@ -2910,7 +2910,6 @@ 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'" },
|
||||
@@ -3159,7 +3158,6 @@ phone = [
|
||||
]
|
||||
pi = [
|
||||
{ name = "scipy" },
|
||||
{ name = "sentencepiece" },
|
||||
{ name = "transformers" },
|
||||
]
|
||||
placo-dep = [
|
||||
@@ -3220,9 +3218,6 @@ sarm = [
|
||||
scipy-dep = [
|
||||
{ name = "scipy" },
|
||||
]
|
||||
sentencepiece-dep = [
|
||||
{ name = "sentencepiece" },
|
||||
]
|
||||
smolvla = [
|
||||
{ name = "accelerate" },
|
||||
{ name = "num2words" },
|
||||
@@ -3423,7 +3418,6 @@ 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'" },
|
||||
@@ -3499,7 +3493,6 @@ 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" },
|
||||
@@ -3516,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", "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"]
|
||||
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"]
|
||||
|
||||
[[package]]
|
||||
name = "librt"
|
||||
@@ -6238,54 +6231,6 @@ 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.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/15/15/2e7a025fc62d764b151ae6d0f2a92f8081755ebe8d4a64099accc6f77ba6/sentencepiece-0.2.1.tar.gz", hash = "sha256:8138cec27c2f2282f4a34d9a016e3374cd40e5c6e9cb335063db66a0a3b71fad", size = 3228515, upload-time = "2025-08-12T07:00:51.718Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/4a/be/32ce495aa1d0e0c323dcb1ba87096037358edee539cac5baf8755a6bd396/sentencepiece-0.2.1-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:57cae326c8727de58c85977b175af132a7138d84c764635d7e71bbee7e774133", size = 1943152, upload-time = "2025-08-12T06:59:40.048Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/88/7e/ff23008899a58678e98c6ff592bf4d368eee5a71af96d0df6b38a039dd4f/sentencepiece-0.2.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:56dd39a3c4d6493db3cdca7e8cc68c6b633f0d4195495cbadfcf5af8a22d05a6", size = 1325651, upload-time = "2025-08-12T06:59:41.536Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/19/84/42eb3ce4796777a1b5d3699dfd4dca85113e68b637f194a6c8d786f16a04/sentencepiece-0.2.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:d9381351182ff9888cc80e41c632e7e274b106f450de33d67a9e8f6043da6f76", size = 1253645, upload-time = "2025-08-12T06:59:42.903Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/89/fa/d3d5ebcba3cb9e6d3775a096251860c41a6bc53a1b9461151df83fe93255/sentencepiece-0.2.1-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:99f955df238021bf11f0fc37cdb54fd5e5b5f7fd30ecc3d93fb48b6815437167", size = 1316273, upload-time = "2025-08-12T06:59:44.476Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/04/88/14f2f4a2b922d8b39be45bf63d79e6cd3a9b2f248b2fcb98a69b12af12f5/sentencepiece-0.2.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0cdfecef430d985f1c2bcbfff3defd1d95dae876fbd0173376012d2d7d24044b", size = 1387881, upload-time = "2025-08-12T06:59:46.09Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/fd/b8/903e5ccb77b4ef140605d5d71b4f9e0ad95d456d6184688073ed11712809/sentencepiece-0.2.1-cp312-cp312-win32.whl", hash = "sha256:a483fd29a34c3e34c39ac5556b0a90942bec253d260235729e50976f5dba1068", size = 999540, upload-time = "2025-08-12T06:59:48.023Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/2d/81/92df5673c067148c2545b1bfe49adfd775bcc3a169a047f5a0e6575ddaca/sentencepiece-0.2.1-cp312-cp312-win_amd64.whl", hash = "sha256:4cdc7c36234fda305e85c32949c5211faaf8dd886096c7cea289ddc12a2d02de", size = 1054671, upload-time = "2025-08-12T06:59:49.895Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/fe/02/c5e3bc518655d714622bec87d83db9cdba1cd0619a4a04e2109751c4f47f/sentencepiece-0.2.1-cp312-cp312-win_arm64.whl", hash = "sha256:daeb5e9e9fcad012324807856113708614d534f596d5008638eb9b40112cd9e4", size = 1033923, upload-time = "2025-08-12T06:59:51.952Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ba/4a/85fbe1706d4d04a7e826b53f327c4b80f849cf1c7b7c5e31a20a97d8f28b/sentencepiece-0.2.1-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:dcd8161eee7b41aae57ded06272905dbd680a0a04b91edd0f64790c796b2f706", size = 1943150, upload-time = "2025-08-12T06:59:53.588Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/c2/83/4cfb393e287509fc2155480b9d184706ef8d9fa8cbf5505d02a5792bf220/sentencepiece-0.2.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:c6c8f42949f419ff8c7e9960dbadcfbc982d7b5efc2f6748210d3dd53a7de062", size = 1325651, upload-time = "2025-08-12T06:59:55.073Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/8d/de/5a007fb53b1ab0aafc69d11a5a3dd72a289d5a3e78dcf2c3a3d9b14ffe93/sentencepiece-0.2.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:097f3394e99456e9e4efba1737c3749d7e23563dd1588ce71a3d007f25475fff", size = 1253641, upload-time = "2025-08-12T06:59:56.562Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/2c/d2/f552be5928105588f4f4d66ee37dd4c61460d8097e62d0e2e0eec41bc61d/sentencepiece-0.2.1-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d7b670879c370d350557edabadbad1f6561a9e6968126e6debca4029e5547820", size = 1316271, upload-time = "2025-08-12T06:59:58.109Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/96/df/0cfe748ace5485be740fed9476dee7877f109da32ed0d280312c94ec259f/sentencepiece-0.2.1-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c7f0fd2f2693309e6628aeeb2e2faf6edd221134dfccac3308ca0de01f8dab47", size = 1387882, upload-time = "2025-08-12T07:00:00.701Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ac/dd/f7774d42a881ced8e1739f393ab1e82ece39fc9abd4779e28050c2e975b5/sentencepiece-0.2.1-cp313-cp313-win32.whl", hash = "sha256:92b3816aa2339355fda2c8c4e021a5de92180b00aaccaf5e2808972e77a4b22f", size = 999541, upload-time = "2025-08-12T07:00:02.709Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/dd/e9/932b9eae6fd7019548321eee1ab8d5e3b3d1294df9d9a0c9ac517c7b636d/sentencepiece-0.2.1-cp313-cp313-win_amd64.whl", hash = "sha256:10ed3dab2044c47f7a2e7b4969b0c430420cdd45735d78c8f853191fa0e3148b", size = 1054669, upload-time = "2025-08-12T07:00:04.915Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/c9/3a/76488a00ea7d6931689cda28726a1447d66bf1a4837943489314593d5596/sentencepiece-0.2.1-cp313-cp313-win_arm64.whl", hash = "sha256:ac650534e2251083c5f75dde4ff28896ce7c8904133dc8fef42780f4d5588fcd", size = 1033922, upload-time = "2025-08-12T07:00:06.496Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/4a/b6/08fe2ce819e02ccb0296f4843e3f195764ce9829cbda61b7513f29b95718/sentencepiece-0.2.1-cp313-cp313t-macosx_10_13_universal2.whl", hash = "sha256:8dd4b477a7b069648d19363aad0cab9bad2f4e83b2d179be668efa672500dc94", size = 1946052, upload-time = "2025-08-12T07:00:08.136Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ab/d9/1ea0e740591ff4c6fc2b6eb1d7510d02f3fb885093f19b2f3abd1363b402/sentencepiece-0.2.1-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:0c0f672da370cc490e4c59d89e12289778310a0e71d176c541e4834759e1ae07", size = 1327408, upload-time = "2025-08-12T07:00:09.572Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/99/7e/1fb26e8a21613f6200e1ab88824d5d203714162cf2883248b517deb500b7/sentencepiece-0.2.1-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:ad8493bea8432dae8d6830365352350f3b4144415a1d09c4c8cb8d30cf3b6c3c", size = 1254857, upload-time = "2025-08-12T07:00:11.021Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/bc/85/c72fd1f3c7a6010544d6ae07f8ddb38b5e2a7e33bd4318f87266c0bbafbf/sentencepiece-0.2.1-cp313-cp313t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:b81a24733726e3678d2db63619acc5a8dccd074f7aa7a54ecd5ca33ca6d2d596", size = 1315722, upload-time = "2025-08-12T07:00:12.989Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/4a/e8/661e5bd82a8aa641fd6c1020bd0e890ef73230a2b7215ddf9c8cd8e941c2/sentencepiece-0.2.1-cp313-cp313t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0a81799d0a68d618e89063fb423c3001a034c893069135ffe51fee439ae474d6", size = 1387452, upload-time = "2025-08-12T07:00:15.088Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/99/5e/ae66c361023a470afcbc1fbb8da722c72ea678a2fcd9a18f1a12598c7501/sentencepiece-0.2.1-cp313-cp313t-win32.whl", hash = "sha256:89a3ea015517c42c0341d0d962f3e6aaf2cf10d71b1932d475c44ba48d00aa2b", size = 1002501, upload-time = "2025-08-12T07:00:16.966Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/c1/03/d332828c4ff764e16c1b56c2c8f9a33488bbe796b53fb6b9c4205ddbf167/sentencepiece-0.2.1-cp313-cp313t-win_amd64.whl", hash = "sha256:33f068c9382dc2e7c228eedfd8163b52baa86bb92f50d0488bf2b7da7032e484", size = 1057555, upload-time = "2025-08-12T07:00:18.573Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/88/14/5aee0bf0864df9bd82bd59e7711362908e4935e3f9cdc1f57246b5d5c9b9/sentencepiece-0.2.1-cp313-cp313t-win_arm64.whl", hash = "sha256:b3616ad246f360e52c85781e47682d31abfb6554c779e42b65333d4b5f44ecc0", size = 1036042, upload-time = "2025-08-12T07:00:20.209Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/24/9c/89eb8b2052f720a612478baf11c8227dcf1dc28cd4ea4c0c19506b5af2a2/sentencepiece-0.2.1-cp314-cp314-macosx_10_13_universal2.whl", hash = "sha256:5d0350b686c320068702116276cfb26c066dc7e65cfef173980b11bb4d606719", size = 1943147, upload-time = "2025-08-12T07:00:21.809Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/82/0b/a1432bc87f97c2ace36386ca23e8bd3b91fb40581b5e6148d24b24186419/sentencepiece-0.2.1-cp314-cp314-macosx_10_13_x86_64.whl", hash = "sha256:c7f54a31cde6fa5cb030370566f68152a742f433f8d2be458463d06c208aef33", size = 1325624, upload-time = "2025-08-12T07:00:23.289Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ea/99/bbe054ebb5a5039457c590e0a4156ed073fb0fe9ce4f7523404dd5b37463/sentencepiece-0.2.1-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:c83b85ab2d6576607f31df77ff86f28182be4a8de6d175d2c33ca609925f5da1", size = 1253670, upload-time = "2025-08-12T07:00:24.69Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/19/ad/d5c7075f701bd97971d7c2ac2904f227566f51ef0838dfbdfdccb58cd212/sentencepiece-0.2.1-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1855f57db07b51fb51ed6c9c452f570624d2b169b36f0f79ef71a6e6c618cd8b", size = 1316247, upload-time = "2025-08-12T07:00:26.435Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/fb/03/35fbe5f3d9a7435eebd0b473e09584bd3cc354ce118b960445b060d33781/sentencepiece-0.2.1-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:01e6912125cb45d3792f530a4d38f8e21bf884d6b4d4ade1b2de5cf7a8d2a52b", size = 1387894, upload-time = "2025-08-12T07:00:28.339Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/dc/aa/956ef729aafb6c8f9c443104c9636489093bb5c61d6b90fc27aa1a865574/sentencepiece-0.2.1-cp314-cp314-win32.whl", hash = "sha256:c415c9de1447e0a74ae3fdb2e52f967cb544113a3a5ce3a194df185cbc1f962f", size = 1096698, upload-time = "2025-08-12T07:00:29.764Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/b8/cb/fe400d8836952cc535c81a0ce47dc6875160e5fedb71d2d9ff0e9894c2a6/sentencepiece-0.2.1-cp314-cp314-win_amd64.whl", hash = "sha256:881b2e44b14fc19feade3cbed314be37de639fc415375cefaa5bc81a4be137fd", size = 1155115, upload-time = "2025-08-12T07:00:32.865Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/32/89/047921cf70f36c7b6b6390876b2399b3633ab73b8d0cb857e5a964238941/sentencepiece-0.2.1-cp314-cp314-win_arm64.whl", hash = "sha256:2005242a16d2dc3ac5fe18aa7667549134d37854823df4c4db244752453b78a8", size = 1133890, upload-time = "2025-08-12T07:00:34.763Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/a1/11/5b414b9fae6255b5fb1e22e2ed3dc3a72d3a694e5703910e640ac78346bb/sentencepiece-0.2.1-cp314-cp314t-macosx_10_13_universal2.whl", hash = "sha256:a19adcec27c524cb7069a1c741060add95f942d1cbf7ad0d104dffa0a7d28a2b", size = 1946081, upload-time = "2025-08-12T07:00:36.97Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/77/eb/7a5682bb25824db8545f8e5662e7f3e32d72a508fdce086029d89695106b/sentencepiece-0.2.1-cp314-cp314t-macosx_10_13_x86_64.whl", hash = "sha256:e37e4b4c4a11662b5db521def4e44d4d30ae69a1743241412a93ae40fdcab4bb", size = 1327406, upload-time = "2025-08-12T07:00:38.669Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/03/b0/811dae8fb9f2784e138785d481469788f2e0d0c109c5737372454415f55f/sentencepiece-0.2.1-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:477c81505db072b3ab627e7eab972ea1025331bd3a92bacbf798df2b75ea86ec", size = 1254846, upload-time = "2025-08-12T07:00:40.611Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/ef/23/195b2e7ec85ebb6a547969f60b723c7aca5a75800ece6cc3f41da872d14e/sentencepiece-0.2.1-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:010f025a544ef770bb395091d57cb94deb9652d8972e0d09f71d85d5a0816c8c", size = 1315721, upload-time = "2025-08-12T07:00:42.914Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/7e/aa/553dbe4178b5f23eb28e59393dddd64186178b56b81d9b8d5c3ff1c28395/sentencepiece-0.2.1-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:733e59ff1794d26db706cd41fc2d7ca5f6c64a820709cb801dc0ea31780d64ab", size = 1387458, upload-time = "2025-08-12T07:00:44.56Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/66/7c/08ff0012507297a4dd74a5420fdc0eb9e3e80f4e88cab1538d7f28db303d/sentencepiece-0.2.1-cp314-cp314t-win32.whl", hash = "sha256:d3233770f78e637dc8b1fda2cd7c3b99ec77e7505041934188a4e7fe751de3b0", size = 1099765, upload-time = "2025-08-12T07:00:46.058Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/91/d5/2a69e1ce15881beb9ddfc7e3f998322f5cedcd5e4d244cb74dade9441663/sentencepiece-0.2.1-cp314-cp314t-win_amd64.whl", hash = "sha256:5e4366c97b68218fd30ea72d70c525e6e78a6c0a88650f57ac4c43c63b234a9d", size = 1157807, upload-time = "2025-08-12T07:00:47.673Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/f3/16/54f611fcfc2d1c46cbe3ec4169780b2cfa7cf63708ef2b71611136db7513/sentencepiece-0.2.1-cp314-cp314t-win_arm64.whl", hash = "sha256:105e36e75cbac1292642045458e8da677b2342dcd33df503e640f0b457cb6751", size = 1136264, upload-time = "2025-08-12T07:00:49.485Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sentry-sdk"
|
||||
version = "2.64.0"
|
||||
|
||||
Reference in New Issue
Block a user