File size: 7,753 Bytes
8991f51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
"""Comprehensive Benchmark Suite for AsyncTensorRLHF.

Benchmarks:
1. In-VRAM Tensor-Native Reward vs. CPU-Roundtrip Baseline across multiple batch sizes.
2. PPO, M2PO, and GRPO loss & backward throughput (samples/sec & tokens/sec).
3. Replay buffer throughput (push/sample ops/sec under concurrency).
4. Staleness robustness benchmark (PPO vs. M2PO gradient variance under policy drift).
"""

import os
import pathlib
import sys
import time
import torch

# Ensure project root in sys.path
sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent.parent))

from src.reward.tensor_native import tensor_native_reward
from src.trainer.ppo_loss import compute_ppo_loss, compute_m2po_loss
from src.trainer.grpo_loss import compute_grpo_loss
from src.buffer.replay_buffer import BoundedReplayBuffer, VersionedReplayBuffer, Experience, VersionedExperience
from src.buffer.group_buffer import GroupAwareReplayBuffer


def run_reward_benchmark(device):
    print("\n" + "=" * 65)
    print("1. TENSOR-NATIVE REWARD vs. CPU-ROUNDTRIP BASELINE BENCHMARK")
    print("=" * 65)
    print(f"{'Batch Size':<12} | {'Seq Len':<10} | {'Tensor-Native (ms)':<20} | {'CPU-Decode (ms)':<18} | {'Speedup':<10}")
    print("-" * 65)

    batch_sizes = [16, 32, 64, 128, 256]
    seq_len = 256
    pat_len = 4

    results = []

    for B in batch_sizes:
        gen_ids = torch.randint(0, 5000, (B, seq_len), device=device)
        patterns = [torch.randint(0, 5000, (pat_len,), device=device) for _ in range(B)]

        # Warmup
        _ = tensor_native_reward(gen_ids, patterns, eos_token_id=99, device=device)
        if device == "cuda":
            torch.cuda.synchronize()

        # Benchmark Tensor-Native (In-VRAM)
        iters = 20
        t0 = time.perf_counter()
        for _ in range(iters):
            _ = tensor_native_reward(gen_ids, patterns, eos_token_id=99, device=device)
        if device == "cuda":
            torch.cuda.synchronize()
        t_native = (time.perf_counter() - t0) / iters * 1000

        # Benchmark CPU-Decode baseline (simulate copy to CPU and string/list search)
        t0 = time.perf_counter()
        for _ in range(iters):
            cpu_ids = gen_ids.cpu().tolist()
            cpu_pats = [p.cpu().tolist() for p in patterns]
            rewards = []
            for b in range(B):
                row = cpu_ids[b]
                pat = cpu_pats[b]
                match = any(row[i:i+len(pat)] == pat for i in range(len(row) - len(pat) + 1))
                rewards.append(1.0 if match else 0.0)
            _ = torch.tensor(rewards, device=device)
        if device == "cuda":
            torch.cuda.synchronize()
        t_cpu = (time.perf_counter() - t0) / iters * 1000

        speedup = t_cpu / max(t_native, 1e-4)
        print(f"{B:<12} | {seq_len:<10} | {t_native:<20.2f} | {t_cpu:<18.2f} | {speedup:<10.1f}x")
        results.append((B, seq_len, t_native, t_cpu, speedup))

    return results


def run_loss_benchmark(device):
    print("\n" + "=" * 65)
    print("2. POLICY LOSS THROUGHPUT (PPO vs. M2PO vs. GRPO)")
    print("=" * 65)
    print(f"{'Loss Type':<12} | {'Batch Size':<12} | {'Forward+Backward (ms)':<25} | {'Tokens/sec':<15}")
    print("-" * 65)

    B = 64
    L = 256
    iters = 30

    # PPO
    p_lp = torch.randn(B, L, device=device, requires_grad=True)
    o_lp = torch.randn(B, L, device=device)
    adv = torch.randn(B, L, device=device)

    # Warmup
    loss = compute_ppo_loss(p_lp, o_lp, adv)
    loss.backward()

    t0 = time.perf_counter()
    for _ in range(iters):
        p_lp.grad = None
        loss = compute_ppo_loss(p_lp, o_lp, adv)
        loss.backward()
    if device == "cuda":
        torch.cuda.synchronize()
    t_ppo = (time.perf_counter() - t0) / iters * 1000
    tok_ppo = (B * L) / (t_ppo / 1000)
    print(f"{'PPO':<12} | {B:<12} | {t_ppo:<25.2f} | {tok_ppo:<15.0f}")

    # M2PO
    p_lp.grad = None
    t0 = time.perf_counter()
    for _ in range(iters):
        p_lp.grad = None
        loss = compute_m2po_loss(p_lp, o_lp, adv, m2_threshold=2.0)
        loss.backward()
    if device == "cuda":
        torch.cuda.synchronize()
    t_m2po = (time.perf_counter() - t0) / iters * 1000
    tok_m2po = (B * L) / (t_m2po / 1000)
    print(f"{'M2PO':<12} | {B:<12} | {t_m2po:<25.2f} | {tok_m2po:<15.0f}")

    # GRPO
    G = 4
    B_grpo = B // G
    p_grp = torch.randn(B_grpo, G, L, device=device, requires_grad=True)
    o_grp = torch.randn(B_grpo, G, L, device=device)
    adv_grp = torch.randn(B_grpo, G, L, device=device)

    t0 = time.perf_counter()
    for _ in range(iters):
        p_grp.grad = None
        loss = compute_grpo_loss(p_grp, o_grp, adv_grp)
        loss.backward()
    if device == "cuda":
        torch.cuda.synchronize()
    t_grpo = (time.perf_counter() - t0) / iters * 1000
    tok_grpo = (B * L) / (t_grpo / 1000)
    print(f"{'GRPO':<12} | {B:<12} | {t_grpo:<25.2f} | {tok_grpo:<15.0f}")


