diff --git a/tests/utils/test_collate.py b/tests/utils/test_collate.py index 2b23b3180..94d87cbcf 100644 --- a/tests/utils/test_collate.py +++ b/tests/utils/test_collate.py @@ -9,7 +9,7 @@ import torch # noqa: E402 from lerobot.utils.collate import lerobot_collate_fn # noqa: E402 -def test_lerobot_collate_preserves_messages_and_drops_raw_language(): +def test_lerobot_collate_preserves_messages_and_raw_language(): batch = [ { "index": torch.tensor(0), @@ -17,14 +17,14 @@ def test_lerobot_collate_preserves_messages_and_drops_raw_language(): "message_streams": ["low_level"], "target_message_indices": [0], "language_persistent": [{"content": "raw"}], - "language_events": [], + "language_events": [{"content": "event a"}], }, { "index": torch.tensor(1), "messages": [{"role": "assistant", "content": "b"}], "message_streams": ["low_level"], "target_message_indices": [0], - "language_persistent": [{"content": "raw"}], + "language_persistent": [{"content": "raw b"}], "language_events": [], }, ] @@ -36,8 +36,8 @@ def test_lerobot_collate_preserves_messages_and_drops_raw_language(): assert out["messages"][1][0]["content"] == "b" assert out["message_streams"] == [["low_level"], ["low_level"]] assert out["target_message_indices"] == [[0], [0]] - assert "language_persistent" not in out - assert "language_events" not in out + assert out["language_persistent"] == [[{"content": "raw"}], [{"content": "raw b"}]] + assert out["language_events"] == [[{"content": "event a"}], []] def test_lerobot_collate_passes_through_standard_batch():