test(collate): expect preserved language columns

This commit is contained in:
Pepijn
2026-07-28 12:39:15 +02:00
parent 3f093d8927
commit 913664c320
+5 -5
View File
@@ -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():