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)