Files
enginex-ascend-910-vllm/tests/e2e/pull_request/one_card/test_guided_decoding.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

292 lines
12 KiB
Python

#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
# Adapted from vllm/tests/entrypoints/llm/test_guided_generate.py
# Copyright 2023 The vLLM team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
import json
import os
from unittest.mock import patch
import jsonschema
import pytest
import regex as re
from vllm.outputs import RequestOutput
from vllm.sampling_params import SamplingParams, StructuredOutputsParams
from tests.e2e.conftest import ModelName
from vllm_ascend.utils import vllm_version_is
os.environ["VLLM_BATCH_INVARIANT"] = "1"
MODEL_NAME = ModelName.QWEN3_06B
GuidedDecodingBackend = ["xgrammar", "guidance", "outlines"]
REGEX_COMPILATION_TIMEOUT_ENV = {"VLLM_REGEX_COMPILATION_TIMEOUT_S": "30"}
@pytest.fixture(params=[False, True], ids=["v1", "v2"])
def model_runner_env(request):
use_v2_model_runner = request.param
if use_v2_model_runner and vllm_version_is("0.23.0"):
pytest.skip("No need to support v2 model runner for vLLM tag version.")
with patch.dict(os.environ, {"VLLM_USE_V2_MODEL_RUNNER": "1" if use_v2_model_runner else "0"}):
yield
@pytest.fixture(scope="module")
def sample_regex():
return (
r"((25[0-5]|(2[0-4]|1\d|[1-9]|)\d)\.){3}"
r"(25[0-5]|(2[0-4]|1\d|[1-9]|)\d)"
)
@pytest.fixture(scope="module")
def sample_json_schema():
return {
"type": "object",
"properties": {
"name": {"type": "string"},
"age": {"type": "integer"},
"skills": {"type": "array", "items": {"type": "string", "maxLength": 10}, "minItems": 3},
"work_history": {
"type": "array",
"items": {
"type": "object",
"properties": {
"company": {"type": "string"},
"duration": {"type": "number"},
"position": {"type": "string"},
},
"required": ["company", "position"],
},
},
},
"required": ["name", "age", "skills", "work_history"],
}
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "xgrammar"}},
)
def test_guided_json_completion_xgrammar(sample_json_schema, request):
sampling_params = SamplingParams(
temperature=1.0, max_tokens=500, structured_outputs=StructuredOutputsParams(json=sample_json_schema)
)
if not vllm_version_is("0.23.0"):
model_marker = request.node.get_closest_marker("model")
model_marker.kwargs["env_vars"] = REGEX_COMPILATION_TIMEOUT_ENV
with patch.dict(os.environ, REGEX_COMPILATION_TIMEOUT_ENV, clear=False):
vllm_runner = request.getfixturevalue("vllm_runner")
prompts = [f"Give an example JSON for an employee profile that fits this schema: {sample_json_schema}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
assert generated_text is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
output_json = json.loads(generated_text)
jsonschema.validate(instance=output_json, schema=sample_json_schema)
else:
vllm_runner = request.getfixturevalue("vllm_runner")
prompts = [f"Give an example JSON for an employee profile that fits this schema: {sample_json_schema}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
assert generated_text is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
output_json = json.loads(generated_text)
jsonschema.validate(instance=output_json, schema=sample_json_schema)
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "xgrammar"}},
)
def test_guided_regex_xgrammar(sample_regex, vllm_runner):
sampling_params = SamplingParams(
temperature=0.8, top_p=0.95, structured_outputs=StructuredOutputsParams(regex=sample_regex)
)
prompts = [f"Give an example IPv4 address with this regex: {sample_regex}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
print(generated_text)
assert generated_text is not None
assert re.fullmatch(".*", generated_text) is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "guidance"}},
)
def test_guided_json_completion_guidance(sample_json_schema, vllm_runner):
sampling_params = SamplingParams(
temperature=1.0, max_tokens=500, structured_outputs=StructuredOutputsParams(json=sample_json_schema)
)
prompts = [f"Give an example JSON for an employee profile that fits this schema: {sample_json_schema}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
assert generated_text is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
output_json = json.loads(generated_text)
jsonschema.validate(instance=output_json, schema=sample_json_schema)
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "guidance"}},
)
def test_guided_regex_guidance(sample_regex, vllm_runner):
sampling_params = SamplingParams(
temperature=0.8, top_p=0.95, structured_outputs=StructuredOutputsParams(regex=sample_regex)
)
prompts = [f"Give an example IPv4 address with this regex: {sample_regex}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
print(generated_text)
assert generated_text is not None
assert re.fullmatch(".*", generated_text) is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "auto"}},
)
def test_guided_auto_rejects_mixed_structured_output_backends(vllm_runner):
xgrammar_schema = {
"type": "object",
"properties": {"name": {"type": "string"}},
"required": ["name"],
}
guidance_schema = {
"type": "object",
"properties": {"count": {"type": "integer", "multipleOf": 2}},
"required": ["count"],
}
xgrammar_params = SamplingParams(
temperature=0.0,
max_tokens=32,
structured_outputs=StructuredOutputsParams(json=xgrammar_schema),
)
prompts = [f"Give an example JSON that fits this schema: {xgrammar_schema}"]
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=xgrammar_params)
assert outputs is not None
assert outputs[0] is not None
guidance_params = SamplingParams(
temperature=0.0,
max_tokens=32,
structured_outputs=StructuredOutputsParams(json=guidance_schema),
)
prompts = [f"Give an example JSON that fits this schema: {guidance_schema}"]
inputs = vllm_runner.get_inputs(prompts)
with pytest.raises(ValueError, match="already using 'xgrammar'.*'guidance'"):
vllm_runner.model.generate(inputs, sampling_params=guidance_params)
@pytest.mark.timeout(1000)
@pytest.mark.model(
model_name=MODEL_NAME,
compilation_config={"cudagraph_capture_sizes": [1, 2, 4, 8]},
extra_kwargs={"seed": 0, "structured_outputs_config": {"backend": "outlines"}},
)
def test_guided_json_completion_outlines(sample_json_schema, request):
sampling_params = SamplingParams(
temperature=1.0, max_tokens=500, structured_outputs=StructuredOutputsParams(json=sample_json_schema)
)
if not vllm_version_is("0.23.0"):
model_marker = request.node.get_closest_marker("model")
model_marker.kwargs["env_vars"] = REGEX_COMPILATION_TIMEOUT_ENV
with patch.dict(os.environ, REGEX_COMPILATION_TIMEOUT_ENV, clear=False):
vllm_runner = request.getfixturevalue("vllm_runner")
prompts = [f"Give an example JSON for an employee profile that fits this schema: {sample_json_schema}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
assert generated_text is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
output_json = json.loads(generated_text)
jsonschema.validate(instance=output_json, schema=sample_json_schema)
else:
vllm_runner = request.getfixturevalue("vllm_runner")
prompts = [f"Give an example JSON for an employee profile that fits this schema: {sample_json_schema}"] * 2
inputs = vllm_runner.get_inputs(prompts)
outputs = vllm_runner.model.generate(inputs, sampling_params=sampling_params)
assert outputs is not None
for output in outputs:
assert output is not None
assert isinstance(output, RequestOutput)
prompt = output.prompt
generated_text = output.outputs[0].text
assert generated_text is not None
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")
output_json = json.loads(generated_text)
jsonschema.validate(instance=output_json, schema=sample_json_schema)