File size: 16,962 Bytes
44745f2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
from __future__ import annotations

import heapq
from dataclasses import asdict, dataclass, field
from itertools import count

from .diagnostics import diagnose_run
from .kv_cache import KVCacheModel
from .latency import AnalyticalLatencyModel
from .metrics import percentile, summarize
from .models import Request, SimulationConfig, SimulationResult, TimelinePoint
from .profiles import get_accelerator, get_model
from .workloads import generate_workload


@dataclass
class PrefillWorker:
    worker_id: int
    latency: AnalyticalLatencyModel
    kv: KVCacheModel
    busy: bool = False
    busy_time_s: float = 0.0


@dataclass
class DecodeWorker:
    worker_id: int
    latency: AnalyticalLatencyModel
    kv: KVCacheModel
    active: list[Request] = field(default_factory=list)
    busy: bool = False
    busy_time_s: float = 0.0
    peak_kv_gb: float = 0.0


class DisaggregatedSimulator:
    """Two-stage prefill/decode discrete-event simulator.

    v0.3 models role-specific worker pools plus a serialized analytical KV-transfer
    link. It is intentionally a systems abstraction, not a distributed-runtime
    emulator: transport, compute and memory timings remain reference-model
    predictions and carry explicit provenance in every result.
    """

    def __init__(self, cfg: SimulationConfig):
        if cfg.prefill_workers < 1 or cfg.decode_workers < 1:
            raise ValueError("prefill_workers and decode_workers must be >= 1")
        if cfg.interconnect_gbps <= 0:
            raise ValueError("interconnect_gbps must be > 0")

        self.cfg = cfg
        self.model = get_model(cfg.model)
        self.prefill_accelerator = get_accelerator(cfg.prefill_accelerator)
        self.decode_accelerator = get_accelerator(cfg.decode_accelerator)
        self.prefill_latency = AnalyticalLatencyModel(self.model, self.prefill_accelerator, cfg.quantization)
        self.decode_latency = AnalyticalLatencyModel(self.model, self.decode_accelerator, cfg.quantization)

        self.prefill_workers = [
            PrefillWorker(i, self.prefill_latency, KVCacheModel(self.prefill_latency, cfg))
            for i in range(cfg.prefill_workers)
        ]
        self.decode_workers = [
            DecodeWorker(i, self.decode_latency, KVCacheModel(self.decode_latency, cfg))
            for i in range(cfg.decode_workers)
        ]

        self.requests = generate_workload(cfg)
        self.prefill_waiting: list[Request] = []
        self.decode_waiting: list[Request] = []
        self.completed: list[Request] = []
        self.events: list[tuple[float, int, str, object]] = []
        self.seq = count()
        self.now = 0.0
        self.transfer_busy_until = 0.0
        self.transfer_busy_time_s = 0.0
        self.transfer_gb = 0.0
        self.transfer_latencies_ms: list[float] = []
        self.transfer_pending = 0
        self.timeline: list[TimelinePoint] = []
        self.peak_total_kv_gb = 0.0
        self.warnings: list[str] = []

        for req in self.requests:
            self._push(req.arrival_time, "arrival", req)

    def _push(self, time_s: float, kind: str, payload: object) -> None:
        heapq.heappush(self.events, (max(time_s, self.now), next(self.seq), kind, payload))

    def _waiting_sorted(self, requests: list[Request]) -> list[Request]:
        if self.cfg.scheduler == "continuous_sjf":
            return sorted(requests, key=lambda r: (r.remaining_prefill + r.output_tokens, r.arrival_time))
        if self.cfg.scheduler in {"continuous_slo", "chunked_slo"}:
            def slack(req: Request) -> tuple[float, float]:
                prefill = self.prefill_latency.prefill_seconds([max(req.remaining_prefill, 1)])
                midpoint_context = req.prompt_tokens + max(req.output_tokens // 2, 1)
                decode = req.output_tokens * self.decode_latency.decode_step_seconds([midpoint_context])
                return (req.deadline_time - self.now - prefill - decode, req.arrival_time)

            return sorted(requests, key=slack)
        return sorted(requests, key=lambda r: r.arrival_time)

    def _record_timeline(self, force: bool = False) -> None:
        target = max(self.cfg.timeline_points, 20)
        total_kv = sum(w.kv.used_gb(w.active) for w in self.decode_workers)
        total_capacity = sum(w.kv.capacity_gb for w in self.decode_workers)
        self.peak_total_kv_gb = max(self.peak_total_kv_gb, total_kv)
        for worker in self.decode_workers:
            worker.peak_kv_gb = max(worker.peak_kv_gb, worker.kv.used_gb(worker.active))
        point = TimelinePoint(
            time_s=self.now,
            waiting=len(self.prefill_waiting),
            prefill_pending=sum(1 for w in self.prefill_workers if w.busy),
            decoding=sum(len(w.active) for w in self.decode_workers),
            completed=len(self.completed),
            kv_used_gb=total_kv,
            kv_capacity_gb=total_capacity,
            transfer_pending=self.transfer_pending,
            decode_ready=len(self.decode_waiting),
            prefill_active=sum(1 for w in self.prefill_workers if w.busy),
        )
        if not self.timeline or force or self.now - self.timeline[-1].time_s >= max(self.cfg.duration_s / target, 0.05):
            self.timeline.append(point)
            if len(self.timeline) > target * 2:
                self.timeline = self.timeline[::2]

    def _prefill_batch_for(self, worker: PrefillWorker) -> list[tuple[Request, int]]:
        if not self.prefill_waiting:
            return []
        selected: list[tuple[Request, int]] = []
        budget = self.cfg.max_batch_tokens
        virtual_selected: list[Request] = []
        for req in self._waiting_sorted(self.prefill_waiting):
            if len(selected) >= self.cfg.max_batch_size or budget <= 0:
                break
            chunk = req.remaining_prefill
            if self.cfg.scheduler == "chunked_slo":
                chunk = min(chunk, self.cfg.chunk_size)
            chunk = min(chunk, budget)
            if chunk <= 0:
                continue
            if not worker.kv.can_admit(req, [], virtual_selected):
                continue
            selected.append((req, chunk))
            virtual_selected.append(req)
            budget -= chunk
        return selected

    def _try_start_prefill(self) -> None:
        for worker in self.prefill_workers:
            if worker.busy:
                continue
            batch = self._prefill_batch_for(worker)
            if not batch:
                continue
            selected_ids = {r.request_id for r, _ in batch}
            self.prefill_waiting = [r for r in self.prefill_waiting if r.request_id not in selected_ids]
            for req, _ in batch:
                if req.first_prefill_time is None:
                    req.first_prefill_time = self.now
            delta = worker.latency.prefill_seconds([chunk for _, chunk in batch])
            worker.busy = True
            worker.busy_time_s += delta
            self._push(self.now + delta, "prefill_done", (worker.worker_id, batch))

    def _schedule_transfer(self, req: Request) -> None:
        # Cached shared-prefix KV is assumed resident in both role pools. Only
        # newly computed prompt state must cross the P/D boundary.
        bytes_to_transfer = req.uncached_prompt_tokens * self.prefill_latency.kv_bytes_per_token()
        gb = bytes_to_transfer / 1e9
        start = max(self.now, self.transfer_busy_until)
        duration = self.cfg.transfer_base_ms / 1000.0 + bytes_to_transfer / (self.cfg.interconnect_gbps * 1e9)
        end = start + duration
        self.transfer_busy_until = end
        self.transfer_busy_time_s += duration
        self.transfer_gb += gb
        self.transfer_latencies_ms.append(duration * 1000.0)
        self.transfer_pending += 1
        req.transfer_start_time = start
        req.transfer_end_time = end
        self._push(end, "transfer_done", req)

    def _decode_order(self) -> list[Request]:
        return self._waiting_sorted(self.decode_waiting)

    def _try_schedule_decode(self) -> None:
        # Admissions occur only between iterations. A worker marked busy has a
        # decode step already in flight and cannot accept work until it completes.
        for worker in self.decode_workers:
            if worker.busy:
                continue
            slots = self.cfg.max_batch_size - len(worker.active)
            if slots > 0 and self.decode_waiting:
                for req in list(self._decode_order()):
                    if slots <= 0:
                        break
                    if worker.kv.can_admit(req, worker.active):
                        self.decode_waiting.remove(req)
                        req.decode_worker_id = worker.worker_id
                        worker.active.append(req)
                        slots -= 1
            if not worker.active:
                continue
            contexts = [r.context_tokens for r in worker.active]
            delta = worker.latency.decode_step_seconds(contexts)
            worker.busy = True
            worker.busy_time_s += delta
            self._push(self.now + delta, "decode_done", worker.worker_id)

    def _handle_event(self, kind: str, payload: object) -> None:
        if kind == "arrival":
            self.prefill_waiting.append(payload)  # type: ignore[arg-type]
            self._try_start_prefill()
            return

        if kind == "prefill_done":
            worker_id, batch = payload  # type: ignore[misc]
            worker = self.prefill_workers[worker_id]
            worker.busy = False
            for req, chunk in batch:
                req.remaining_prefill = max(0, req.remaining_prefill - chunk)
                if req.remaining_prefill > 0:
                    self.prefill_waiting.append(req)
                else:
                    req.prefill_complete_time = self.now
                    self._schedule_transfer(req)
            self._try_start_prefill()
            return

        if kind == "transfer_done":
            req = payload  # type: ignore[assignment]
            self.transfer_pending = max(0, self.transfer_pending - 1)
            self.decode_waiting.append(req)
            self._try_schedule_decode()
            return

        if kind == "decode_done":
            worker = self.decode_workers[int(payload)]
            worker.busy = False
            for req in worker.active:
                req.generated_tokens += 1
                if req.first_token_time is None:
                    req.first_token_time = self.now
            done = [r for r in worker.active if r.complete]
            for req in done:
                req.completion_time = self.now
                self.completed.append(req)
            if done:
                done_ids = {r.request_id for r in done}
                worker.active = [r for r in worker.active if r.request_id not in done_ids]
            self._try_schedule_decode()
            return

        raise RuntimeError(f"Unknown event kind: {kind}")

    def run(self) -> SimulationResult:
        if self.cfg.scheduler == "static_fcfs":
            self.warnings.append("Static FCFS is not defined for P/D disaggregation; using continuous FCFS semantics.")
        self._record_timeline(force=True)
        processed = 0
        max_events = max(10000, len(self.requests) * max(self.cfg.output_tokens_mean, 1) * 20)

        while self.events and len(self.completed) < len(self.requests):
            time_s, _, kind, payload = heapq.heappop(self.events)
            self.now = max(self.now, time_s)
            self._handle_event(kind, payload)
            self._record_timeline()
            processed += 1
            if processed > max_events:
                self.warnings.append("Simulation stopped at the event safety limit.")
                break

        # If queued work remains with no events, the configuration is memory- or
        # admission-constrained rather than silently considered complete.
        if len(self.completed) < len(self.requests) and not self.events:
            self.warnings.append("Disaggregated pipeline stalled before all requests completed.")

        self._record_timeline(force=True)
        makespan = max(self.now, self.cfg.duration_s if self.requests else 0.0)
        prefill_util = sum(w.busy_time_s for w in self.prefill_workers) / max(makespan * len(self.prefill_workers), 1e-9)
        decode_util = sum(w.busy_time_s for w in self.decode_workers) / max(makespan * len(self.decode_workers), 1e-9)
        transfer_util = self.transfer_busy_time_s / max(makespan, 1e-9)
        hottest = min(1.0, max(prefill_util, decode_util, transfer_util))
        summary, latency = summarize(self.completed, self.cfg, makespan, hottest * makespan)
        summary["requests_generated"] = len(self.requests)
        summary["requests_unfinished"] = len(self.requests) - len(self.completed)

        total_kv_capacity = sum(w.kv.capacity_gb for w in self.decode_workers)
        peak_kv_gb = self.peak_total_kv_gb
        peak_worker_util = max(
            (w.peak_kv_gb / w.kv.capacity_gb if w.kv.capacity_gb > 0 else 0.0 for w in self.decode_workers),
            default=0.0,
        )
        resource = {
            "topology": "disaggregated_pd",
            "model_weight_gb": self.decode_latency.model_weight_gb,
            "kv_capacity_gb": total_kv_capacity,
            "peak_kv_gb": peak_kv_gb,
            "peak_kv_utilization": peak_worker_util,
            "prefill_accelerator_vram_gb": self.prefill_accelerator.vram_gb,
            "decode_accelerator_vram_gb": self.decode_accelerator.vram_gb,
            "prefill_workers": len(self.prefill_workers),
            "decode_workers": len(self.decode_workers),
            "accelerator_instances": len(self.prefill_workers) + len(self.decode_workers),
            "prefill_busy_fraction": min(1.0, prefill_util),
            "decode_busy_fraction": min(1.0, decode_util),
            "transfer_busy_fraction": min(1.0, transfer_util),
            "kv_transfer_gb": self.transfer_gb,
            "mean_transfer_ms": sum(self.transfer_latencies_ms) / len(self.transfer_latencies_ms) if self.transfer_latencies_ms else 0.0,
            "p95_transfer_ms": percentile(self.transfer_latencies_ms, 0.95),
            "interconnect_gbps": self.cfg.interconnect_gbps,
            "prefix_cache_gb_per_decode_worker": self.decode_workers[0].kv.shared_prefix_gb if self.decode_workers else 0.0,
            "prefix_cache_hits": sum(1 for r in self.requests if r.prefix_cache_hit),
            "prefix_cache_hit_rate": (sum(1 for r in self.requests if r.prefix_cache_hit) / len(self.requests)) if self.requests else 0.0,
            "prefill_tokens_saved": sum(r.cached_prefix_tokens for r in self.requests),
        }

        request_rows = []
        for req in self.completed[:2000]:
            transfer_ms = 0.0
            if req.transfer_start_time is not None and req.transfer_end_time is not None:
                transfer_ms = (req.transfer_end_time - req.transfer_start_time) * 1000.0
            request_rows.append(
                {
                    "request_id": req.request_id,
                    "arrival_time": req.arrival_time,
                    "prompt_tokens": req.prompt_tokens,
                    "output_tokens": req.output_tokens,
                    "cached_prefix_tokens": req.cached_prefix_tokens,
                    "decode_worker_id": req.decode_worker_id,
                    "transfer_ms": transfer_ms,
                    "ttft_ms": (req.first_token_time - req.arrival_time) * 1000.0 if req.first_token_time is not None else None,
                    "e2e_ms": (req.completion_time - req.arrival_time) * 1000.0 if req.completion_time is not None else None,
                }
            )

        diagnostics = diagnose_run(summary, latency, resource, self.cfg)
        provenance = {
            "simulator": "InferScale-Sim",
            "version": "0.3.0",
            "latency_profile_type": "analytical-reference",
            "profile_warning": "Reference profiles are analytical proxies, not measured hardware benchmarks.",
            "model_profile_source": self.model.source,
            "prefill_accelerator_profile_source": self.prefill_accelerator.source,
            "decode_accelerator_profile_source": self.decode_accelerator.source,
            "topology": "disaggregated_pd",
            "transfer_model": "serialized-reference-link",
        }
        return SimulationResult(
            config=self.cfg.to_dict(),
            provenance=provenance,
            summary=summary,
            latency=latency,
            resource=resource,
            diagnostics=diagnostics,
            requests=request_rows,
            timeline=[asdict(p) for p in self.timeline],
            warnings=self.warnings,
        )


def run_disaggregated(config: dict) -> dict:
    return DisaggregatedSimulator(SimulationConfig.from_dict(config)).run().to_dict()