def run_buffer_benchmark():
    print("\n" + "=" * 65)
    print("3. REPLAY BUFFER THROUGHPUT BENCHMARK")
    print("=" * 65)

    buf = BoundedReplayBuffer(max_size=50000)
    N = 20000
    exp = Experience(
        prompt_ids=torch.tensor([1, 2, 3]),
        generated_ids=torch.tensor([4, 5, 6, 7]),
        log_probs=torch.tensor([-0.1, -0.2, -0.1, -0.3]),
        reward=1.0,
        version=0,
    )

    t0 = time.perf_counter()
    for _ in range(N):
        buf.push(exp)
    t_push = time.perf_counter() - t0
    push_rate = N / t_push

    t0 = time.perf_counter()
    sampled = 0
    while sampled < N:
        batch = buf.sample(64)
        if not batch:
            break
        sampled += len(batch)
    t_sample = time.perf_counter() - t0
    sample_rate = sampled / t_sample

    print(f"Push Throughput:   {push_rate:,.0f} ops/sec ({N} items pushed in {t_push*1000:.1f} ms)")
    print(f"Sample Throughput: {sample_rate:,.0f} ops/sec ({sampled} items sampled in {t_sample*1000:.1f} ms)")


def run_staleness_benchmark(device):
    print("\n" + "=" * 65)
    print("4. ASYNCHRONOUS STALENESS ROBUSTNESS (PPO vs. M2PO)")
    print("=" * 65)
    print(f"{'Staleness tau':<15} | {'PPO Grad Norm':<18} | {'M2PO Grad Norm':<18} | {'Variance Reduction':<20}")
    print("-" * 65)

    B, L = 32, 64
    staleness_levels = [0, 1, 2, 3, 5, 8]

    for tau in staleness_levels:
        # Drift std proportional to staleness tau
        drift = 0.15 * tau
        old_lp = torch.randn(B, L, device=device)
        policy_lp_ppo = (old_lp + torch.randn(B, L, device=device) * drift).clone().detach().requires_grad_(True)
        policy_lp_m2po = policy_lp_ppo.clone().detach().requires_grad_(True)
        adv = torch.randn(B, L, device=device)

        loss_p = compute_ppo_loss(policy_lp_ppo, old_lp, adv)
        loss_p.backward()
        p_norm = policy_lp_ppo.grad.norm().item()

        loss_m = compute_m2po_loss(policy_lp_m2po, old_lp, adv, m2_threshold=2.0)
        loss_m.backward()
        m_norm = policy_lp_m2po.grad.norm().item()

        red = ((p_norm - m_norm) / max(p_norm, 1e-6)) * 100
        print(f"tau = {tau:<10} | {p_norm:<18.4f} | {m_norm:<18.4f} | {red:<20.1f}%")


def main():
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print("=" * 65)
    print("AsyncTensorRLHF Comprehensive System Benchmark")
    print(f"Hardware Platform: {device.upper()}")
    if device == "cuda":
        print(f"GPU: {torch.cuda.get_device_name(0)} (Capability: {torch.cuda.get_device_capability(0)})")
    print("=" * 65)

    run_reward_benchmark(device)
    run_loss_benchmark(device)
    run_buffer_benchmark()
    run_staleness_benchmark(device)
    print("\n" + "=" * 65)
    print("ALL COMPREHENSIVE BENCHMARKS COMPLETED SUCCESSFULLY")
    print("=" * 65)


if __name__ == "__main__":
    main()