Spaces:
Sleeping
Sleeping
| from concurrent.futures import ThreadPoolExecutor | |
| from server.task_curriculum import ( | |
| TaskCurriculumScheduler, | |
| select_reset_task, | |
| should_enable_task_curriculum, | |
| ) | |
| class ReverseShuffle: | |
| def __init__(self) -> None: | |
| self.calls = 0 | |
| def shuffle(self, values): | |
| self.calls += 1 | |
| values.reverse() | |
| def test_scheduler_visits_every_task_once_before_repetition(): | |
| tasks = ("a", "b", "c") | |
| scheduler = TaskCurriculumScheduler(tasks=tasks, rng=ReverseShuffle()) | |
| first_cycle = [scheduler.next_task() for _ in tasks] | |
| assert sorted(first_cycle) == sorted(tasks) | |
| next_task = scheduler.next_task() | |
| assert next_task in tasks | |
| assert scheduler.cycle == 2 | |
| def test_scheduler_reshuffles_after_complete_cycle(): | |
| tasks = ("a", "b", "c") | |
| rng = ReverseShuffle() | |
| scheduler = TaskCurriculumScheduler(tasks=tasks, rng=rng) | |
| assert rng.calls == 1 | |
| for _ in tasks: | |
| scheduler.next_task() | |
| scheduler.next_task() | |
| assert rng.calls == 2 | |
| def test_curriculum_enabled_non_miniwob_does_not_sample(): | |
| scheduler = TaskCurriculumScheduler(tasks=("scheduled",), rng=ReverseShuffle()) | |
| scheduler_before = scheduler.cycle | |
| enabled = should_enable_task_curriculum("webarena", enabled=True) | |
| selected_task = select_reset_task( | |
| current_task="0", | |
| scheduler=scheduler if enabled else None, | |
| ) | |
| assert enabled is False | |
| assert selected_task == "0" | |
| assert scheduler.cycle == scheduler_before | |
| def test_explicit_reset_task_overrides_scheduler_selection(): | |
| scheduler = TaskCurriculumScheduler(tasks=("scheduled",), rng=ReverseShuffle()) | |
| scheduler_before = scheduler.cycle | |
| selected_task = select_reset_task( | |
| current_task="click-test", | |
| scheduler=scheduler, | |
| explicit_task="enter-text", | |
| ) | |
| assert selected_task == "enter-text" | |
| assert scheduler.cycle == scheduler_before | |
| def test_scheduler_is_thread_safe_across_complete_cycles(): | |
| tasks = tuple(str(index) for index in range(10)) | |
| scheduler = TaskCurriculumScheduler(tasks=tasks, rng=ReverseShuffle()) | |
| with ThreadPoolExecutor(max_workers=20) as executor: | |
| results = list(executor.map(lambda _: scheduler.next_task(), range(100))) | |
| assert sorted(results) == sorted(tasks * 10) | |