Spaces:
Running
Running
| 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 | |
| class PrefillWorker: | |
| worker_id: int | |
| latency: AnalyticalLatencyModel | |
| kv: KVCacheModel | |
| busy: bool = False | |
| busy_time_s: float = 0.0 | |
| 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() | |