From 3f093d89273ba935a4cec1f5747cc418623f5db6 Mon Sep 17 00:00:00 2001 From: Pepijn Date: Tue, 28 Jul 2026 10:04:28 +0200 Subject: [PATCH] feat(data): add recipe-driven language supervision --- docs/source/language_and_recipes.mdx | 11 ++ src/lerobot/configs/default.py | 6 + src/lerobot/configs/recipe.py | 14 +- src/lerobot/configs/recipes/subtask.yaml | 16 ++ .../configs/recipes/subtask_joint.yaml | 13 ++ src/lerobot/configs/recipes/subtask_mem.yaml | 30 ++++ .../recipes/subtask_mem_vqa_speech.yaml | 70 ++++++++ src/lerobot/datasets/dataset_reader.py | 30 ++++ src/lerobot/datasets/factory.py | 18 +- src/lerobot/datasets/language_render.py | 83 +++++++-- src/lerobot/processor/batch_processor.py | 3 - src/lerobot/processor/pipeline.py | 89 +++++++++- .../processor/render_messages_processor.py | 162 ++++++++++++++++-- src/lerobot/processor/tokenizer_processor.py | 67 ++++++-- src/lerobot/utils/collate.py | 2 +- src/lerobot/utils/constants.py | 2 + tests/configs/test_recipe.py | 7 + tests/datasets/test_language_render.py | 78 +++++++++ tests/datasets/test_sampler.py | 4 +- .../test_render_messages_processor.py | 89 +++++++++- tests/processor/test_tokenizer_processor.py | 42 ++++- 21 files changed, 779 insertions(+), 57 deletions(-) create mode 100644 src/lerobot/configs/recipes/subtask.yaml create mode 100644 src/lerobot/configs/recipes/subtask_joint.yaml create mode 100644 src/lerobot/configs/recipes/subtask_mem.yaml create mode 100644 src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml 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."""