InferScale-Sim / py /inferscale /disaggregated.py
ArchitSharma's picture
Commiting v0.3
44745f2
Raw
History Blame Contribute Delete
17 kB
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()