File size: 3,795 Bytes
2dd5f57
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
from unittest.mock import patch

import pytest

from speculators.train.trainer import _StepTimer


def test_disabled_timer_returns_none():
    timer = _StepTimer(enabled=False)
    timer.mark_value("start", 0.0)
    timer.mark("fetch")
    timer.mark("fwd")
    timer.mark("bwd")
    timer.mark("opt")
    assert timer.now() is None
    assert timer.profile(num_tokens=1024) is None


@patch("speculators.train.trainer.torch.accelerator.synchronize")
def test_enabled_timer_returns_profile(mock_sync):
    timer = _StepTimer(enabled=True)

    with patch(
        "speculators.train.trainer.time.perf_counter",
        side_effect=[
            0.1,
            0.3,
            0.5,
            0.6,
            0.6,
        ],
    ):
        timer.mark_value("start", 0.0)
        timer.mark("fetch")
        timer.mark("fwd")
        timer.mark("bwd")
        timer.mark("opt")
        t_next = timer.now()

    assert mock_sync.call_count == 5
    assert t_next == 0.6

    profile = timer.profile(num_tokens=4096)
    assert profile is not None
    assert profile["fetch_ms"] == (0.1 - 0.0) * 1000
    assert profile["fwd_ms"] == (0.3 - 0.1) * 1000
    assert profile["bwd_ms"] == (0.5 - 0.3) * 1000
    assert profile["opt_ms"] == (0.6 - 0.5) * 1000
    assert profile["step_ms"] == (0.6 - 0.0) * 1000
    assert profile["tokens_per_s"] == 4096 / 0.6
    assert profile["fetch_frac"] == 100 / 600


def test_disabled_to_enabled_transition():
    """Simulate log_freq=2: a disabled step followed by an enabled step.

    The disabled step's ``timer.now()`` returns None so the training loop
    falls back to ``time.perf_counter()`` for ``t_before_fetch``.  The
    subsequent enabled step feeds that value into ``mark_value("start", ...)``.
    Verify the profile is valid with a realistic ``start`` mark.
    """
    timer = _StepTimer()

    # --- disabled step (global_step=1, log_freq=2 → 1%2 != 0) ---
    timer.reset(enabled=False)
    timer.mark_value("start", 1.0)
    timer.mark("fetch")
    timer.mark("fwd")
    timer.mark("bwd")
    timer.mark("opt")
    assert timer.now() is None
    assert timer.profile(num_tokens=512) is None

    # Simulate the fallback: t_before_fetch = timer.now() or time.perf_counter()
    t_before_fetch = 2.0  # stands in for the perf_counter() fallback

    # --- enabled step (global_step=2, log_freq=2 → 2%2 == 0) ---
    timer.reset(enabled=True)
    timer.mark_value("start", t_before_fetch)

    with (
        patch("speculators.train.trainer.torch.accelerator.synchronize"),
        patch(
            "speculators.train.trainer.time.perf_counter",
            side_effect=[2.1, 2.4, 2.5, 2.6, 2.6],
        ),
    ):
        timer.mark("fetch")
        timer.mark("fwd")
        timer.mark("bwd")
        timer.mark("opt")
        t_next = timer.now()

    assert t_next == 2.6
    profile = timer.profile(num_tokens=2048)
    assert profile is not None
    assert profile["fetch_ms"] == (2.1 - 2.0) * 1000
    assert profile["fwd_ms"] == (2.4 - 2.1) * 1000
    assert profile["bwd_ms"] == (2.5 - 2.4) * 1000
    assert profile["opt_ms"] == (2.6 - 2.5) * 1000
    assert profile["step_ms"] == (2.6 - 2.0) * 1000
    assert profile["tokens_per_s"] == pytest.approx(2048 / 0.6)


def test_zero_step_ms_returns_zero_throughput():
    timer = _StepTimer(enabled=True)
    timer.mark_value("start", 1.0)
    with (
        patch("speculators.train.trainer.torch.accelerator.synchronize"),
        patch("speculators.train.trainer.time.perf_counter", return_value=1.0),
    ):
        timer.mark("fetch")
        timer.mark("fwd")
        timer.mark("bwd")
        timer.mark("opt")
    profile = timer.profile(num_tokens=4096)
    assert profile is not None
    assert profile["tokens_per_s"] == 0.0
    assert profile["fetch_frac"] == 0.0