mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
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:
@@ -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"])
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user