Skip to content
Closed
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
79 changes: 78 additions & 1 deletion tests/entrypoints/scale_out/derender/test_derender.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ def _make_generate_response(
token_ids: list[int] | None,
request_id: str = "chatcmpl-test-id",
finish_reason: str = "stop",
output_text: str | None = None,
logprobs: dict | None = None,
prompt_logprobs: list | None = None,
kv_transfer_params: dict | None = None,
Expand All @@ -60,6 +61,8 @@ def _make_generate_response(
"finish_reason": finish_reason,
"logprobs": logprobs,
}
if output_text is not None:
choice["output_text"] = output_text
return {
"request_id": request_id,
"choices": [choice],
Expand Down Expand Up @@ -108,6 +111,40 @@ async def test_derender_chat_roundtrip(client):
assert data["choices"][0]["message"]["role"] == "assistant"


@pytest.mark.asyncio
@pytest.mark.parametrize(
("output_text", "include_stop", "min_tokens"),
[
("answer", False, 0),
("answer<END>", True, 0),
("keep<END>allowed", False, 2),
],
)
async def test_derender_chat_uses_engine_stop_text(
client, output_text, include_stop, min_tokens
):
gen_req = await _render_chat(client)
token_ids = gen_req["token_ids"][:5]
response = await client.post(
"/v1/chat/completions/derender",
json={
"model": MODEL_NAME,
"generate_response": _make_generate_response(
token_ids, output_text=output_text
),
"chat_request": {
"model": MODEL_NAME,
"messages": [{"role": "user", "content": "Hello"}],
"stop": ["<END>"],
"include_stop_str_in_output": include_stop,
"min_tokens": min_tokens,
},
},
)
assert response.status_code == 200
assert response.json()["choices"][0]["message"]["content"] == output_text


@pytest.mark.asyncio
async def test_derender_chat_usage(client):
"""Supplied prompt_tokens flows through into usage correctly."""
Expand Down Expand Up @@ -392,10 +429,11 @@ async def _render_completion(client: httpx.AsyncClient, prompt: str) -> dict:
def _make_completion_generate_response(
token_ids: list[int],
request_id: str,
output_text: str | None = None,
kv_transfer_params: dict | None = None,
logprobs: dict | None = None,
) -> dict:
return {
response = {
"request_id": request_id,
"choices": [
{
Expand All @@ -408,6 +446,9 @@ def _make_completion_generate_response(
"prompt_logprobs": None,
"kv_transfer_params": kv_transfer_params,
}
if output_text is not None:
response["choices"][0]["output_text"] = output_text
return response


@pytest.mark.asyncio
Expand Down Expand Up @@ -440,6 +481,42 @@ async def test_derender_completion_roundtrip(client):
assert choices[1]["text"]


@pytest.mark.asyncio
@pytest.mark.parametrize(
("output_text", "include_stop", "min_tokens"),
[
("answer", False, 0),
("answer<END>", True, 0),
("keep<END>allowed", False, 2),
],
)
async def test_derender_completion_uses_engine_stop_text(
client, output_text, include_stop, min_tokens
):
gen_req = await _render_completion(client, "Hello")
token_ids = gen_req["token_ids"][:4]
response = await client.post(
"/v1/completions/derender",
json={
"model": MODEL_NAME,
"generate_responses": [
_make_completion_generate_response(
token_ids, gen_req["request_id"], output_text=output_text
)
],
"completion_request": {
"model": MODEL_NAME,
"prompt": "Hello",
"stop": ["<END>"],
"include_stop_str_in_output": include_stop,
"min_tokens": min_tokens,
},
},
)
assert response.status_code == 200
assert response.json()["choices"][0]["text"] == output_text


@pytest.mark.asyncio
async def test_derender_completion_usage_aggregation(client):
"""prompt_tokens=[5, 10] is aggregated correctly into usage."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,7 @@ async def _fake_preprocess(*args, **kwargs):
def _make_request_output(
request_id: str,
token_ids: list[int],
text: str = "",
finish_reason: str | None = None,
finished: bool = False,
prompt_token_ids: list[int] | None = None,
Expand All @@ -142,7 +143,7 @@ def _make_request_output(
outputs=[
CompletionOutput(
index=index,
text="",
text=text,
token_ids=token_ids,
cumulative_logprob=None,
logprobs=logprobs,
Expand Down Expand Up @@ -285,6 +286,34 @@ async def mock_generate(*args, **kwargs):
await serving.serve_tokens(request)


@pytest.mark.asyncio
async def test_non_stream_preserves_engine_stop_text_and_token_ids():
engine = _mock_engine()

async def mock_generate(*args, **kwargs):
yield _make_request_output(
"req-1",
token_ids=[10, 20, 30],
text="keep<END>ignored earlier, final",
finish_reason="stop",
finished=True,
)

engine.generate = MagicMock(side_effect=mock_generate)
serving = _build_serving_tokens(engine)
request = GenerateRequest(
token_ids=[1, 2, 3],
sampling_params=SamplingParams(max_tokens=10, min_tokens=2, stop=["<END>"]),
model=MODEL_NAME,
)

response = await serving.serve_tokens(request)

assert isinstance(response, GenerateResponse)
assert response.choices[0].output_text == "keep<END>ignored earlier, final"
assert response.choices[0].token_ids == [10, 20, 30]


@pytest.mark.asyncio
async def test_stream_basic():
"""Streaming returns SSE chunks with correct token_ids and ends with [DONE]."""
Expand Down
2 changes: 2 additions & 0 deletions vllm/entrypoints/scale_out/token_in_token_out/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,6 +322,8 @@ class GenerateResponseChoice(BaseModel):
# per OpenAI spec this is the default
finish_reason: str | None = "stop"
token_ids: list[int] | None = None
output_text: str | None = None
"""Engine-truncated text when stop strings require detokenization."""
# Per-token expert routing decisions, base64-encoded ``.npy`` bytes
# (numpy serialization). Shape after decode:
# (num_tokens - 1, num_layers, num_experts_per_tok) dtype uint8/uint16/int32
Expand Down
5 changes: 5 additions & 0 deletions vllm/entrypoints/scale_out/token_in_token_out/serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -349,6 +349,11 @@ async def serve_tokens_full_generator(
logprobs=logprobs,
finish_reason=output.finish_reason if output.finish_reason else "stop",
token_ids=as_list(output.token_ids),
output_text=(
output.text
if sampling_params.stop and sampling_params.detokenize
else None
),
routed_experts=routed_experts_b64,
sampling_mask=sampling_mask,
)
Expand Down
12 changes: 12 additions & 0 deletions vllm/renderers/online_derenderer.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,6 +205,12 @@ def _derender_chat(
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=skip_special
)
if (
chat_request is not None
and chat_request.stop
and choice.output_text is not None
):
decoded_text = choice.output_text
message = ChatMessage(role="assistant", content=decoded_text)

choices.append(
Expand Down Expand Up @@ -645,6 +651,12 @@ def _derender_completion(
decoded_text = tokenizer.decode(
choice.token_ids, skip_special_tokens=skip_special
)
if (
completion_request is not None
and completion_request.stop
and choice.output_text is not None
):
decoded_text = choice.output_text
completion_logprobs = None
if choice.logprobs is not None:
resolved = _resolve_logprobs(choice.logprobs, tokenizer)
Expand Down
Loading