mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 01:41:54 +00:00
Add joint-sequence subtask training and collision-free FAST vocab mapping
- recipes/subtask_joint.yaml: paper-style single sequence (pi0.5 §IV-B) — the supervised subtask span gets text CE and conditions the FAST and flow losses in the same forward. - joint_subtask_conditioning config flag rebuilds the same layout at inference: state on the task turn, generated subtask as a causal assistant turn (encode_prompt_with_targets + lang_causal_marks through sample_actions), in both the policy select_action path and the runtime adapter. - fast_skip_tokens default 128 -> 1152 so FAST codes land below the <loc> range and never collide with VQA loc targets; _FAST_ACTION_VOCAB_SIZE tightened to the universal tokenizer's 1024 codes. - Strip the trailing space from the 'Assistant:' generation prefill — SentencePiece folds the space into the first target token, so the space-suffixed prefill ended in a lone '▁' never seen in training. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -146,6 +146,12 @@ The renderer does not apply a tokenizer chat template. Policy processors decide
|
|||||||
Blend recipes select one weighted sub-recipe deterministically from the sample index.
|
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.
|
`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.
|
||||||
|
|
||||||
|
A message recipe with a supervised assistant turn on the `low_level` stream trains
|
||||||
|
the π0.5 paper's joint sequence instead of a blend: the target span gets text CE
|
||||||
|
while also conditioning the action losses in the same forward.
|
||||||
|
`recipes/subtask_joint.yaml` is the provided example; pair it with
|
||||||
|
`--policy.joint_subtask_conditioning=true` at inference.
|
||||||
|
|
||||||
## Graceful absence
|
## Graceful absence
|
||||||
|
|
||||||
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
|
||||||
|
|||||||
@@ -55,9 +55,20 @@ The provided recipes are:
|
|||||||
| Recipe | Required annotations | Trains |
|
| Recipe | Required annotations | Trains |
|
||||||
| ------------------------------------- | ----------------------------------------------------------------------- | -------------------------------------------------- |
|
| ------------------------------------- | ----------------------------------------------------------------------- | -------------------------------------------------- |
|
||||||
| `recipes/subtask.yaml` | `subtask` | Subtask prediction and subtask-conditioned actions |
|
| `recipes/subtask.yaml` | `subtask` | Subtask prediction and subtask-conditioned actions |
|
||||||
|
| `recipes/subtask_joint.yaml` | `subtask` | Paper-style joint sequence: subtask text and actions in one sample |
|
||||||
| `recipes/subtask_mem.yaml` | `subtask`, `memory` | Subtasks, actions, and memory updates |
|
| `recipes/subtask_mem.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 |
|
| `recipes/subtask_mem_vqa_speech.yaml` | `subtask`, `memory`, `vqa`; interjection/speech rows for those branches | Subtasks, actions, memory, VQA, and spoken replies |
|
||||||
|
|
||||||
|
The blend recipes factorize training into separate high-level (task → subtask)
|
||||||
|
and low-level (subtask → actions) samples, matching how inference decomposes
|
||||||
|
π(a|o, subtask)·π(subtask|o, task). `recipes/subtask_joint.yaml` instead uses
|
||||||
|
the π0.5 paper's single-sequence layout — the supervised subtask span is
|
||||||
|
attended causally and conditions the FAST and flow losses in the same forward.
|
||||||
|
Checkpoints trained with the joint recipe must set
|
||||||
|
`--policy.joint_subtask_conditioning=true` at inference so the flow prefix
|
||||||
|
rebuilds the same layout (task turn with state, then the generated subtask as a
|
||||||
|
causal assistant turn); leave it `false` for the blend recipes.
|
||||||
|
|
||||||
Use `lerobot-annotate` to generate these columns. The repository includes a
|
Use `lerobot-annotate` to generate these columns. The repository includes a
|
||||||
Hugging Face Jobs launcher that you can edit for your source and destination
|
Hugging Face Jobs launcher that you can edit for your source and destination
|
||||||
datasets. For a local annotation run, first install
|
datasets. For a local annotation run, first install
|
||||||
@@ -117,6 +128,14 @@ the expected prompt, text target, and action endpoints before scaling up.
|
|||||||
| `policy.knowledge_insulation` | `true` | Blocks action-loss gradients through the VLM K/V path |
|
| `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.flow_num_repeats` | `5` | Reuses one VLM prefix for independent denoising targets |
|
||||||
| `policy.lm_head_lr_scale` | `1.0` | Scales language-head learning rate; `1.0` uses the base rate |
|
| `policy.lm_head_lr_scale` | `1.0` | Scales language-head learning rate; `1.0` uses the base rate |
|
||||||
|
| `policy.fast_skip_tokens` | `1152` | FAST id offset; skips `<seg>`+`<loc>` so VQA and FAST never collide |
|
||||||
|
| `policy.joint_subtask_conditioning` | `false` | Rebuilds the joint-sequence prefix at inference (see recipes) |
|
||||||
|
|
||||||
|
`fast_skip_tokens=1152` places FAST codes below PaliGemma's `<loc>` range.
|
||||||
|
openpi's pi0-FAST convention is `128` (FAST occupies the `<loc>` ids); use that
|
||||||
|
value only to stay weight-compatible with checkpoints trained that way, and
|
||||||
|
avoid combining it with the VQA recipe, whose `<loc>` targets would share
|
||||||
|
embedding rows with FAST codes.
|
||||||
|
|
||||||
The loss weights are starting points, not dataset-independent constants. Track
|
The loss weights are starting points, not dataset-independent constants. Track
|
||||||
flow loss and text/FAST losses separately, and inspect generated subtasks rather
|
flow loss and text/FAST losses separately, and inspect generated subtasks rather
|
||||||
|
|||||||
@@ -0,0 +1,13 @@
|
|||||||
|
# Paper-style joint sequence (pi0.5 §IV-B): one sample supervises the subtask
|
||||||
|
# text with CE and, because the assistant turn is part of the prefix, conditions
|
||||||
|
# the FAST and flow action losses on the same annotated subtask in one forward.
|
||||||
|
# The supervised span is attended causally; the action losses see task + subtask.
|
||||||
|
#
|
||||||
|
# Pair with `--policy.joint_subtask_conditioning=true` at inference so the flow
|
||||||
|
# prefix reproduces this layout (task turn with state + causal generated subtask).
|
||||||
|
# Samples without a `subtask` annotation fall back to a plain task-prompt
|
||||||
|
# low-level sample via `if_present`.
|
||||||
|
|
||||||
|
messages:
|
||||||
|
- {role: user, content: "${task}", stream: low_level}
|
||||||
|
- {role: assistant, content: "${subtask}", stream: low_level, target: true, if_present: subtask}
|
||||||
@@ -63,9 +63,14 @@ class PI052Config(PI05Config):
|
|||||||
max_action_tokens: int = 256
|
max_action_tokens: int = 256
|
||||||
"""Maximum number of FAST tokens per action chunk."""
|
"""Maximum number of FAST tokens per action chunk."""
|
||||||
|
|
||||||
fast_skip_tokens: int = 128
|
fast_skip_tokens: int = 1152
|
||||||
"""Number of low-vocab tokens the FAST tokenizer skips to avoid
|
"""Number of top-of-vocab tokens the FAST id mapping skips.
|
||||||
collisions with PaliGemma's text vocabulary."""
|
|
||||||
|
1152 skips PaliGemma's 128 ``<seg>`` and 1024 ``<loc>`` special tokens so
|
||||||
|
FAST codes land in plain-text ids below 256000 and never collide with the
|
||||||
|
``<loc>`` targets used for VQA. openpi's pi0-FAST convention is 128 (FAST
|
||||||
|
occupies the ``<loc>`` range); use 128 only to stay weight-compatible with
|
||||||
|
checkpoints trained that way."""
|
||||||
|
|
||||||
fast_action_loss_weight: float = 1.0
|
fast_action_loss_weight: float = 1.0
|
||||||
"""Weight on FAST action-token CE relative to continuous-flow supervision."""
|
"""Weight on FAST action-token CE relative to continuous-flow supervision."""
|
||||||
@@ -76,6 +81,16 @@ class PI052Config(PI05Config):
|
|||||||
Non-positive values regenerate each action chunk while still refreshing the action prompt every chunk.
|
Non-positive values regenerate each action chunk while still refreshing the action prompt every chunk.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
joint_subtask_conditioning: bool = False
|
||||||
|
"""Condition low-level action inference on the task plus the generated subtask.
|
||||||
|
|
||||||
|
Matches paper-style joint-sequence recipes (``recipes/subtask_joint.yaml``)
|
||||||
|
where one sample supervises the subtask text and conditions the action
|
||||||
|
losses on it: the inference prefix becomes
|
||||||
|
``User: {task}, State: ...;\\nAssistant: {subtask}<eos>`` with the subtask
|
||||||
|
span attended causally, exactly as trained. Leave ``False`` for the blend
|
||||||
|
recipes, whose low-level samples use ``User: {subtask}, State: ...;``."""
|
||||||
|
|
||||||
auto_fit_fast_tokenizer: bool = False
|
auto_fit_fast_tokenizer: bool = False
|
||||||
"""Fit and cache a dataset-specific FAST tokenizer before training.
|
"""Fit and cache a dataset-specific FAST tokenizer before training.
|
||||||
|
|
||||||
|
|||||||
@@ -41,22 +41,53 @@ class PI052PolicyAdapter(BaseLanguageAdapter):
|
|||||||
|
|
||||||
subtask = state.language_context.get("subtask") or state.task or ""
|
subtask = state.language_context.get("subtask") or state.task or ""
|
||||||
# Match the training prompt by conditioning on both subtask and discretized state.
|
# Match the training prompt by conditioning on both subtask and discretized state.
|
||||||
content = subtask
|
state_str = None
|
||||||
obs_state = observation.get(OBS_STATE)
|
obs_state = observation.get(OBS_STATE)
|
||||||
if isinstance(obs_state, torch.Tensor) and obs_state.numel() > 0:
|
if isinstance(obs_state, torch.Tensor) and obs_state.numel() > 0:
|
||||||
from lerobot.policies.pi052.text_processor_pi052 import discretize_state_str # noqa: PLC0415
|
from lerobot.policies.pi052.text_processor_pi052 import discretize_state_str # noqa: PLC0415
|
||||||
|
|
||||||
state_row = obs_state[0] if obs_state.ndim > 1 else obs_state
|
state_row = obs_state[0] if obs_state.ndim > 1 else obs_state
|
||||||
content = f"{subtask}, State: {discretize_state_str(state_row)};"
|
state_str = discretize_state_str(state_row)
|
||||||
|
|
||||||
text_batch = _build_text_batch(
|
|
||||||
self.policy,
|
|
||||||
[{"role": "user", "content": content}],
|
|
||||||
add_generation_prompt=False,
|
|
||||||
)
|
|
||||||
batch = dict(observation)
|
batch = dict(observation)
|
||||||
batch[OBS_LANGUAGE_TOKENS] = text_batch["lang_tokens"]
|
if getattr(self.policy.config, "joint_subtask_conditioning", False):
|
||||||
batch[OBS_LANGUAGE_ATTENTION_MASK] = text_batch["lang_masks"]
|
# Joint sequences keep the task turn (with state) and render the
|
||||||
|
# subtask as a causal assistant turn, exactly as trained.
|
||||||
|
from transformers import AutoTokenizer # noqa: PLC0415
|
||||||
|
|
||||||
|
from lerobot.policies.pi052.text_processor_pi052 import ( # noqa: PLC0415
|
||||||
|
encode_prompt_with_targets,
|
||||||
|
register_paligemma_loc_tokens,
|
||||||
|
)
|
||||||
|
from lerobot.utils.constants import OBS_LANGUAGE_CAUSAL_MARKS # noqa: PLC0415
|
||||||
|
|
||||||
|
task = state.task or ""
|
||||||
|
task_content = task if state_str is None else f"{task}, State: {state_str};"
|
||||||
|
tok_name = getattr(self.policy.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
|
||||||
|
tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
|
||||||
|
ids, attn, marks = encode_prompt_with_targets(
|
||||||
|
tokenizer,
|
||||||
|
[
|
||||||
|
{"role": "user", "content": task_content},
|
||||||
|
{"role": "assistant", "content": subtask},
|
||||||
|
],
|
||||||
|
target_indices=[1],
|
||||||
|
)
|
||||||
|
device = getattr(self.policy.config, "device", None)
|
||||||
|
if device is not None:
|
||||||
|
ids, attn, marks = ids.to(device), attn.to(device), marks.to(device)
|
||||||
|
batch[OBS_LANGUAGE_TOKENS] = ids
|
||||||
|
batch[OBS_LANGUAGE_ATTENTION_MASK] = attn
|
||||||
|
batch[OBS_LANGUAGE_CAUSAL_MARKS] = marks
|
||||||
|
else:
|
||||||
|
content = subtask if state_str is None else f"{subtask}, State: {state_str};"
|
||||||
|
text_batch = _build_text_batch(
|
||||||
|
self.policy,
|
||||||
|
[{"role": "user", "content": content}],
|
||||||
|
add_generation_prompt=False,
|
||||||
|
)
|
||||||
|
batch[OBS_LANGUAGE_TOKENS] = text_batch["lang_tokens"]
|
||||||
|
batch[OBS_LANGUAGE_ATTENTION_MASK] = text_batch["lang_masks"]
|
||||||
return self.policy.predict_action_chunk(batch)
|
return self.policy.predict_action_chunk(batch)
|
||||||
|
|
||||||
def generate_text(
|
def generate_text(
|
||||||
@@ -67,6 +98,21 @@ class PI052PolicyAdapter(BaseLanguageAdapter):
|
|||||||
user_text: str | None = None,
|
user_text: str | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
messages = self.build_messages(kind, state, user_text=user_text)
|
messages = self.build_messages(kind, state, user_text=user_text)
|
||||||
|
if kind == "subtask" and getattr(self.policy.config, "joint_subtask_conditioning", False):
|
||||||
|
# Joint samples carry state on the task turn, so the subtask must be
|
||||||
|
# generated from the same state-bearing prompt.
|
||||||
|
import torch # noqa: PLC0415
|
||||||
|
|
||||||
|
from lerobot.policies.pi052.text_processor_pi052 import discretize_state_str # noqa: PLC0415
|
||||||
|
from lerobot.utils.constants import OBS_STATE # noqa: PLC0415
|
||||||
|
|
||||||
|
obs_state = (observation or {}).get(OBS_STATE)
|
||||||
|
if isinstance(obs_state, torch.Tensor) and obs_state.numel() > 0:
|
||||||
|
state_row = obs_state[0] if obs_state.ndim > 1 else obs_state
|
||||||
|
for m in reversed(messages):
|
||||||
|
if m.get("role") == "user":
|
||||||
|
m["content"] = f"{m.get('content', '')}, State: {discretize_state_str(state_row)};"
|
||||||
|
break
|
||||||
return _generate_with_policy(
|
return _generate_with_policy(
|
||||||
self.policy,
|
self.policy,
|
||||||
messages,
|
messages,
|
||||||
@@ -141,7 +187,10 @@ def _build_text_batch(
|
|||||||
messages = [_strip_blocks(_flatten_say_tool_calls(m)) for m in prompt_messages]
|
messages = [_strip_blocks(_flatten_say_tool_calls(m)) for m in prompt_messages]
|
||||||
prompt, _spans = _format_messages(messages)
|
prompt, _spans = _format_messages(messages)
|
||||||
if add_generation_prompt:
|
if add_generation_prompt:
|
||||||
prompt = prompt + "Assistant: "
|
# No trailing space: SentencePiece folds it into the first target token
|
||||||
|
# ("▁move"), so a space-suffixed prefill ends in a lone "▁" the model
|
||||||
|
# never saw at this position during training.
|
||||||
|
prompt = prompt + "Assistant:"
|
||||||
|
|
||||||
encoded = tokenizer(prompt, return_tensors="pt")
|
encoded = tokenizer(prompt, return_tensors="pt")
|
||||||
ids = encoded["input_ids"]
|
ids = encoded["input_ids"]
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from torch.nn import functional
|
|||||||
from lerobot.utils.constants import (
|
from lerobot.utils.constants import (
|
||||||
ACTION,
|
ACTION,
|
||||||
OBS_LANGUAGE_ATTENTION_MASK,
|
OBS_LANGUAGE_ATTENTION_MASK,
|
||||||
|
OBS_LANGUAGE_CAUSAL_MARKS,
|
||||||
OBS_LANGUAGE_TOKENS,
|
OBS_LANGUAGE_TOKENS,
|
||||||
OBS_STATE,
|
OBS_STATE,
|
||||||
)
|
)
|
||||||
@@ -143,9 +144,42 @@ class PI05Pytorch(PI05PytorchBase): # see openpi `PI0Pytorch`
|
|||||||
suffix_out = suffix_out.to(dtype=torch.float32)
|
suffix_out = suffix_out.to(dtype=torch.float32)
|
||||||
return self.action_out_proj(suffix_out)
|
return self.action_out_proj(suffix_out)
|
||||||
|
|
||||||
|
def sample_actions(
|
||||||
|
self,
|
||||||
|
images,
|
||||||
|
img_masks,
|
||||||
|
tokens,
|
||||||
|
masks,
|
||||||
|
noise=None,
|
||||||
|
num_steps=None,
|
||||||
|
lang_causal_marks=None,
|
||||||
|
**kwargs,
|
||||||
|
) -> Tensor:
|
||||||
|
"""Sample actions, optionally marking trailing language positions causal.
|
||||||
|
|
||||||
# FAST tokens occupy the high vocabulary range and must be masked during text generation.
|
``lang_causal_marks`` (B, L_lang bool) flags generated-subtask tokens so
|
||||||
_FAST_ACTION_VOCAB_SIZE = 2048
|
joint-sequence checkpoints see the same causal prefix layout at
|
||||||
|
inference as during training (``_mark_target_span_causal``).
|
||||||
|
"""
|
||||||
|
self._lang_causal_marks = lang_causal_marks
|
||||||
|
try:
|
||||||
|
return super().sample_actions(
|
||||||
|
images, img_masks, tokens, masks, noise=noise, num_steps=num_steps, **kwargs
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
self._lang_causal_marks = None
|
||||||
|
|
||||||
|
def embed_prefix(self, images, img_masks, tokens, masks):
|
||||||
|
prefix_embs, prefix_pad, prefix_att = super().embed_prefix(images, img_masks, tokens, masks)
|
||||||
|
marks = getattr(self, "_lang_causal_marks", None)
|
||||||
|
if marks is not None:
|
||||||
|
prefix_att = _apply_causal_language_marks(prefix_att, marks.to(prefix_att.device))
|
||||||
|
return prefix_embs, prefix_pad, prefix_att
|
||||||
|
|
||||||
|
|
||||||
|
# The universal `physical-intelligence/fast` tokenizer (and dataset refits of it)
|
||||||
|
# uses 1024 BPE codes; text generation must mask any that map below the <loc> range.
|
||||||
|
_FAST_ACTION_VOCAB_SIZE = 1024
|
||||||
|
|
||||||
|
|
||||||
_HF_KERNELS_ENABLED = False
|
_HF_KERNELS_ENABLED = False
|
||||||
@@ -331,6 +365,17 @@ def _mark_target_span_causal(
|
|||||||
return att
|
return att
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_causal_language_marks(prefix_att_masks: Tensor, marks: Tensor) -> Tensor:
|
||||||
|
"""OR per-token causal marks into the trailing language segment of a prefix."""
|
||||||
|
att = prefix_att_masks.clone()
|
||||||
|
n = min(marks.shape[1], att.shape[1])
|
||||||
|
if n <= 0:
|
||||||
|
return att
|
||||||
|
seg = att[:, -n:].bool()
|
||||||
|
att[:, -n:] = (seg | marks[:, -n:].bool()).to(att.dtype)
|
||||||
|
return att
|
||||||
|
|
||||||
|
|
||||||
def _fast_lin_ce(
|
def _fast_lin_ce(
|
||||||
hidden: Tensor,
|
hidden: Tensor,
|
||||||
lm_head_weight: Tensor,
|
lm_head_weight: Tensor,
|
||||||
@@ -1622,8 +1667,12 @@ class PI052Policy(PI05Policy):
|
|||||||
return decoded
|
return decoded
|
||||||
|
|
||||||
def _prepare_action_batch(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
|
def _prepare_action_batch(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
|
||||||
from .inference.pi052_adapter import _build_text_batch # noqa: PLC0415
|
from .inference.pi052_adapter import _build_text_batch, _get_loc_tokenizer # noqa: PLC0415
|
||||||
from .text_processor_pi052 import discretize_state_str # noqa: PLC0415
|
from .text_processor_pi052 import ( # noqa: PLC0415
|
||||||
|
discretize_state_str,
|
||||||
|
encode_prompt_with_targets,
|
||||||
|
register_paligemma_loc_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
n = self._batch_size_from_observation(batch)
|
n = self._batch_size_from_observation(batch)
|
||||||
self._ensure_subtask_state(n)
|
self._ensure_subtask_state(n)
|
||||||
@@ -1631,6 +1680,14 @@ class PI052Policy(PI05Policy):
|
|||||||
# Mirror training by appending the already normalized state to low-level prompts.
|
# Mirror training by appending the already normalized state to low-level prompts.
|
||||||
state_all = batch.get(OBS_STATE)
|
state_all = batch.get(OBS_STATE)
|
||||||
|
|
||||||
|
joint = bool(getattr(self.config, "joint_subtask_conditioning", False))
|
||||||
|
joint_tokenizer = None
|
||||||
|
if joint:
|
||||||
|
from transformers import AutoTokenizer # noqa: PLC0415
|
||||||
|
|
||||||
|
tok_name = getattr(self.config, "tokenizer_name", None) or "google/paligemma-3b-pt-224"
|
||||||
|
joint_tokenizer = _get_loc_tokenizer(tok_name, AutoTokenizer, register_paligemma_loc_tokens)
|
||||||
|
|
||||||
# Hold subtasks for the configured interval to match training and avoid rapid replanning.
|
# Hold subtasks for the configured interval to match training and avoid rapid replanning.
|
||||||
replan = int(getattr(self.config, "subtask_replan_steps", 0) or 0)
|
replan = int(getattr(self.config, "subtask_replan_steps", 0) or 0)
|
||||||
hold_chunks = max(1, round(replan / self.config.n_action_steps)) if replan > 0 else 1
|
hold_chunks = max(1, round(replan / self.config.n_action_steps)) if replan > 0 else 1
|
||||||
@@ -1638,7 +1695,7 @@ class PI052Policy(PI05Policy):
|
|||||||
self._subtask_chunk_counter += 1
|
self._subtask_chunk_counter += 1
|
||||||
|
|
||||||
# Generate and batch one independently conditioned subtask per environment.
|
# Generate and batch one independently conditioned subtask per environment.
|
||||||
rows: list[tuple[Tensor, Tensor | None]] = []
|
rows: list[tuple[Tensor, Tensor | None, Tensor | None]] = []
|
||||||
tokenizer = None
|
tokenizer = None
|
||||||
for i in range(n):
|
for i in range(n):
|
||||||
if regenerate or not self.last_subtasks[i]:
|
if regenerate or not self.last_subtasks[i]:
|
||||||
@@ -1648,32 +1705,62 @@ class PI052Policy(PI05Policy):
|
|||||||
# Hold the previously generated subtask; only the state in the
|
# Hold the previously generated subtask; only the state in the
|
||||||
# prompt below is refreshed to the current observation.
|
# prompt below is refreshed to the current observation.
|
||||||
subtask = self.last_subtasks[i]
|
subtask = self.last_subtasks[i]
|
||||||
content = subtask
|
state_str = discretize_state_str(state_all[i]) if torch.is_tensor(state_all) else None
|
||||||
if torch.is_tensor(state_all):
|
if joint:
|
||||||
content = f"{subtask}, State: {discretize_state_str(state_all[i])};"
|
# Joint sequences keep the task turn (with state) and render the
|
||||||
text_batch = _build_text_batch(
|
# subtask as a causal assistant turn, exactly as trained.
|
||||||
self,
|
task_content = tasks[i]
|
||||||
[{"role": "user", "content": content}],
|
if state_str is not None:
|
||||||
add_generation_prompt=False,
|
task_content = f"{task_content}, State: {state_str};"
|
||||||
)
|
ids, attn, marks = encode_prompt_with_targets(
|
||||||
rows.append((text_batch["lang_tokens"], text_batch["lang_masks"]))
|
joint_tokenizer,
|
||||||
tokenizer = text_batch["tokenizer"]
|
[
|
||||||
|
{"role": "user", "content": task_content},
|
||||||
|
{"role": "assistant", "content": subtask},
|
||||||
|
],
|
||||||
|
target_indices=[1],
|
||||||
|
)
|
||||||
|
device = getattr(self.config, "device", None)
|
||||||
|
if device is not None:
|
||||||
|
ids, attn, marks = ids.to(device), attn.to(device), marks.to(device)
|
||||||
|
rows.append((ids, attn, marks))
|
||||||
|
tokenizer = joint_tokenizer
|
||||||
|
else:
|
||||||
|
content = subtask if state_str is None else f"{subtask}, State: {state_str};"
|
||||||
|
text_batch = _build_text_batch(
|
||||||
|
self,
|
||||||
|
[{"role": "user", "content": content}],
|
||||||
|
add_generation_prompt=False,
|
||||||
|
)
|
||||||
|
rows.append((text_batch["lang_tokens"], text_batch["lang_masks"], None))
|
||||||
|
tokenizer = text_batch["tokenizer"]
|
||||||
|
|
||||||
tokens, masks = self._stack_token_rows(rows, tokenizer)
|
tokens, masks, marks = self._stack_token_rows(rows, tokenizer)
|
||||||
|
|
||||||
out = dict(batch)
|
out = dict(batch)
|
||||||
out[OBS_LANGUAGE_TOKENS] = tokens
|
out[OBS_LANGUAGE_TOKENS] = tokens
|
||||||
out[OBS_LANGUAGE_ATTENTION_MASK] = masks
|
out[OBS_LANGUAGE_ATTENTION_MASK] = masks
|
||||||
|
if marks is not None:
|
||||||
|
out[OBS_LANGUAGE_CAUSAL_MARKS] = marks
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def _generate_low_level_subtask(self, obs_i: dict[str, Tensor], task: str, i: int) -> str:
|
def _generate_low_level_subtask(self, obs_i: dict[str, Tensor], task: str, i: int) -> str:
|
||||||
from .inference.pi052_adapter import _generate_with_policy # noqa: PLC0415
|
from .inference.pi052_adapter import _generate_with_policy # noqa: PLC0415
|
||||||
|
from .text_processor_pi052 import discretize_state_str # noqa: PLC0415
|
||||||
|
|
||||||
msg = ""
|
msg = ""
|
||||||
if task:
|
if task:
|
||||||
|
content = task
|
||||||
|
if getattr(self.config, "joint_subtask_conditioning", False):
|
||||||
|
# Joint samples carry state on the task turn, so the subtask
|
||||||
|
# must be generated from the same state-bearing prompt.
|
||||||
|
state = obs_i.get(OBS_STATE)
|
||||||
|
if torch.is_tensor(state) and state.numel() > 0:
|
||||||
|
state_row = state[0] if state.ndim > 1 else state
|
||||||
|
content = f"{task}, State: {discretize_state_str(state_row)};"
|
||||||
msg = _generate_with_policy(
|
msg = _generate_with_policy(
|
||||||
self,
|
self,
|
||||||
[{"role": "user", "content": task}],
|
[{"role": "user", "content": content}],
|
||||||
observation=obs_i,
|
observation=obs_i,
|
||||||
label=f"eval subtask gen[{i}]",
|
label=f"eval subtask gen[{i}]",
|
||||||
suppress_loc_tokens=True,
|
suppress_loc_tokens=True,
|
||||||
@@ -1740,21 +1827,27 @@ class PI052Policy(PI05Policy):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _stack_token_rows(rows: list[tuple[Tensor, Tensor | None]], tokenizer: Any) -> tuple[Tensor, Tensor]:
|
def _stack_token_rows(
|
||||||
"""Right-pad per-env ``(1, L_i)`` token/mask rows and stack to ``(n, L)``.
|
rows: list[tuple[Tensor, Tensor | None, Tensor | None]], tokenizer: Any
|
||||||
|
) -> tuple[Tensor, Tensor, Tensor | None]:
|
||||||
|
"""Right-pad per-env ``(1, L_i)`` token/mask/marks rows and stack to ``(n, L)``.
|
||||||
|
|
||||||
Right-padding with a False attention mask matches the training-time
|
Right-padding with a False attention mask matches the training-time
|
||||||
tokenizer (``padding_side="right"``), so the action expert treats pad
|
tokenizer (``padding_side="right"``), so the action expert treats pad
|
||||||
positions as masked.
|
positions as masked. Causal marks (third element, optional) pad False.
|
||||||
"""
|
"""
|
||||||
max_len = max(t.shape[1] for t, _ in rows)
|
max_len = max(t.shape[1] for t, _, _ in rows)
|
||||||
pad_id = getattr(tokenizer, "pad_token_id", None) or 0
|
pad_id = getattr(tokenizer, "pad_token_id", None) or 0
|
||||||
|
has_marks = any(m is not None for _, _, m in rows)
|
||||||
tok_rows: list[Tensor] = []
|
tok_rows: list[Tensor] = []
|
||||||
mask_rows: list[Tensor] = []
|
mask_rows: list[Tensor] = []
|
||||||
for tokens, masks in rows:
|
marks_rows: list[Tensor] = []
|
||||||
|
for tokens, masks, marks in rows:
|
||||||
length = tokens.shape[1]
|
length = tokens.shape[1]
|
||||||
if masks is None:
|
if masks is None:
|
||||||
masks = torch.ones((1, length), dtype=torch.bool, device=tokens.device)
|
masks = torch.ones((1, length), dtype=torch.bool, device=tokens.device)
|
||||||
|
if has_marks and marks is None:
|
||||||
|
marks = torch.zeros((1, length), dtype=torch.bool, device=tokens.device)
|
||||||
if length < max_len:
|
if length < max_len:
|
||||||
pad = max_len - length
|
pad = max_len - length
|
||||||
tokens = torch.cat(
|
tokens = torch.cat(
|
||||||
@@ -1765,9 +1858,17 @@ class PI052Policy(PI05Policy):
|
|||||||
[masks, torch.zeros((1, pad), dtype=masks.dtype, device=masks.device)],
|
[masks, torch.zeros((1, pad), dtype=masks.dtype, device=masks.device)],
|
||||||
dim=1,
|
dim=1,
|
||||||
)
|
)
|
||||||
|
if has_marks:
|
||||||
|
marks = torch.cat(
|
||||||
|
[marks, torch.zeros((1, pad), dtype=marks.dtype, device=marks.device)],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
tok_rows.append(tokens)
|
tok_rows.append(tokens)
|
||||||
mask_rows.append(masks)
|
mask_rows.append(masks)
|
||||||
return torch.cat(tok_rows, dim=0), torch.cat(mask_rows, dim=0)
|
if has_marks:
|
||||||
|
marks_rows.append(marks)
|
||||||
|
stacked_marks = torch.cat(marks_rows, dim=0) if has_marks else None
|
||||||
|
return torch.cat(tok_rows, dim=0), torch.cat(mask_rows, dim=0), stacked_marks
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _fallback_subtask_from_task(task: str) -> str:
|
def _fallback_subtask_from_task(task: str) -> str:
|
||||||
@@ -1915,4 +2016,22 @@ class PI052Policy(PI05Policy):
|
|||||||
if self.config.use_flashrt_fp8_mlp and not getattr(self, "_fp8_applied", False):
|
if self.config.use_flashrt_fp8_mlp and not getattr(self, "_fp8_applied", False):
|
||||||
self._fp8_applied = True
|
self._fp8_applied = True
|
||||||
self.apply_flashrt_fp8_mlp(batch)
|
self.apply_flashrt_fp8_mlp(batch)
|
||||||
return super().predict_action_chunk(batch, **kwargs)
|
marks = batch.get(OBS_LANGUAGE_CAUSAL_MARKS)
|
||||||
|
if marks is None:
|
||||||
|
return super().predict_action_chunk(batch, **kwargs)
|
||||||
|
return self._predict_action_chunk_with_marks(batch, marks, **kwargs)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def _predict_action_chunk_with_marks(
|
||||||
|
self, batch: dict[str, Tensor], marks: Tensor, **kwargs: Unpack[ActionSelectKwargs]
|
||||||
|
) -> Tensor:
|
||||||
|
"""Base ``predict_action_chunk`` plus causal marks on the generated-subtask span."""
|
||||||
|
self.eval()
|
||||||
|
images, img_masks = self._preprocess_images(batch)
|
||||||
|
tokens = batch[OBS_LANGUAGE_TOKENS]
|
||||||
|
masks = batch[OBS_LANGUAGE_ATTENTION_MASK]
|
||||||
|
actions = self.model.sample_actions(
|
||||||
|
images, img_masks, tokens, masks, lang_causal_marks=marks, **kwargs
|
||||||
|
)
|
||||||
|
original_action_dim = self.config.output_features[ACTION].shape[0]
|
||||||
|
return actions[:, :, :original_action_dim]
|
||||||
|
|||||||
@@ -251,6 +251,44 @@ def _format_messages(
|
|||||||
return "".join(parts), spans
|
return "".join(parts), spans
|
||||||
|
|
||||||
|
|
||||||
|
def encode_prompt_with_targets(
|
||||||
|
tokenizer: Any, messages: list[dict[str, Any]], target_indices: list[int]
|
||||||
|
) -> tuple[Tensor, Tensor, Tensor]:
|
||||||
|
"""Tokenize a flat prompt and mark the token positions of target spans.
|
||||||
|
|
||||||
|
Inference-side twin of ``PI052TextTokenizerStep._encode_messages``: same
|
||||||
|
serialization (role headers, target EOS) and the same offset-overlap span
|
||||||
|
arithmetic, but unpadded and returning a boolean target mask instead of
|
||||||
|
labels. Used to rebuild joint-sequence prompts whose target spans must be
|
||||||
|
attended causally, matching ``_mark_target_span_causal`` at train time.
|
||||||
|
|
||||||
|
Returns ``(input_ids, attention_mask, target_marks)``, each ``(1, L)``.
|
||||||
|
"""
|
||||||
|
prompt, spans = _format_messages(messages, target_indices, getattr(tokenizer, "eos_token", None))
|
||||||
|
encoded = tokenizer(prompt, return_tensors="pt", return_offsets_mapping=True)
|
||||||
|
input_ids = encoded["input_ids"][0]
|
||||||
|
attention_mask = encoded.get("attention_mask")
|
||||||
|
if attention_mask is None:
|
||||||
|
attention_mask = torch.ones_like(input_ids, dtype=torch.bool)
|
||||||
|
else:
|
||||||
|
attention_mask = attention_mask[0].bool()
|
||||||
|
offsets = encoded["offset_mapping"][0]
|
||||||
|
|
||||||
|
marks = torch.zeros_like(input_ids, dtype=torch.bool)
|
||||||
|
for idx in target_indices:
|
||||||
|
if idx >= len(spans):
|
||||||
|
continue
|
||||||
|
char_start, char_end = spans[idx]
|
||||||
|
for token_pos in range(input_ids.shape[0]):
|
||||||
|
if not attention_mask[token_pos]:
|
||||||
|
continue
|
||||||
|
tok_start, tok_end = int(offsets[token_pos, 0]), int(offsets[token_pos, 1])
|
||||||
|
if tok_end <= char_start or tok_start >= char_end:
|
||||||
|
continue
|
||||||
|
marks[token_pos] = True
|
||||||
|
return input_ids.unsqueeze(0), attention_mask.unsqueeze(0), marks.unsqueeze(0)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@ProcessorStepRegistry.register(name="pi052_text_tokenizer")
|
@ProcessorStepRegistry.register(name="pi052_text_tokenizer")
|
||||||
class PI052TextTokenizerStep(ProcessorStep):
|
class PI052TextTokenizerStep(ProcessorStep):
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ OBS_IMAGES = OBS_IMAGE + "s"
|
|||||||
OBS_LANGUAGE = OBS_STR + ".language"
|
OBS_LANGUAGE = OBS_STR + ".language"
|
||||||
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
||||||
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
||||||
|
OBS_LANGUAGE_CAUSAL_MARKS = OBS_LANGUAGE + ".causal_marks"
|
||||||
OBS_LANGUAGE_SUBTASK = OBS_STR + ".subtask"
|
OBS_LANGUAGE_SUBTASK = OBS_STR + ".subtask"
|
||||||
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
|
OBS_LANGUAGE_SUBTASK_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
|
||||||
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
|
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
|
||||||
|
|||||||
@@ -0,0 +1,148 @@
|
|||||||
|
#!/usr/bin/env python
|
||||||
|
|
||||||
|
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Tests for PI052 joint-sequence (paper-style) subtask conditioning.
|
||||||
|
|
||||||
|
Joint recipes train the subtask text and the action losses in one sequence,
|
||||||
|
with the supervised subtask span attended causally. At inference the same
|
||||||
|
layout is rebuilt around the *generated* subtask, so these tests pin:
|
||||||
|
|
||||||
|
- the inference-side encoder produces the same token ids and target positions
|
||||||
|
as the training-time tokenizer step for the same messages;
|
||||||
|
- OR-ing causal marks into a prefix reproduces the training-time attention
|
||||||
|
pattern (prompt cannot see the subtask; subtask is causal over itself);
|
||||||
|
- the joint recipe file stays a valid message recipe;
|
||||||
|
- the FAST id mapping with the default ``fast_skip_tokens`` stays clear of
|
||||||
|
PaliGemma's ``<loc>`` range so VQA and FAST supervision never collide.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from lerobot.configs.recipe import TrainingRecipe
|
||||||
|
from lerobot.policies.pi052.text_processor_pi052 import (
|
||||||
|
PI052TextTokenizerStep,
|
||||||
|
encode_prompt_with_targets,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _CharTokenizer:
|
||||||
|
"""Char-level stub: 1 char = 1 token, so offsets are trivially aligned."""
|
||||||
|
|
||||||
|
pad_token_id = 0
|
||||||
|
eos_token = "\x1f" # unit separator — a 1-char "EOS" for testing
|
||||||
|
|
||||||
|
def __call__(self, text, max_length=None, padding=None, return_tensors=None, **kwargs):
|
||||||
|
limit = max_length if max_length is not None else len(text)
|
||||||
|
ids = [ord(c) % 251 + 1 for c in text[:limit]]
|
||||||
|
offsets = [(i, i + 1) for i in range(len(ids))]
|
||||||
|
attention = [1] * len(ids)
|
||||||
|
if padding == "max_length" and max_length is not None and len(ids) < max_length:
|
||||||
|
pad = max_length - len(ids)
|
||||||
|
ids += [self.pad_token_id] * pad
|
||||||
|
offsets += [(0, 0)] * pad
|
||||||
|
attention += [0] * pad
|
||||||
|
return {
|
||||||
|
"input_ids": torch.tensor([ids], dtype=torch.long),
|
||||||
|
"attention_mask": torch.tensor([attention], dtype=torch.long),
|
||||||
|
"offset_mapping": torch.tensor([offsets], dtype=torch.long),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
_MESSAGES = [
|
||||||
|
{"role": "user", "content": "fold the towel"},
|
||||||
|
{"role": "assistant", "content": "grab the near corner"},
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_encode_prompt_with_targets_matches_training_labels():
|
||||||
|
tokenizer = _CharTokenizer()
|
||||||
|
|
||||||
|
step = PI052TextTokenizerStep(max_length=120)
|
||||||
|
step._tokenizer = tokenizer
|
||||||
|
train_ids, train_attn, labels, predict_actions, _prompt = step._encode_messages(
|
||||||
|
tokenizer,
|
||||||
|
[dict(m) for m in _MESSAGES],
|
||||||
|
message_streams=["low_level", "low_level"],
|
||||||
|
target_indices=[1],
|
||||||
|
complementary={},
|
||||||
|
)
|
||||||
|
assert bool(predict_actions)
|
||||||
|
|
||||||
|
ids, attn, marks = encode_prompt_with_targets(tokenizer, [dict(m) for m in _MESSAGES], [1])
|
||||||
|
|
||||||
|
n = int(attn.sum())
|
||||||
|
assert n == int(train_attn.sum())
|
||||||
|
assert torch.equal(ids[0, :n], train_ids[:n])
|
||||||
|
# Causal marks at inference must cover exactly the supervised label span.
|
||||||
|
assert torch.equal(marks[0, :n], labels[:n] != -100)
|
||||||
|
assert marks.any(), "the assistant target span must be marked"
|
||||||
|
# The user turn must stay unmarked (bidirectional prompt).
|
||||||
|
user_len = len("User: fold the towel\n")
|
||||||
|
assert not marks[0, :user_len].any()
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_causal_language_marks_reproduces_training_mask():
|
||||||
|
from lerobot.policies.pi05.modeling_pi05 import make_att_2d_masks
|
||||||
|
from lerobot.policies.pi052.modeling_pi052 import _apply_causal_language_marks
|
||||||
|
|
||||||
|
n_img, n_lang = 4, 8
|
||||||
|
prefix_len = n_img + n_lang
|
||||||
|
pad = torch.ones((1, prefix_len), dtype=torch.bool)
|
||||||
|
att = torch.zeros((1, prefix_len), dtype=torch.bool)
|
||||||
|
# Subtask span = language positions 5..7 (prefix positions 9..11).
|
||||||
|
marks = torch.zeros((1, n_lang), dtype=torch.bool)
|
||||||
|
marks[0, 5:8] = True
|
||||||
|
|
||||||
|
att_marked = _apply_causal_language_marks(att, marks)
|
||||||
|
att_2d = make_att_2d_masks(pad, att_marked)[0]
|
||||||
|
|
||||||
|
subtask = [n_img + 5, n_img + 6, n_img + 7]
|
||||||
|
# Prompt and images never see the subtask.
|
||||||
|
for q in range(n_img + 5):
|
||||||
|
for k in subtask:
|
||||||
|
assert not att_2d[q, k], f"prompt position {q} must not attend subtask position {k}"
|
||||||
|
# Subtask tokens see the full prompt and earlier subtask tokens only.
|
||||||
|
for qi, q in enumerate(subtask):
|
||||||
|
for k in range(n_img + 5):
|
||||||
|
assert att_2d[q, k]
|
||||||
|
for ki, k in enumerate(subtask):
|
||||||
|
assert bool(att_2d[q, k]) == (ki <= qi)
|
||||||
|
|
||||||
|
|
||||||
|
def test_joint_recipe_is_a_valid_message_recipe():
|
||||||
|
recipe_path = Path(__file__).parents[3] / "src" / "lerobot" / "configs" / "recipes" / "subtask_joint.yaml"
|
||||||
|
recipe = TrainingRecipe.from_yaml(recipe_path)
|
||||||
|
assert recipe.messages is not None and len(recipe.messages) == 2
|
||||||
|
assert all(turn.stream == "low_level" for turn in recipe.messages)
|
||||||
|
assert not recipe.messages[0].target
|
||||||
|
assert recipe.messages[1].target
|
||||||
|
assert recipe.messages[1].if_present == "subtask"
|
||||||
|
|
||||||
|
|
||||||
|
def test_default_fast_mapping_clears_loc_and_seg_ranges():
|
||||||
|
from lerobot.policies.pi052.configuration_pi052 import PI052Config
|
||||||
|
from lerobot.policies.pi052.modeling_pi052 import _FAST_ACTION_VOCAB_SIZE
|
||||||
|
|
||||||
|
skip = PI052Config.__dataclass_fields__["fast_skip_tokens"].default
|
||||||
|
assert skip == 1152
|
||||||
|
|
||||||
|
paligemma_vocab = 257152
|
||||||
|
fast_ids = paligemma_vocab - 1 - skip - torch.arange(_FAST_ACTION_VOCAB_SIZE)
|
||||||
|
# Below the <loc> range [256000, 257024) and the <seg> range [257024, 257152).
|
||||||
|
assert int(fast_ids.max()) < 256000
|
||||||
|
assert int(fast_ids.min()) >= 0
|
||||||
Reference in New Issue
Block a user