Spaces:
Sleeping
Sleeping
File size: 2,302 Bytes
a484d33 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 | 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)
|