diff --git a/docs/source/language_and_recipes.mdx b/docs/source/language_and_recipes.mdx
index 4181dbe34..8d2aed59d 100644
--- a/docs/source/language_and_recipes.mdx
+++ b/docs/source/language_and_recipes.mdx
@@ -141,6 +141,17 @@ 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.
+
+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
If both language columns are missing, `None`, or empty, `RenderMessagesStep` is a no-op.
diff --git a/src/lerobot/configs/default.py b/src/lerobot/configs/default.py
index 38991a665..32134ccc9 100644
--- a/src/lerobot/configs/default.py
+++ b/src/lerobot/configs/default.py
@@ -33,6 +33,8 @@ 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
+ # Episode indices to drop (e.g. corrupt or heterogeneous ones). Applied on top of `episodes`.
+ exclude_episodes: list[int] | None = None
image_transforms: ImageTransformsConfig = field(default_factory=ImageTransformsConfig)
revision: str | None = None
use_imagenet_stats: bool = True
@@ -62,6 +64,10 @@ class DatasetConfig:
if len(self.episodes) != len(set(self.episodes)):
duplicates = sorted({ep for ep in self.episodes if self.episodes.count(ep) > 1})
raise ValueError(f"Episode indices contain duplicates: {duplicates}")
+ if self.exclude_episodes is not None and any(ep < 0 for ep in self.exclude_episodes):
+ raise ValueError(
+ f"exclude_episodes must be non-negative, got: {[ep for ep in self.exclude_episodes if ep < 0]}"
+ )
@dataclass
diff --git a/src/lerobot/configs/recipe.py b/src/lerobot/configs/recipe.py
index 28e5a0db3..43d90c1ce 100644
--- a/src/lerobot/configs/recipe.py
+++ b/src/lerobot/configs/recipe.py
@@ -78,7 +78,7 @@ class MessageTurn:
raise ValueError(f"Unsupported message stream: {self.stream!r}")
if self.content is None and self.tool_calls_from is None:
raise ValueError("MessageTurn.content is required unless tool_calls_from is set.")
- if self.content is not None and not isinstance(self.content, (str, list)):
+ if self.content is not None and not isinstance(self.content, str | list):
raise TypeError("MessageTurn.content must be a string, a list of HF-style blocks, or None.")
if isinstance(self.content, list):
for block in self.content:
@@ -147,7 +147,7 @@ class TrainingRecipe:
return cls.from_dict(data)
def _validate_message_recipe(self) -> None:
- """Ensure every templated binding is known and at least one turn is a target."""
+ """Validate bindings and require text or low-level action supervision."""
assert self.messages is not None
known_bindings = set(DEFAULT_BINDINGS) | set(self.bindings or {}) | {"task"}
@@ -156,8 +156,14 @@ class TrainingRecipe:
if missing:
raise ValueError(f"MessageTurn references unknown binding(s): {sorted(missing)}")
- if not any(turn.target for turn in self.messages):
- raise ValueError("Message recipes must contain at least one target turn.")
+ 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)."
+ )
def _validate_blend_recipe(self) -> None:
"""Ensure each blend component is a non-empty, weighted message recipe."""
diff --git a/src/lerobot/configs/recipes/subtask.yaml b/src/lerobot/configs/recipes/subtask.yaml
new file mode 100644
index 000000000..c90ca8f78
--- /dev/null
+++ b/src/lerobot/configs/recipes/subtask.yaml
@@ -0,0 +1,16 @@
+# 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}
diff --git a/src/lerobot/configs/recipes/subtask_joint.yaml b/src/lerobot/configs/recipes/subtask_joint.yaml
new file mode 100644
index 000000000..5258ccf23
--- /dev/null
+++ b/src/lerobot/configs/recipes/subtask_joint.yaml
@@ -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}
diff --git a/src/lerobot/configs/recipes/subtask_mem.yaml b/src/lerobot/configs/recipes/subtask_mem.yaml
new file mode 100644
index 000000000..d12fe5009
--- /dev/null
+++ b/src/lerobot/configs/recipes/subtask_mem.yaml
@@ -0,0 +1,30 @@
+# 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}
diff --git a/src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml b/src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml
new file mode 100644
index 000000000..1a3d39952
--- /dev/null
+++ b/src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml
@@ -0,0 +1,70 @@
+# 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 `...`.
+
+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 `...` 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}
diff --git a/src/lerobot/datasets/dataset_reader.py b/src/lerobot/datasets/dataset_reader.py
index f4e1f6a31..4b32eb66b 100644
--- a/src/lerobot/datasets/dataset_reader.py
+++ b/src/lerobot/datasets/dataset_reader.py
@@ -163,10 +163,40 @@ 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:
diff --git a/src/lerobot/datasets/factory.py b/src/lerobot/datasets/factory.py
index da7b4365a..ad7f207ec 100644
--- a/src/lerobot/datasets/factory.py
+++ b/src/lerobot/datasets/factory.py
@@ -66,6 +66,17 @@ def resolve_delta_timestamps(
return delta_timestamps
+def _resolve_episodes(
+ episodes: list[int] | None, exclude_episodes: list[int] | None, total_episodes: int
+) -> list[int] | None:
+ """Apply an episode exclusion list on top of an optional allowlist."""
+ if not exclude_episodes:
+ return episodes
+ base = episodes if episodes is not None else list(range(total_episodes))
+ excluded = set(exclude_episodes)
+ return [episode for episode in base if episode not in excluded]
+
+
def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDataset:
"""Handles the logic of setting up delta timestamps and image transforms before creating a dataset.
@@ -87,11 +98,14 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
cfg.dataset.repo_id, root=cfg.dataset.root, revision=cfg.dataset.revision
)
delta_timestamps = resolve_delta_timestamps(cfg.trainable_config, ds_meta)
+ episodes = _resolve_episodes(
+ cfg.dataset.episodes, cfg.dataset.exclude_episodes, ds_meta.total_episodes
+ )
if not cfg.dataset.streaming:
dataset = LeRobotDataset(
cfg.dataset.repo_id,
root=cfg.dataset.root,
- episodes=cfg.dataset.episodes,
+ episodes=episodes,
delta_timestamps=delta_timestamps,
image_transforms=image_transforms,
revision=cfg.dataset.revision,
@@ -104,7 +118,7 @@ def make_dataset(cfg: TrainPipelineConfig) -> LeRobotDataset | MultiLeRobotDatas
dataset = StreamingLeRobotDataset(
cfg.dataset.repo_id,
root=cfg.dataset.root,
- episodes=cfg.dataset.episodes,
+ episodes=episodes,
delta_timestamps=delta_timestamps,
image_transforms=image_transforms,
revision=cfg.dataset.revision,
diff --git a/src/lerobot/datasets/language_render.py b/src/lerobot/datasets/language_render.py
index 999fa19ad..624f19e9b 100644
--- a/src/lerobot/datasets/language_render.py
+++ b/src/lerobot/datasets/language_render.py
@@ -162,14 +162,28 @@ def render_sample(
task: str | None = None,
dataset_ctx: Any | None = None,
) -> RenderedMessages | None:
- """Render the chat-style messages for a single dataset sample.
+ """Resolve one sample's bindings and render its message recipe.
- 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.
+ Returns ``None`` when no text or low-level action supervision applies.
"""
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,
@@ -183,6 +197,55 @@ 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:
@@ -346,7 +409,9 @@ def _render_message_recipe(
if turn.target:
target_indices.append(message_idx)
- if not target_indices:
+ # 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:
return None
rendered = {
@@ -403,14 +468,12 @@ def _validate_rendered(rendered: RenderedMessages) -> None:
if len(streams) != len(messages):
raise ValueError("message_streams must be aligned with messages.")
- if not target_indices:
- raise ValueError("Rendered samples must contain at least one target message.")
+ # 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.")
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(
diff --git a/src/lerobot/processor/batch_processor.py b/src/lerobot/processor/batch_processor.py
index 669c68a0a..804a3aaf0 100644
--- a/src/lerobot/processor/batch_processor.py
+++ b/src/lerobot/processor/batch_processor.py
@@ -175,9 +175,6 @@ 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)):
diff --git a/src/lerobot/processor/pipeline.py b/src/lerobot/processor/pipeline.py
index e40a7c479..b30493db9 100644
--- a/src/lerobot/processor/pipeline.py
+++ b/src/lerobot/processor/pipeline.py
@@ -41,7 +41,7 @@ from pathlib import Path
from typing import Any, TypedDict, TypeVar, cast
import torch
-from huggingface_hub import hf_hub_download
+from huggingface_hub import hf_hub_download, snapshot_download
from safetensors.torch import load_file, save_file
from lerobot.configs import PipelineFeatureType, PolicyFeature
@@ -205,6 +205,10 @@ class ProcessorStep(ABC):
"""
return None
+ def save_artifacts(self, save_directory: Path) -> dict[str, str]:
+ """Save non-tensor assets and map constructor arguments to relative paths."""
+ return {}
+
def reset(self) -> None:
"""Resets the internal state of the processor step, if any."""
return None
@@ -549,6 +553,22 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
pipeline_config = self.get_config()
pipeline_state_dict = self.state_dict()
+ for processor_step, step_entry in zip(self.steps, pipeline_config["steps"], strict=True):
+ artifacts = processor_step.save_artifacts(save_directory)
+ if artifacts:
+ for config_key, relative_path in artifacts.items():
+ artifact_path = Path(relative_path)
+ if artifact_path.is_absolute() or ".." in artifact_path.parts:
+ raise ValueError(
+ f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
+ )
+ if not (save_directory / artifact_path).exists():
+ raise FileNotFoundError(
+ f"Processor step did not save declared artifact '{relative_path}'"
+ )
+ step_entry["config"][config_key] = artifact_path.as_posix()
+ step_entry["artifacts"] = artifacts
+
for state_key, step_state_dict in pipeline_state_dict.items():
state_filename = f"{state_key}.safetensors"
save_file(step_state_dict, save_directory / state_filename)
@@ -733,7 +753,13 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
# 3. Build steps with overrides
steps, validated_overrides = cls._build_steps_with_overrides(
- loaded_config, overrides or {}, model_id, base_path, hub_download_kwargs, is_local_source
+ loaded_config,
+ overrides or {},
+ model_id,
+ base_path,
+ config_filename,
+ hub_download_kwargs,
+ is_local_source,
)
# 4. Validate that all overrides were used
@@ -922,6 +948,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
overrides: dict[str, Any],
model_id: str,
base_path: Path | None,
+ config_filename: str,
hub_download_kwargs: dict[str, Any],
is_local_source: bool = False,
) -> tuple[list[ProcessorStep], set[str]]:
@@ -976,15 +1003,68 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
ImportError: If a step class cannot be imported or found in registry
ValueError: If a step cannot be instantiated with its configuration
"""
+ loaded_config = deepcopy(loaded_config)
+ cls._resolve_artifact_paths(
+ loaded_config,
+ model_id,
+ base_path,
+ config_filename,
+ hub_download_kwargs,
+ )
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
cls._load_step_state(
- step_instance, step_entry, model_id, base_path, hub_download_kwargs, is_local_source
+ step_instance,
+ step_entry,
+ model_id,
+ base_path,
+ config_filename,
+ hub_download_kwargs,
+ is_local_source,
)
return steps, remaining_override_keys
+ @classmethod
+ def _resolve_artifact_paths(
+ cls,
+ loaded_config: dict[str, Any],
+ model_id: str,
+ base_path: Path | None,
+ config_filename: str,
+ hub_download_kwargs: dict[str, Any],
+ ) -> None:
+ """Resolve declared relative processor artifacts before step construction."""
+ is_local = Path(model_id).is_dir() or Path(model_id).is_file()
+
+ for step_entry in loaded_config["steps"]:
+ artifacts = step_entry.get("artifacts", {})
+ for config_key, relative_path in artifacts.items():
+ artifact_path = Path(relative_path)
+ if artifact_path.is_absolute() or ".." in artifact_path.parts:
+ raise ValueError(
+ f"Processor artifact path must be relative to the checkpoint: {relative_path!r}"
+ )
+
+ resolved_path = base_path / artifact_path if base_path is not None else artifact_path
+ if not resolved_path.exists() and not is_local:
+ repository_path = Path(config_filename).parent / artifact_path
+ snapshot_download(
+ repo_id=model_id,
+ repo_type="model",
+ allow_patterns=f"{repository_path.as_posix()}/**",
+ **hub_download_kwargs,
+ )
+
+ if not resolved_path.exists():
+ step_name = step_entry.get("registry_name", step_entry.get("class", "unknown"))
+ raise FileNotFoundError(
+ f"Missing processor artifact '{relative_path}' for step '{step_name}' "
+ f"next to '{config_filename}'. Checkpoint artifacts are incomplete."
+ )
+ step_entry["config"][config_key] = str(resolved_path)
+
@classmethod
def _build_steps_from_config(
cls,
@@ -1144,6 +1224,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
step_entry: dict[str, Any],
model_id: str,
base_path: Path | None,
+ config_filename: str,
hub_download_kwargs: dict[str, Any],
is_local_source: bool = False,
) -> None:
@@ -1209,7 +1290,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
# Download from Hub
state_path = hf_hub_download(
repo_id=model_id,
- filename=state_filename,
+ filename=(Path(config_filename).parent / state_filename).as_posix(),
repo_type="model",
**hub_download_kwargs,
)
diff --git a/src/lerobot/processor/render_messages_processor.py b/src/lerobot/processor/render_messages_processor.py
index 140592f0e..2fca46e7e 100644
--- a/src/lerobot/processor/render_messages_processor.py
+++ b/src/lerobot/processor/render_messages_processor.py
@@ -16,7 +16,7 @@
from __future__ import annotations
-from dataclasses import dataclass
+from dataclasses import asdict, dataclass
from typing import Any
from lerobot.configs import PipelineFeatureType, PolicyFeature
@@ -32,17 +32,18 @@ from .pipeline import ProcessorStep, ProcessorStepRegistry
@dataclass
@ProcessorStepRegistry.register(name="render_messages_processor")
class RenderMessagesStep(ProcessorStep):
- """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.
- """
+ """Render language columns into recipe-defined messages and supervision metadata."""
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 {}
@@ -50,7 +51,17 @@ class RenderMessagesStep(ProcessorStep):
events = complementary_data.get(LANGUAGE_EVENTS) or []
if not persistent and not events:
- return transition
+ 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)
timestamp = complementary_data.get("timestamp")
if timestamp is None:
@@ -67,18 +78,147 @@ class RenderMessagesStep(ProcessorStep):
dataset_ctx=self.dataset_ctx,
)
if rendered is None:
- return None
+ rendered = _fallback_low_level_render(complementary_data.get("task"))
+ if rendered is None:
+ return None
new_transition = transition.copy()
- new_complementary_data = dict(complementary_data)
+ 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.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": [],
+ }
diff --git a/src/lerobot/processor/tokenizer_processor.py b/src/lerobot/processor/tokenizer_processor.py
index a808e6127..71b5ccee3 100644
--- a/src/lerobot/processor/tokenizer_processor.py
+++ b/src/lerobot/processor/tokenizer_processor.py
@@ -25,6 +25,7 @@ from __future__ import annotations
import logging
from dataclasses import dataclass, field
+from pathlib import Path
from typing import TYPE_CHECKING, Any
import torch
@@ -32,6 +33,7 @@ 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,
@@ -136,7 +138,7 @@ class TokenizerProcessorStep(ObservationProcessorStep):
# Standardize to a list of strings for the tokenizer
if isinstance(task, str):
return [task]
- elif isinstance(task, (list, tuple)) and all(isinstance(t, str) for t in task):
+ elif isinstance(task, list | tuple) and all(isinstance(t, str) for t in task):
return list(task)
return None
@@ -349,6 +351,7 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
max_action_tokens: int = 256
fast_skip_tokens: int = 128
paligemma_tokenizer_name: str = "google/paligemma-3b-pt-224"
+ allow_truncation: bool = True
# Internal tokenizer instance (not part of the config)
action_tokenizer: Any = field(default=None, init=False, repr=False)
_paligemma_tokenizer: Any = field(default=None, init=False, repr=False)
@@ -412,14 +415,15 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
# During inference, no action is available, skip tokenization
return new_transition
- # Tokenize and get both tokens and mask
- tokens, mask = self._tokenize_action(action)
+ # Tokenize and get masks for the full formatted sequence and the discrete action codes.
+ tokens, mask, code_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
@@ -430,7 +434,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]:
+ def _tokenize_action(self, action: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Tokenizes the action tensor and creates a mask.
@@ -459,6 +463,7 @@ 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)
@@ -476,65 +481,82 @@ 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
- # add bos
+ 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)
tokens = torch.cat(
[
torch.tensor([bos_id], device=action.device),
- 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),
+ prompt_tokens,
+ action_code_tokens,
+ end_tokens,
]
)
+ 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:
+ if not self.allow_truncation:
+ raise ValueError(
+ f"FAST action sequence has {len(tokens)} tokens, exceeding "
+ f"max_action_tokens={self.max_action_tokens}."
+ )
logging.warning(
f"Token length ({len(tokens)}) exceeds max length ({self.max_action_tokens}), truncating. "
"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(
- self.max_action_tokens - len(tokens), dtype=torch.bool, device=action.device
- ),
+ torch.zeros(pad_len, 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, self.max_action_tokens - len(tokens)), value=0)
+ tokens = torch.nn.functional.pad(tokens, (0, pad_len), 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
+ return tokens_batch, masks_batch, code_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]:
@@ -550,6 +572,9 @@ 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,
+ "allow_truncation": self.allow_truncation,
}
# Only save tokenizer_name if it was used to create the tokenizer
@@ -558,6 +583,14 @@ class ActionTokenizerProcessorStep(ActionProcessorStep):
return config
+ def save_artifacts(self, save_directory: Path) -> dict[str, str]:
+ artifact_path = Path("action_tokenizer")
+ save_pretrained = getattr(self.action_tokenizer, "save_pretrained", None)
+ if save_pretrained is None:
+ raise TypeError("Action tokenizer must implement save_pretrained() to save a portable pipeline.")
+ save_pretrained(save_directory / artifact_path)
+ return {"action_tokenizer_name": artifact_path.as_posix()}
+
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
diff --git a/src/lerobot/utils/collate.py b/src/lerobot/utils/collate.py
index fce7e6b42..d159a9b23 100644
--- a/src/lerobot/utils/collate.py
+++ b/src/lerobot/utils/collate.py
@@ -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"}
+_PYTHON_LIST_KEYS = {"messages", "message_streams", "target_message_indices", *LANGUAGE_COLUMNS}
def lerobot_collate_fn(batch: list[dict[str, Any] | None]) -> dict[str, Any] | None:
diff --git a/src/lerobot/utils/constants.py b/src/lerobot/utils/constants.py
index 8f735fe6d..fae865b16 100644
--- a/src/lerobot/utils/constants.py
+++ b/src/lerobot/utils/constants.py
@@ -26,6 +26,7 @@ OBS_IMAGES = OBS_IMAGE + "s"
OBS_LANGUAGE = OBS_STR + ".language"
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
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_TOKENS = OBS_LANGUAGE_SUBTASK + ".tokens"
OBS_LANGUAGE_SUBTASK_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK + ".attention_mask"
@@ -34,6 +35,7 @@ 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"
diff --git a/tests/configs/test_recipe.py b/tests/configs/test_recipe.py
index b4954efbf..53520bfa3 100644
--- a/tests/configs/test_recipe.py
+++ b/tests/configs/test_recipe.py
@@ -29,6 +29,13 @@ 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.
diff --git a/tests/datasets/test_language_render.py b/tests/datasets/test_language_render.py
index fcef41fd8..4ab4a2f2c 100644
--- a/tests/datasets/test_language_render.py
+++ b/tests/datasets/test_language_render.py
@@ -343,6 +343,84 @@ 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
diff --git a/tests/datasets/test_sampler.py b/tests/datasets/test_sampler.py
index 7a5fc0fe0..dbe2eb7f0 100644
--- a/tests/datasets/test_sampler.py
+++ b/tests/datasets/test_sampler.py
@@ -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
+from lerobot.datasets.sampler import EpisodeAwareSampler, compute_sampler_state
def calculate_episode_data_index(hf_dataset: Dataset) -> dict[str, torch.Tensor]:
@@ -154,8 +154,6 @@ 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
diff --git a/tests/processor/test_render_messages_processor.py b/tests/processor/test_render_messages_processor.py
index f96e3c0ab..783c2f7d8 100644
--- a/tests/processor/test_render_messages_processor.py
+++ b/tests/processor/test_render_messages_processor.py
@@ -12,7 +12,9 @@ from lerobot.processor.render_messages_processor import RenderMessagesStep # no
from lerobot.types import TransitionKey # noqa: E402
-def test_render_messages_step_noops_without_language_columns():
+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."""
recipe = TrainingRecipe(
messages=[
MessageTurn(role="user", content="${task}", stream="high_level"),
@@ -21,6 +23,24 @@ def test_render_messages_step_noops_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
@@ -58,3 +78,70 @@ 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"] == [[], []]
diff --git a/tests/processor/test_tokenizer_processor.py b/tests/processor/test_tokenizer_processor.py
index 5708e6e81..aab0c14ec 100644
--- a/tests/processor/test_tokenizer_processor.py
+++ b/tests/processor/test_tokenizer_processor.py
@@ -25,7 +25,7 @@ import pytest
import torch
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
-from lerobot.processor import DataProcessorPipeline, TokenizerProcessorStep
+from lerobot.processor import ActionTokenizerProcessorStep, DataProcessorPipeline, TokenizerProcessorStep
from lerobot.processor.converters import create_transition, identity_transition
from lerobot.types import TransitionKey
from lerobot.utils.constants import (
@@ -88,6 +88,46 @@ 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.allow_truncation = False
+ 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",
+ "allow_truncation": False,
+ "action_tokenizer_name": "custom/fast",
+ }
+
+
+def test_action_tokenizer_can_reject_truncated_sequences():
+ processor = object.__new__(ActionTokenizerProcessorStep)
+ processor.max_action_tokens = 4
+ processor.fast_skip_tokens = 128
+ processor.allow_truncation = False
+ processor.action_tokenizer = lambda _actions: [1, 2, 3]
+ processor._paligemma_tokenizer = type(
+ "Tokenizer",
+ (),
+ {
+ "vocab_size": 1000,
+ "bos_token_id": 2,
+ "encode": lambda _self, text, **_kwargs: [10, 11] if text == "Action: " else [12, 1],
+ },
+ )()
+
+ with pytest.raises(ValueError, match="max_action_tokens=4"):
+ processor._tokenize_action(torch.zeros(1, 2, 1))
+
+
@pytest.fixture
def mock_tokenizer():
"""Provide a mock tokenizer for testing."""