135 lines
4.9 KiB
Python
135 lines
4.9 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
#
|
|
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
|
|
|
from __future__ import annotations
|
|
|
|
from inspect import Signature, signature
|
|
from typing import Any
|
|
|
|
from vllm.exceptions import VLLMValidationError
|
|
from vllm.sampling_params import SamplingParams
|
|
from vllm.v1.structured_output import StructuredOutputManager
|
|
|
|
_BACKEND_ATTR = "_vllm_ascend_structured_output_backend"
|
|
_ORIGINAL_GRAMMAR_INIT_ATTR = "_vllm_ascend_original_grammar_init"
|
|
_ORIGINAL_VALIDATE_ATTR = "_vllm_ascend_original_validate_structured_outputs"
|
|
|
|
|
|
def _request_backend(request: Any) -> str | None:
|
|
if getattr(request, "structured_output_request", None) is None:
|
|
return None
|
|
|
|
sampling_params = getattr(request, "sampling_params", None)
|
|
structured_outputs = getattr(sampling_params, "structured_outputs", None)
|
|
backend = getattr(structured_outputs, "_backend", None)
|
|
return backend if isinstance(backend, str) else None
|
|
|
|
|
|
def _backend_name_from_instance(backend: Any) -> str | None:
|
|
if backend is None:
|
|
return None
|
|
|
|
backend_names = {
|
|
"XgrammarBackend": "xgrammar",
|
|
"GuidanceBackend": "guidance",
|
|
"OutlinesBackend": "outlines",
|
|
"LMFormatEnforcerBackend": "lm-format-enforcer",
|
|
}
|
|
for backend_cls in type(backend).__mro__:
|
|
for class_name, backend_name in backend_names.items():
|
|
if class_name in backend_cls.__name__:
|
|
return backend_name
|
|
return None
|
|
|
|
|
|
def _raise_mixed_backend(initialized_backend: str, request_backend: str) -> None:
|
|
raise VLLMValidationError(
|
|
"V1 structured outputs only supports one backend per engine. "
|
|
f"The engine is already using '{initialized_backend}', but "
|
|
f"this request resolved to '{request_backend}'. Configure "
|
|
"`structured_outputs_config.backend` explicitly or use schemas "
|
|
"supported by the initialized backend."
|
|
)
|
|
|
|
|
|
def _sampling_params_backend(sampling_params: SamplingParams) -> str | None:
|
|
structured_outputs = getattr(sampling_params, "structured_outputs", None)
|
|
backend = getattr(structured_outputs, "_backend", None)
|
|
return backend if isinstance(backend, str) else None
|
|
|
|
|
|
def _structured_outputs_config_from_call(
|
|
validate_signature: Signature,
|
|
sampling_params: SamplingParams,
|
|
args: tuple[Any, ...],
|
|
kwargs: dict[str, Any],
|
|
) -> Any:
|
|
bound_arguments = validate_signature.bind_partial(
|
|
sampling_params,
|
|
*args,
|
|
**kwargs,
|
|
)
|
|
return bound_arguments.arguments.get("structured_outputs_config")
|
|
|
|
|
|
def _patch_sampling_params_validation() -> None:
|
|
original_validate = SamplingParams._validate_structured_outputs
|
|
validate_signature = signature(original_validate)
|
|
setattr(SamplingParams, _ORIGINAL_VALIDATE_ATTR, original_validate)
|
|
|
|
def _validate_structured_outputs(
|
|
self: SamplingParams,
|
|
*args: Any,
|
|
**kwargs: Any,
|
|
) -> None:
|
|
result = original_validate(self, *args, **kwargs)
|
|
structured_outputs_config = _structured_outputs_config_from_call(
|
|
validate_signature,
|
|
self,
|
|
args,
|
|
kwargs,
|
|
)
|
|
request_backend = _sampling_params_backend(self)
|
|
if structured_outputs_config is None or request_backend is None:
|
|
return result
|
|
|
|
initialized_backend = getattr(structured_outputs_config, _BACKEND_ATTR, None)
|
|
if initialized_backend is not None and request_backend != initialized_backend:
|
|
_raise_mixed_backend(initialized_backend, request_backend)
|
|
|
|
setattr(structured_outputs_config, _BACKEND_ATTR, request_backend)
|
|
return result
|
|
|
|
SamplingParams._validate_structured_outputs = _validate_structured_outputs
|
|
|
|
|
|
def _patch_structured_output_manager() -> None:
|
|
original_grammar_init = StructuredOutputManager.grammar_init
|
|
setattr(StructuredOutputManager, _ORIGINAL_GRAMMAR_INIT_ATTR, original_grammar_init)
|
|
|
|
def grammar_init(self: StructuredOutputManager, request: Any) -> None:
|
|
request_backend = _request_backend(request)
|
|
if request_backend is None:
|
|
return original_grammar_init(self, request)
|
|
|
|
initialized_backend = getattr(self, _BACKEND_ATTR, None)
|
|
if initialized_backend is None:
|
|
initialized_backend = _backend_name_from_instance(getattr(self, "backend", None))
|
|
if initialized_backend is not None:
|
|
setattr(self, _BACKEND_ATTR, initialized_backend)
|
|
|
|
if initialized_backend is not None and request_backend != initialized_backend:
|
|
_raise_mixed_backend(initialized_backend, request_backend)
|
|
|
|
result = original_grammar_init(self, request)
|
|
if getattr(self, "backend", None) is not None:
|
|
setattr(self, _BACKEND_ATTR, request_backend)
|
|
return result
|
|
|
|
StructuredOutputManager.grammar_init = grammar_init
|
|
|
|
|
|
_patch_sampling_params_validation()
|
|
_patch_structured_output_manager()
|