browsergym_form_task / tests /test_task_curriculum.py
morty649's picture
Deploy BrowserGym form task Space
a484d33
Raw
History Blame Contribute Delete
2.3 kB
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)