Source code for stormlog.infer.vllm_hook.engine

"""Engine-core snapshots: admissions, scheduled steps, their outputs, and exits.

Every value is a scalar copied at the moment it is observed; nothing keeps a
live vLLM object. Per-step context comes from the scheduler output's own fields
because ``Scheduler.schedule`` has already advanced each request's
``num_computed_tokens`` by the time it returns, and under async scheduling the
next step has been scheduled before this one's output is processed.
"""

from __future__ import annotations

import time
from dataclasses import dataclass, field
from typing import Any

from .writer import EpochWriter

ITERATION_ATTRIBUTE = "_stormlog_iteration"
# vLLM keeps at most a few steps in flight; more pending means outputs were
# never processed (an error or shutdown), and the oldest are forgotten.
MAX_PENDING = 64


@dataclass
class _Member:
    internal: str
    scheduled: int
    computed_before: int
    prompt_tokens: int | None
    drafts: int


@dataclass
class _Pending:
    iteration: str
    members: dict[str, _Member]


[docs] @dataclass class EngineRecorder: """Records for one scheduler instance.""" writer: EpochWriter producer: str next_iteration: int = 0 # Per live request: its prompt length, and its committed context. Both are # dropped when vLLM frees the request, so they never outgrow its live set. prompt_tokens: dict[str, int] = field(default_factory=dict) committed: dict[str, int] = field(default_factory=dict) pending: dict[str, _Pending] = field(default_factory=dict) # ------------------------------------------------------------ admission
[docs] def on_admit(self, request: Any) -> None: self.writer.emit( "alias", { "internal": str(request.request_id), "external": _optional_str(getattr(request, "external_req_id", None)), **_stamp(), }, )
# ------------------------------------------------------------ schedule
[docs] def on_schedule(self, scheduler: Any, output: Any, start: tuple[int, int]) -> None: iteration = str(self.next_iteration) self.next_iteration += 1 setattr(output, ITERATION_ATTRIBUTE, (self.producer, iteration)) members = [ *self._new_members(scheduler, output), *self._cached_members(scheduler, output), ] self.pending[iteration] = _Pending( iteration, {member.internal: member for member, _ in members} ) while len(self.pending) > MAX_PENDING: self.pending.pop(next(iter(self.pending))) end = _stamp() self.writer.emit( "scheduled", { "iteration": iteration, "start_wall_ns": start[0], "start_mono_ns": start[1], "end_wall_ns": end["wall_ns"], "end_mono_ns": end["mono_ns"], "total_tokens": int(output.total_num_scheduled_tokens), "zero_token": int(output.total_num_scheduled_tokens) == 0, "preempted": sorted(getattr(output, "preempted_req_ids", None) or ()), "members": [fields for _, fields in members], }, )
def _new_members( self, scheduler: Any, output: Any ) -> list[tuple[_Member, dict[str, Any]]]: members = [] for data in output.scheduled_new_reqs: internal = str(data.req_id) first = internal not in self.prompt_tokens if first: self.prompt_tokens[internal] = int(data.prompt_len) members.append( self._member( scheduler, output, internal, computed_before=int(data.num_computed_tokens), first=first, context=True, ) ) return members def _cached_members( self, scheduler: Any, output: Any ) -> list[tuple[_Member, dict[str, Any]]]: cached = output.scheduled_cached_reqs is_context = getattr(cached, "is_context_phase", None) return [ self._member( scheduler, output, str(internal), computed_before=int(cached.num_computed_tokens[index]), first=str(internal) not in self.prompt_tokens, context=bool(is_context(internal)) if callable(is_context) else None, ) for index, internal in enumerate(cached.req_ids) ] def _member( self, scheduler: Any, output: Any, internal: str, *, computed_before: int, first: bool, context: bool | None, ) -> tuple[_Member, dict[str, Any]]: scheduled = int(output.num_scheduled_tokens.get(internal, 0)) drafts = len(output.scheduled_spec_decode_tokens.get(internal, ()) or ()) request = scheduler.requests.get(internal) if request is not None: # A resumable (streaming-input) request grows its prompt between # sessions, so the live length wins over the first sighting's. self.prompt_tokens[internal] = int(request.num_prompt_tokens) prompt = self.prompt_tokens.get(internal) prefill = max(0, min(scheduled, prompt - computed_before)) if prompt else 0 member = _Member(internal, scheduled, computed_before, prompt, drafts) fields = { "internal": internal, "sighting": "first" if first else "repeat", # vLLM's own classification: new in this output, or a cached request # still in its context phase. A recomputed request is context. "phase": ( None if context is None else ("context" if context else "generation") ), "scheduled": scheduled, "computed_before": computed_before, "prompt_tokens": prompt, "prefill_scheduled": prefill, "past_prompt_scheduled": scheduled - prefill, "drafts_scheduled": drafts, "cached_at_admission": computed_before if first else None, "recompute": (not first) and computed_before < self.committed.get(internal, computed_before), "output_before": ( int(request.num_output_tokens) if request is not None else None ), "resumable": bool(getattr(request, "resumable", False)), } return member, fields # ------------------------------------------------------------ output
[docs] def before_update( self, scheduler: Any, output: Any, model_output: Any ) -> dict[str, dict[str, Any]]: """Per member, what vLLM's own output loop is about to read.""" sampled = getattr(model_output, "sampled_token_ids", None) or [] index_of = getattr(model_output, "req_id_to_index", None) or {} snapshot = {} for internal in output.num_scheduled_tokens: request = scheduler.requests.get(internal) index = index_of.get(internal) snapshot[internal] = { "exists": request is not None, "finished": request is not None and bool(request.is_finished()), "stale": request is not None and int(getattr(request, "num_stale_output_tokens", 0)) > 0, "drop_stale": bool(getattr(request, "drop_stale_output", False)), # Copied now: stop handling trims the sampled list in place. "sampled": ( len(sampled[index]) if index is not None and index < len(sampled) else 0 ), } return snapshot
[docs] def after_update( self, scheduler: Any, output: Any, before: dict[str, dict[str, Any]], *, result: Any = None, failed: bool = False, ) -> None: """Record the step's outcome; a failed update is recorded as unknown.""" identity = getattr(output, ITERATION_ATTRIBUTE, None) if identity is None: return pending = self.pending.pop(identity[1], None) emitted = None if failed else _emitted(result) per_step = int(getattr(scheduler, "num_sampled_tokens_per_step", 1)) members = [ ( _failed_member(name) if failed else self._completed_member( scheduler, member, before.get(name, {}), emitted, per_step ) ) for name, member in (pending.members.items() if pending else ()) ] fields: dict[str, Any] = {"iteration": identity[1], **_stamp()} if failed: fields["update_failed"] = True fields["members"] = members self.writer.emit("completed", fields)
def _completed_member( self, scheduler: Any, member: _Member, before: dict[str, Any], emitted: dict[str, tuple[int, str | None]] | None, per_step: int, ) -> dict[str, Any]: outcome = _outcome(before) sampled = int(before.get("sampled", 0)) accepted = ( max(sampled - per_step, 0) if member.drafts and (sampled or per_step == 0) else 0 ) stale = bool(before.get("stale")) computed_after = None retained = None finish = None if outcome == "kept": rejected = 0 if stale else member.drafts - accepted computed_after = member.computed_before + member.scheduled - rejected if emitted is not None: retained, finish = emitted.get(member.internal, (0, None)) if member.internal in scheduler.requests: self.committed[member.internal] = computed_after return { "internal": member.internal, "outcome": outcome, "stale": stale, "sampled": sampled, "accepted_drafts": accepted, "retained": retained, "finish_reason": finish, "computed_after": computed_after, } # ------------------------------------------------------------ exit
[docs] def on_free(self, request: Any) -> None: internal = str(request.request_id) output_tokens = int(request.num_output_tokens) self.prompt_tokens.pop(internal, None) self.committed.pop(internal, None) status = getattr(request, "status", None) reason = ( request.get_finished_reason() if hasattr(request, "get_finished_reason") else None ) self.writer.emit( "terminal", { "internal": internal, "status": getattr(status, "name", None) or _optional_str(status), "finish_reason": _optional_str(reason), "output_tokens": output_tokens, **_stamp(), }, )
def _emitted(result: Any) -> dict[str, tuple[int, str | None]] | None: """Tokens and finish reason vLLM sent each request in this step's output. ``update_from_output`` returns ``{client: EngineCoreOutputs}``; counting the emitted tokens is right whatever vLLM's own counters do: stop trimming, a finished request already freed, or a streaming-input session reset. Any other shape is unknown, not zero. """ if not isinstance(result, dict): return None emitted: dict[str, tuple[int, str | None]] = {} for outputs in result.values(): for item in getattr(outputs, "outputs", None) or (): emitted[str(item.request_id)] = ( len(item.new_token_ids or ()), # vLLM's FinishReason prints as "stop", "length" and so on. _optional_str(getattr(item, "finish_reason", None)), ) return emitted def _failed_member(internal: str) -> dict[str, Any]: return { "internal": internal, "outcome": "unknown", "stale": None, "sampled": None, "accepted_drafts": None, "retained": None, "finish_reason": None, "computed_after": None, } def _outcome(before: dict[str, Any]) -> str: if not before: return "unknown" if not before["exists"] or before["finished"]: return "discarded_finished" if before["stale"] and before["drop_stale"]: return "dropped_stale" return "kept" def _stamp() -> dict[str, int]: return {"wall_ns": time.time_ns(), "mono_ns": time.monotonic_ns()} def _optional_str(value: Any) -> str | None: return None if value is None else str(value) __all__ = ["ITERATION_ATTRIBUTE", "EngineRecorder"]