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. the current model 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, cfg.prefill_time_scale, cfg.decode_time_scale ) self.decode_latency = AnalyticalLatencyModel( self.model, self.decode_accelerator, cfg.quantization, cfg.prefill_time_scale, cfg.decode_time_scale ) 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) ) * max(self.cfg.transfer_time_scale, 1e-6) 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) workload_horizon = 0.0 if self.cfg.arrival_process == "trace" else (self.cfg.duration_s if self.requests else 0.0) makespan = max(self.now, workload_horizon) 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", "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()