Skip to content
Draft
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
2 changes: 2 additions & 0 deletions docs/serving/online_serving/token_in_token_out.md
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@ The request field `output_mode` selects how much of the postprocessing the serve

Every response and every stream chunk echoes `output_mode`. Servers that predate the field ignore it and return token IDs only with a 200, so check that the `output_mode` in the response matches the one you sent. Older servers don't return it.

`weight_version` reports the committed serving version on each output. A request resumed after `pause(mode="keep")` and a weight update can therefore emit chunks with different versions. Read the version from each chunk rather than assuming it is fixed for the request.

```python
import httpx

Expand Down
3 changes: 3 additions & 0 deletions rust/src/engine-core-client/src/protocol/output.rs
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,8 @@ pub struct EngineCoreOutput {
/// the first output of a request that asked for them.
#[serde(default)]
pub prompt_token_id_logprobs: Option<WireNdArray>,
#[serde(default)]
pub weight_version: Option<String>,
}

/// Raw per-sequence speculative-decoding accumulator.
Expand Down Expand Up @@ -542,6 +544,7 @@ mod tests {
new_sampling_mask: None,
spec_decode_metrics: None,
prompt_token_id_logprobs: None,
weight_version: None,
},
],
scheduler_stats: None,
Expand Down
3 changes: 3 additions & 0 deletions rust/src/engine-core-client/src/tests/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2858,6 +2858,9 @@ fn python_msgpack_fixtures_match_rust_encoding() {
new_sampling_mask: None,
spec_decode_metrics: None,
prompt_token_id_logprobs: None,
weight_version: Some(
"step-7",
),
},
],
scheduler_stats: None,
Expand Down
2 changes: 2 additions & 0 deletions rust/src/engine-core-client/src/tests/python_compat.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,7 @@ class EngineCoreOutput(
new_sampling_mask: object | None = None
spec_decode_metrics: object | None = None
prompt_token_id_logprobs: object | None = None
weight_version: str | None = None


class ExtendedEngineCoreOutput(EngineCoreOutput):
Expand Down Expand Up @@ -237,6 +238,7 @@ class EngineCoreOutputs(
request_id="req-1",
new_token_ids=[7, 8],
finish_reason=FinishReason.LENGTH,
weight_version="step-7",
)
],
finished_requests={"req-1"},
Expand Down
6 changes: 6 additions & 0 deletions rust/src/llm/src/output.rs
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ pub struct CollectedGenerateOutput {
pub sampling_mask: Option<SamplingMask>,
/// Per-request speculative-decoding metrics from the terminal output.
pub spec_decode_metrics: Option<RequestSpecDecodeMetrics>,
pub weight_version: Option<String>,
}

/// Prompt-scoped metadata emitted only once on the first [`GenerateOutput`] for
Expand Down Expand Up @@ -170,6 +171,7 @@ pub struct GenerateOutput {
pub sampling_mask: Option<SamplingMask>,
/// Per-request speculative-decoding metrics, present on terminal outputs.
pub spec_decode_metrics: Option<RequestSpecDecodeMetrics>,
pub weight_version: Option<String>,
}

impl GenerateOutput {
Expand Down Expand Up @@ -220,6 +222,7 @@ impl GenerateOutput {
ec_transfer_params: None,
sampling_mask: None,
spec_decode_metrics: None,
weight_version: None,
}
}
}
Expand Down Expand Up @@ -327,6 +330,7 @@ impl Stream for GenerateOutputStream {
ec_transfer_params: raw.ec_transfer_params,
sampling_mask,
spec_decode_metrics: raw.spec_decode_metrics,
weight_version: raw.weight_version,
};

Poll::Ready(Some(Ok(output)))
Expand Down Expand Up @@ -421,6 +425,7 @@ impl<T: Stream<Item = Result<GenerateOutput>> + Send> T {
ec_transfer_params: None,
sampling_mask,
spec_decode_metrics: None,
weight_version: output.weight_version.clone(),
});
}

Expand All @@ -435,6 +440,7 @@ impl<T: Stream<Item = Result<GenerateOutput>> + Send> T {
collected.kv_transfer_params = output.kv_transfer_params;
collected.ec_transfer_params = output.ec_transfer_params;
collected.spec_decode_metrics = output.spec_decode_metrics;
collected.weight_version = output.weight_version;
if let Some(mask) = collected.sampling_mask.as_ref()
&& mask.rows.len() != collected.token_ids.len()
{
Expand Down
1 change: 1 addition & 0 deletions rust/src/llm/tests/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -424,6 +424,7 @@ async fn collect_output_rejects_partial_sampling_mask() {
rows: vec![vec![1, 33, 99]],
}),
spec_decode_metrics: None,
weight_version: None,
};

let error = futures::stream::iter([Ok(output)]).collect_output().await.unwrap_err();
Expand Down
28 changes: 28 additions & 0 deletions rust/src/server/src/routes/inference/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -170,10 +170,12 @@ async fn generate_chunk_stream(
let mut prompt_token_ids = None;
let mut usage = TokenUsage::default();
let mut spec_decode_metrics = None;
let mut weight_version = None;

while let Some(next) = stream.next().await {
match next {
Ok(output) => {
weight_version = output.weight_version.clone();
if let Some(metrics) = output.spec_decode_metrics {
spec_decode_metrics = Some(SpeculativeDecodingMetrics::from(metrics));
}
Expand Down Expand Up @@ -238,6 +240,7 @@ async fn generate_chunk_stream(
mm_placeholders: prompt_token_ids.as_ref().and_then(|_| mm_placeholders.take()),
prompt_token_ids,
metrics: None,
weight_version: output.weight_version,
})
.await;
}
Expand All @@ -263,6 +266,7 @@ async fn generate_chunk_stream(
speculative_decoding,
),
}),
weight_version,
})
.await;
}
Expand Down Expand Up @@ -347,6 +351,7 @@ fn collect_generate(
metrics: collected.spec_decode_metrics.map(|metrics| PerRequestMetrics {
speculative_decoding: SpeculativeDecodingMetrics::from(metrics),
}),
weight_version: collected.weight_version,
})
}

Expand Down Expand Up @@ -537,6 +542,7 @@ mod tests {
ec_transfer_params: None,
sampling_mask: None,
spec_decode_metrics: None,
weight_version: None,
}),
Ok(GenerateOutput {
request_id: String::new(),
Expand All @@ -553,6 +559,7 @@ mod tests {
ec_transfer_params: None,
sampling_mask: None,
spec_decode_metrics: None,
weight_version: Some("step-2".to_string()),
}),
]);

Expand Down Expand Up @@ -583,6 +590,8 @@ mod tests {
.expect("collect chunks");

assert_eq!(chunks.len(), 2);
assert_eq!(chunks[0].weight_version.as_deref(), Some("step-2"));
assert_eq!(chunks[1].weight_version.as_deref(), Some("step-2"));
assert_eq!(chunks[0].prompt_token_ids, Some(vec![11, 22]));
assert_eq!(chunks[0].mm_placeholders, Some(mm_placeholders));
assert!(chunks[1].prompt_token_ids.is_none());
Expand Down Expand Up @@ -637,6 +646,7 @@ mod tests {
ec_transfer_params: None,
sampling_mask: None,
spec_decode_metrics: None,
weight_version: None,
}
}

Expand Down Expand Up @@ -672,6 +682,19 @@ mod tests {
assert_eq!(chunks[0].choices[0].sampling_mask, Some(vec![vec![33, 44]]));
}

#[tokio::test]
async fn generate_chunk_stream_tracks_weight_updates() {
let mut first = stream_output(Some(&[11]), vec![33], None);
first.weight_version = Some("step-1".to_string());
let mut second = stream_output(None, vec![44], Some(FinishReason::Length));
second.weight_version = Some("step-2".to_string());

let chunks = collect_chunks(vec![first, second], false, None).await;

assert_eq!(chunks[0].weight_version.as_deref(), Some("step-1"));
assert_eq!(chunks[1].weight_version.as_deref(), Some("step-2"));
}

#[tokio::test]
async fn generate_chunk_stream_omits_prompt_metadata_by_default() {
let chunks = collect_chunks(
Expand Down Expand Up @@ -811,6 +834,7 @@ mod tests {
rows: vec![vec![30, 40]],
}),
spec_decode_metrics: None,
weight_version: Some("step-2".to_string()),
};

let response = collect_generate(
Expand All @@ -828,6 +852,7 @@ mod tests {
)
.expect("response");

assert_eq!(response.weight_version.as_deref(), Some("step-2"));
assert!(response.prompt_token_ids.is_none());
assert!(response.mm_placeholders.is_none());
assert_eq!(response.choices[0].sampling_mask, Some(vec![vec![30, 40]]));
Expand Down Expand Up @@ -865,6 +890,7 @@ mod tests {
prompt_token_ids: vec![10],
sampling_mask: None,
spec_decode_metrics: Some(metrics.clone()),
weight_version: None,
};

let response = collect_generate(
Expand Down Expand Up @@ -957,6 +983,7 @@ mod tests {
prompt_token_ids: vec![10, 20],
sampling_mask: None,
spec_decode_metrics: None,
weight_version: None,
};

let response = collect_generate(
Expand Down Expand Up @@ -992,6 +1019,7 @@ mod tests {
ec_transfer_params: None,
sampling_mask: None,
spec_decode_metrics: None,
weight_version: None,
prompt_token_ids,
};

Expand Down
2 changes: 2 additions & 0 deletions rust/src/server/src/routes/inference/generate/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,7 @@ pub(super) struct GenerateStreamResponse {
pub prompt_token_ids: Option<Vec<u32>>,
pub mm_placeholders: Option<MultiModalPlaceholders>,
pub metrics: Option<PerRequestMetrics<StreamingSpeculativeDecodingMetrics>>,
pub weight_version: Option<String>,
}

/// Mirrors the Python vLLM `GenerateResponse` class.
Expand All @@ -103,6 +104,7 @@ pub(super) struct GenerateResponse {
pub kv_transfer_params: Option<Value>,
pub ec_transfer_params: Option<Value>,
pub metrics: Option<PerRequestMetrics<SpeculativeDecodingMetrics>>,
pub weight_version: Option<String>,
}

#[derive(Debug, Clone, PartialEq, Serialize)]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,7 @@ def _make_request_output(
index: int = 0,
spec_decode_metrics: RequestSpecDecodeMetrics | None = None,
text: str = "",
weight_version: str | None = None,
) -> RequestOutput:
return RequestOutput(
request_id=request_id,
Expand All @@ -169,6 +170,7 @@ def _make_request_output(
encoder_prompt=None,
encoder_prompt_token_ids=None,
num_cached_tokens=num_cached_tokens,
weight_version=weight_version,
)


Expand Down Expand Up @@ -226,7 +228,11 @@ async def test_serve_tokens_skips_mm_cache_for_remote_engine_execution():

async def mock_generate(*args, **kwargs):
yield _make_request_output(
"req-1", token_ids=[10], finish_reason="stop", finished=True
"req-1",
token_ids=[10],
finish_reason="stop",
finished=True,
weight_version="step-7",
)

engine.generate = MagicMock(side_effect=mock_generate)
Expand All @@ -242,6 +248,7 @@ async def mock_generate(*args, **kwargs):
response = await serving.serve_tokens(request)

assert isinstance(response, GenerateTokensResponse)
assert response.weight_version == "step-7"
assert (
serving.online_renderer.preprocess_completion.call_args.kwargs["skip_mm_cache"]
is True
Expand Down Expand Up @@ -344,10 +351,14 @@ async def test_stream_basic():
engine = _mock_engine()

async def mock_generate(*args, **kwargs):
yield _make_request_output("req-1", token_ids=[10])
yield _make_request_output("req-1", token_ids=[20, 30])
yield _make_request_output("req-1", token_ids=[10], weight_version="step-1")
yield _make_request_output("req-1", token_ids=[20, 30], weight_version="step-2")
yield _make_request_output(
"req-1", token_ids=[40], finish_reason="stop", finished=True
"req-1",
token_ids=[40],
finish_reason="stop",
finished=True,
weight_version="step-2",
)

engine.generate = MagicMock(side_effect=mock_generate)
Expand Down Expand Up @@ -376,6 +387,11 @@ async def mock_generate(*args, **kwargs):
assert data_chunks[1]["choices"][0]["token_ids"] == [20, 30]
assert data_chunks[2]["choices"][0]["token_ids"] == [40]
assert data_chunks[2]["choices"][0]["finish_reason"] == "stop"
assert [chunk["weight_version"] for chunk in data_chunks] == [
"step-1",
"step-2",
"step-2",
]


@pytest.mark.asyncio
Expand Down Expand Up @@ -957,7 +973,11 @@ async def test_stream_include_usage():
async def mock_generate(*args, **kwargs):
yield _make_request_output("req-1", token_ids=[10])
yield _make_request_output(
"req-1", token_ids=[20], finish_reason="stop", finished=True
"req-1",
token_ids=[20],
finish_reason="stop",
finished=True,
weight_version="step-2",
)

engine.generate = MagicMock(side_effect=mock_generate)
Expand All @@ -980,11 +1000,13 @@ async def mock_generate(*args, **kwargs):
assert parsed[-1] == "[DONE]"

# The chunk before [DONE] should be the usage-only chunk
assert "weight_version" not in parsed[0]
usage_chunk = parsed[-2]
assert usage_chunk["choices"] == []
assert usage_chunk["usage"]["prompt_tokens"] == 3
assert usage_chunk["usage"]["completion_tokens"] == 2
assert usage_chunk["usage"]["total_tokens"] == 5
assert usage_chunk["weight_version"] == "step-2"


@pytest.mark.asyncio
Expand Down
5 changes: 4 additions & 1 deletion tests/v1/core/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -898,7 +898,9 @@ def test_update_from_output_routes_sampling_masks_by_request():
),
)

outputs = scheduler.update_from_output(scheduler_output, model_output)[0].outputs
outputs = scheduler.update_from_output(
scheduler_output, model_output, weight_version="step-2"
)[0].outputs

assert [out.request_id for out in outputs] == [req.request_id for req in requests]
assert [out.new_sampling_mask.token_ids.tolist() for out in outputs] == [
Expand All @@ -907,6 +909,7 @@ def test_update_from_output_routes_sampling_masks_by_request():
[4, 5, 6],
]
assert all(out.new_sampling_mask.offsets is None for out in outputs)
assert all(out.weight_version == "step-2" for out in outputs)


def test_update_from_output_routes_multi_position_sampling_masks():
Expand Down
11 changes: 11 additions & 0 deletions tests/v1/engine/test_engine_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -701,6 +701,17 @@ def _pausable_engine_core_proc() -> EngineCoreProc:
return core


def test_abort_output_carries_committed_weight_version():
core = _pausable_engine_core_proc()
core._weight_version = "2"
core.output_queue = MagicMock()

core._send_abort_outputs_to_client(["request-0"], 0)

_, outputs = core.output_queue.put_nowait.call_args.args[0]
assert outputs.outputs[0].weight_version == "2"


@pytest.mark.parametrize(
"pause_state,has_requests,has_batches",
[
Expand Down
Loading
Loading