mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 21:19:40 +00:00
feat(task modification precedence): Improving task modification precedence so that all modes can be used in a single run. Adapting tests accordinginly.
This commit is contained in:
@@ -1161,6 +1161,20 @@ def test_modify_tasks_replacements_with_episode_overrides(sample_dataset):
|
||||
assert len(modified_dataset.meta.tasks) == 3
|
||||
|
||||
|
||||
def test_modify_tasks_default_task_and_replacements(sample_dataset):
|
||||
"""Test that new_task acts as the default for episodes not matched by task_replacements."""
|
||||
modified_dataset = modify_tasks(
|
||||
sample_dataset,
|
||||
new_task="Default task",
|
||||
task_replacements={"task_0": "Pick the cube"},
|
||||
)
|
||||
|
||||
for ep_idx in range(5):
|
||||
expected_task = "Pick the cube" if ep_idx % 2 == 0 else "Default task"
|
||||
assert modified_dataset.meta.episodes[ep_idx]["tasks"][0] == expected_task
|
||||
assert len(modified_dataset.meta.tasks) == 2
|
||||
|
||||
|
||||
def test_modify_tasks_no_task_specified(sample_dataset):
|
||||
"""Test error when no task is specified."""
|
||||
with pytest.raises(ValueError, match="Must specify at least one of new_task, episode_tasks, or task_replacements"):
|
||||
@@ -1179,16 +1193,6 @@ def test_modify_tasks_invalid_task_replacements(sample_dataset):
|
||||
modify_tasks(sample_dataset, task_replacements={"missing_task": "New task"})
|
||||
|
||||
|
||||
def test_modify_tasks_rejects_default_task_and_replacements(sample_dataset):
|
||||
"""Test that default-task assignment cannot be combined with find-and-replace."""
|
||||
with pytest.raises(ValueError, match="Cannot combine new_task with task_replacements"):
|
||||
modify_tasks(
|
||||
sample_dataset,
|
||||
new_task="Default task",
|
||||
task_replacements={"task_0": "Pick the cube"},
|
||||
)
|
||||
|
||||
|
||||
def test_modify_tasks_updates_info_json(sample_dataset):
|
||||
"""Test that total_tasks is updated in info.json."""
|
||||
episode_tasks = {0: "Task A", 1: "Task B", 2: "Task C", 3: "Task A", 4: "Task B"}
|
||||
|
||||
@@ -1,83 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset
|
||||
from lerobot.scripts.lerobot_edit_dataset import EditDatasetConfig, ModifyTasksConfig, handle_modify_tasks
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_dataset(tmp_path, empty_lerobot_dataset_factory):
|
||||
features = {
|
||||
"action": {"dtype": "float32", "shape": (6,), "names": None},
|
||||
"observation.state": {"dtype": "float32", "shape": (4,), "names": None},
|
||||
"observation.images.top": {"dtype": "image", "shape": (224, 224, 3), "names": None},
|
||||
}
|
||||
|
||||
dataset = empty_lerobot_dataset_factory(
|
||||
root=tmp_path / "test_dataset",
|
||||
features=features,
|
||||
)
|
||||
|
||||
for ep_idx in range(5):
|
||||
for _ in range(10):
|
||||
frame = {
|
||||
"action": np.random.randn(6).astype(np.float32),
|
||||
"observation.state": np.random.randn(4).astype(np.float32),
|
||||
"observation.images.top": np.random.randint(0, 255, size=(224, 224, 3), dtype=np.uint8),
|
||||
"task": f"task_{ep_idx % 2}",
|
||||
}
|
||||
dataset.add_frame(frame)
|
||||
dataset.save_episode()
|
||||
|
||||
dataset.finalize()
|
||||
return dataset
|
||||
|
||||
|
||||
def test_handle_modify_tasks_with_replacements(sample_dataset):
|
||||
cfg = EditDatasetConfig(
|
||||
repo_id=sample_dataset.repo_id,
|
||||
root=str(sample_dataset.root),
|
||||
operation=ModifyTasksConfig(
|
||||
task_replacements={
|
||||
"task_0": "Pick the cube",
|
||||
"task_1": "Place the cube",
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
handle_modify_tasks(cfg)
|
||||
|
||||
modified_dataset = LeRobotDataset(cfg.repo_id, root=sample_dataset.root)
|
||||
assert modified_dataset.meta.episodes[0]["tasks"][0] == "Pick the cube"
|
||||
assert modified_dataset.meta.episodes[1]["tasks"][0] == "Place the cube"
|
||||
assert len(modified_dataset.meta.tasks) == 2
|
||||
|
||||
|
||||
def test_handle_modify_tasks_rejects_default_task_and_replacements(sample_dataset):
|
||||
cfg = EditDatasetConfig(
|
||||
repo_id=sample_dataset.repo_id,
|
||||
root=str(sample_dataset.root),
|
||||
operation=ModifyTasksConfig(
|
||||
new_task="Default task",
|
||||
task_replacements={"task_0": "Pick the cube"},
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Cannot combine new_task with task_replacements"):
|
||||
handle_modify_tasks(cfg)
|
||||
Reference in New Issue
Block a user