File size: 2,566 Bytes
0c6c82c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
44745f2
0c6c82c
 
 
 
44745f2
 
 
 
0c6c82c
 
 
 
 
 
 
44745f2
 
0c6c82c
 
 
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
from __future__ import annotations

import math
import random

from .models import Request, SimulationConfig


def _sample_lognormal(mean: float, cv: float, rng: random.Random, minimum: int = 1) -> int:
    if cv <= 1e-9:
        return max(minimum, int(round(mean)))
    variance_ratio = cv * cv
    sigma2 = math.log(1.0 + variance_ratio)
    sigma = math.sqrt(sigma2)
    mu = math.log(max(mean, 1e-6)) - sigma2 / 2.0
    return max(minimum, int(round(rng.lognormvariate(mu, sigma))))


def _arrival_times(cfg: SimulationConfig, rng: random.Random) -> list[float]:
    rate = max(cfg.request_rate_rps, 1e-9)
    arrivals: list[float] = []
    t = 0.0

    if cfg.arrival_process == "constant":
        step = 1.0 / rate
        while t < cfg.duration_s:
            arrivals.append(t)
            t += step
        return arrivals

    if cfg.arrival_process == "bursty":
        while t < cfg.duration_s:
            phase = int(t // max(cfg.burst_period_s, 0.1)) % 2
            local_rate = rate * (cfg.burst_multiplier if phase else 0.55)
            t += rng.expovariate(max(local_rate, 1e-9))
            if t < cfg.duration_s:
                arrivals.append(t)
        return arrivals

    if cfg.arrival_process != "poisson":
        raise ValueError(f"Unknown arrival process: {cfg.arrival_process}")

    while t < cfg.duration_s:
        t += rng.expovariate(rate)
        if t < cfg.duration_s:
            arrivals.append(t)
    return arrivals


def generate_workload(cfg: SimulationConfig) -> list[Request]:
    rng = random.Random(cfg.seed)
    cache_rng = random.Random(cfg.seed ^ 0x5A17CACE)
    requests: list[Request] = []
    for idx, arrival in enumerate(_arrival_times(cfg, rng)):
        prompt = _sample_lognormal(cfg.prompt_tokens_mean, cfg.prompt_tokens_cv, rng)
        output = _sample_lognormal(cfg.output_tokens_mean, cfg.output_tokens_cv, rng)
        cached = 0
        if cfg.prefix_cache_enabled and cfg.shared_prefix_tokens > 0 and cfg.prefix_reuse_fraction > 0:
            if cache_rng.random() < min(max(cfg.prefix_reuse_fraction, 0.0), 1.0):
                cached = min(cfg.shared_prefix_tokens, max(prompt - 1, 0))
        requests.append(
            Request(
                request_id=idx,
                arrival_time=arrival,
                prompt_tokens=prompt,
                output_tokens=output,
                deadline_time=arrival + cfg.slo_e2e_ms / 1000.0,
                remaining_prefill=max(0, prompt - cached),
                cached_prefix_tokens=cached,
            )
        )
    return requests