feat(data): add recipe-driven language supervision (#4182)

* feat(data): add recipe-driven language supervision

* test(collate): expect preserved language columns

* Address PR review feedback

* Address Claude review feedback
This commit is contained in:
Pepijn
2026-08-04 16:48:47 +02:00
committed by GitHub
parent f66e5128ec
commit 64b23178d5
28 changed files with 1062 additions and 79 deletions
@@ -4,15 +4,22 @@ import pytest
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
import numpy as np # noqa: E402
import torch # noqa: E402
from lerobot.configs.recipe import MessageTurn, TrainingRecipe # noqa: E402
from lerobot.lerobot_types import TransitionKey # noqa: E402
from lerobot.processor.converters import create_transition # noqa: E402
from lerobot.processor.render_messages_processor import RenderMessagesStep # noqa: E402
from lerobot.processor.render_messages_processor import ( # noqa: E402
RenderMessagesStep,
_fallback_low_level_render,
_select_batch_indices,
)
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 +28,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 +83,129 @@ 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"] == [[], []]
def test_render_messages_step_rejects_mismatched_non_empty_language_batches():
recipe = TrainingRecipe(
messages=[
MessageTurn(
role="assistant",
content="${subtask}",
stream="high_level",
target=True,
if_present="subtask",
)
]
)
transition = create_transition(
complementary_data={
"timestamp": torch.tensor([0.0, 1.0, 2.0]),
"language_persistent": [[], []],
"language_events": [[{"style": "unmatched"}], [], []],
}
)
with pytest.raises(ValueError, match="must have equal lengths"):
RenderMessagesStep(recipe)(transition)
def test_select_batch_indices_slices_numpy_action():
action = np.arange(6).reshape(3, 2)
transition = create_transition(action=action)
selected = _select_batch_indices(transition, [2, 0], batch_size=3)
np.testing.assert_array_equal(selected[TransitionKey.ACTION], action[[2, 0]])
def test_select_batch_indices_slices_robot_action_dict():
transition = create_transition(
action={
"joints": np.arange(6).reshape(3, 2),
"gripper": torch.tensor([[0.0], [1.0], [2.0]]),
}
)
selected = _select_batch_indices(transition, [2, 0], batch_size=3)
np.testing.assert_array_equal(selected[TransitionKey.ACTION]["joints"], np.array([[4, 5], [0, 1]]))
assert torch.equal(selected[TransitionKey.ACTION]["gripper"], torch.tensor([[2.0], [0.0]]))
def test_select_batch_indices_rejects_misaligned_list():
transition = create_transition(complementary_data={"task": ["one", "two"]})
with pytest.raises(ValueError, match="expected 3 values, got 2"):
_select_batch_indices(transition, [2, 0], batch_size=3)
def test_fallback_low_level_render_rejects_partially_missing_task_batch():
with pytest.raises(ValueError, match=r"missing task at indices \[1\]"):
_fallback_low_level_render(["pick cube", None, "place cube"])
+54 -5
View File
@@ -19,6 +19,7 @@ Tests for the TokenizerProcessorStep class.
"""
import tempfile
from pathlib import Path
from unittest.mock import patch
import pytest
@@ -26,7 +27,7 @@ import torch
from lerobot.configs.types import FeatureType, PipelineFeatureType, PolicyFeature
from lerobot.lerobot_types import TransitionKey
from lerobot.processor import DataProcessorPipeline, TokenizerProcessorStep
from lerobot.processor import ActionTokenizerProcessorStep, DataProcessorPipeline, TokenizerProcessorStep
from lerobot.processor.converters import create_transition, identity_transition
from lerobot.utils.constants import (
ACTION,
@@ -87,6 +88,51 @@ class MockTokenizer:
return result
def save_pretrained(self, save_directory: str | Path) -> None:
save_directory = Path(save_directory)
save_directory.mkdir(parents=True, exist_ok=True)
(save_directory / "tokenizer_config.json").write_text("{}")
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():
@@ -490,9 +536,11 @@ def test_save_and_load_pretrained_with_tokenizer_name(mock_auto_tokenizer):
@skip_if_package_missing("transformers")
def test_save_and_load_pretrained_with_tokenizer_object():
"""Test saving and loading processor with tokenizer object using overrides."""
@patch("lerobot.processor.tokenizer_processor.AutoTokenizer")
def test_save_and_load_pretrained_with_tokenizer_object(mock_auto_tokenizer):
"""Test that a tokenizer object is saved and reloads from its local artifact."""
mock_tokenizer = MockTokenizer(vocab_size=100)
mock_auto_tokenizer.from_pretrained.return_value = mock_tokenizer
original_processor = TokenizerProcessorStep(
tokenizer=mock_tokenizer, max_length=32, task_key="instruction"
@@ -506,11 +554,12 @@ def test_save_and_load_pretrained_with_tokenizer_object():
# Save processor
robot_processor.save_pretrained(temp_dir)
# Load processor with tokenizer override (since tokenizer object wasn't saved)
assert (Path(temp_dir) / "tokenizer" / "tokenizer_config.json").is_file()
# Load processor without an object override: the saved artifact is portable.
loaded_processor = DataProcessorPipeline.from_pretrained(
temp_dir,
config_filename="dataprocessorpipeline.json",
overrides={"tokenizer_processor": {"tokenizer": mock_tokenizer}},
to_transition=identity_transition,
to_output=identity_transition,
)