Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 77 additions & 10 deletions server/executor/engine.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import logging
import random
import uuid
from collections import deque
from dataclasses import dataclass
Expand Down Expand Up @@ -26,9 +27,10 @@
SequenceBatchTask,
)
from server.metrics.timers import now_ns
from server.model.determinism import uniforms_from_seeds
from server.model.inference_context import InferenceContext, inference_context
from server.model.prefill_helpers import build_prefill_inputs
from server.model.sampling import sample_token
from server.model.sampling import sample_token, sample_tokens
from server.model.types import ModelBackend

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -364,10 +366,29 @@ def __init__(

def _admit(self, req: GenerationRequestState, seq: Sequence) -> None:
"""Track a request and hand its sequence to the scheduler's waiting list."""
# Unseeded requests still need a concrete integer seed for the batched
# counter-based noise source. Draw a stable per-request salt once here
# (engine thread, each request passes through exactly once): reproducible
# across the request's decode steps, random across requests.
if req.sampling_params.seed is None and req.noise_salt is None:
req.noise_salt = random.getrandbits(63)
self._all_requests[req.request_id] = RequestContext(request=req, sequence=seq)
self._seq_to_request[seq.sequence_id] = req
self._scheduler.add(seq)

def _noise_seed(self, request_state: GenerationRequestState) -> int:
"""Effective sampling seed: the request's seed, else its per-request salt.

The salt is normally assigned at admission; assign it lazily here too so
any request reaching the sampler always has a concrete integer seed.
"""
seed = request_state.sampling_params.seed
if seed is not None:
return seed
if request_state.noise_salt is None:
request_state.noise_salt = random.getrandbits(63)
return request_state.noise_salt
Comment on lines +385 to +390

def _try_admit(
self, candidates: deque[tuple[GenerationRequestState, Sequence]]
) -> deque[tuple[GenerationRequestState, Sequence]]:
Expand Down Expand Up @@ -482,9 +503,7 @@ def _reap_cancelled(self) -> None:
"""
if self._pending:
self._pending = deque(
(req, seq)
for (req, seq) in self._pending
if not req.cancelled.is_set()
(req, seq) for (req, seq) in self._pending if not req.cancelled.is_set()
)

# Snapshot: _cleanup_request mutates _all_requests during iteration.
Expand Down Expand Up @@ -635,6 +654,16 @@ def _sample_one(
token_id = sample_token(
logits, request_state.sampling_params, request_state.generator
)
return self._decode_result_for_token(token_id, request_state)
Comment on lines 654 to +657

def _decode_result_for_token(
self, token_id: int, request_state: GenerationRequestState
) -> DecodeResult:
"""Turn a sampled ``token_id`` into a ``DecodeResult`` (EOS/detok/max-len).

Shared by the per-row prefill path (``_sample_one``) and the batched
decode path (``_post_decode``), so both apply identical finish logic.
"""
if token_id == self._backend.tokenizer.eos_token_id:
return DecodeResult(
token_id=token_id, token="", finish_reason=FinishReason.EOS
Expand Down Expand Up @@ -727,15 +756,53 @@ def _post_decode(
(the value passed as ``position_id`` in ``_prepare_decode``), so we
advance ``num_tokens`` by one to reflect the new cached length before
the next decode step. ``out.logits`` has one row per sequence.

The whole batch is sampled in a single ``sample_tokens`` call with one
``.tolist()`` sync, replacing the old per-row ``.item()`` loop (cost C1
in the engine-loop plan). Randomness is deterministic and
batch-invariant: each row draws from ``(its seed, its output-token
count)`` via ``uniforms_from_seeds``, independent of batch composition.
"""
for i, seq in enumerate(sequences):
request_state = self._seq_to_request[seq.sequence_id]
batch_size = len(sequences)
if batch_size == 0:
return

request_states = [self._seq_to_request[seq.sequence_id] for seq in sequences]
device = out.logits.device
logits = out.logits[0, :batch_size, :] # [B, V]

temperature = torch.tensor(
[rs.sampling_params.temperature for rs in request_states],
dtype=torch.float32,
device=device,
)
top_k = torch.tensor(
[rs.sampling_params.top_k for rs in request_states],
dtype=torch.long,
device=device,
)
top_p = torch.tensor(
[rs.sampling_params.top_p for rs in request_states],
dtype=torch.float32,
device=device,
)
# step = the request's output-token count so far; output_tokens is
# appended inside on_token *after* this call, so it is the correct
# 0-based step index for the token being drawn now.
seeds = [self._noise_seed(rs) for rs in request_states]
steps = [rs.num_output_tokens for rs in request_states]
uniforms = uniforms_from_seeds(seeds, steps, device=device)

# One batched sample, one sync. A failure inside this math would have
# poisoned every row anyway, so it is left to propagate to run()'s fatal
# handler; per-row params are already validated at the API layer.
token_ids = sample_tokens(logits, temperature, top_k, top_p, uniforms).tolist()

for seq, request_state, token_id in zip(sequences, request_states, token_ids):
# The input token was committed to the cache during the forward.
seq.num_tokens += 1

logit = out.logits[0, i, :].unsqueeze(0)
try:
result = self._sample_one(logit, request_state)
result = self._decode_result_for_token(token_id, request_state)
seq.generated_token_ids.append(result.token_id)
self._emitter.on_token(
request_state,
Expand All @@ -752,7 +819,7 @@ def _post_decode(
except Exception:
# Fail just this request instead of tearing down the worker.
logger.exception(
"Failed to sample next token for request %s",
"Failed to finalize token for request %s",
request_state.request_id,
)
seq.finished = True
Expand Down
5 changes: 5 additions & 0 deletions server/executor/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,6 +193,11 @@ class GenerationRequestState:

generator: torch.Generator | None = None

# Stable per-request salt assigned at admission for unseeded requests, so the
# batched counter-based noise source has a concrete integer seed. Reproducible
# across the request's decode steps, random across requests.
noise_salt: int | None = None

# Set by the HTTP thread (Worker.cancel) on timeout/disconnect; polled by the
# engine thread. An Event is thread-safe, so no lock is needed.
cancelled: threading.Event = field(default_factory=threading.Event)
Expand Down
69 changes: 69 additions & 0 deletions tests/executor/test_schedule_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -692,6 +692,75 @@ def decode(self, token_ids, skip_special_tokens=True) -> str:
assert any(isinstance(e, ErrorEvent) for e in drain_events(req_a))


def test_post_decode_mixed_greedy_sampled_topk_batch() -> None:
"""A batch mixing a greedy row, a sampled top-k=1 row, and a sampled
top-k=2 row all decode correctly in one batched sample_tokens call."""
engine, *_ = _make_engine(prompt_tokens=[1, 2, 3])

req_greedy = GenerationRequestState(
request_id="greedy",
sampling_params=SamplingParams(max_new_tokens=5, temperature=0.0, top_p=1.0),
prompt="hello",
enqueued_ns=0,
)
# top_k=1 collapses the nucleus to the single argmax, so its sampled token
# is deterministic even though temperature > 0.
req_k1 = GenerationRequestState(
request_id="k1",
sampling_params=SamplingParams(
max_new_tokens=5, temperature=1.0, top_p=1.0, top_k=1, seed=123
),
prompt="hello",
enqueued_ns=0,
)
# top_k=2 leaves two survivors; the picked token must be one of them.
req_k2 = GenerationRequestState(
request_id="k2",
sampling_params=SamplingParams(
max_new_tokens=5, temperature=1.0, top_p=1.0, top_k=2, seed=7
),
prompt="hello",
enqueued_ns=0,
)

reqs = [req_greedy, req_k1, req_k2]
seqs = []
for i, req in enumerate(reqs):
seq = Sequence(
sequence_id=f"s{i}",
prompt_token_ids=[1, 2, 3],
generated_token_ids=[4],
num_prompt_tokens=3,
num_tokens=3,
max_new_tokens=5,
block_table=[i],
state=SequenceState.RUNNING,
)
engine._seq_to_request[seq.sequence_id] = req
seqs.append(seq)

# Row 0 argmax -> 5; row 1 argmax -> 9; row 2 top-2 -> {12, 13}.
logits = torch.full((1, 3, VOCAB), -10.0, dtype=torch.float32)
logits[0, 0, 5] = 1.0
logits[0, 1, 9] = 1.0
logits[0, 2, 12] = 2.0
logits[0, 2, 13] = 1.5
out = _FakeOutput(logits)

engine._post_decode(out, seqs)

assert seqs[0].generated_token_ids == [4, 5] # greedy -> exact argmax
assert seqs[1].generated_token_ids == [4, 9] # top_k=1 -> forced argmax
assert seqs[2].generated_token_ids[-1] in {12, 13} # top_k=2 -> allowed set

for req in reqs:
assert req.status != RequestStatus.FAILED
assert req.num_output_tokens == 1
for seq in seqs:
assert seq.num_tokens == 4 # advanced by one
assert seq.finished is False


def _make_resumable_request(
request_id: str, max_new_tokens: int, start_ns: int, num_prompt_tokens: int
) -> GenerationRequestState:
Expand Down