Files
zenml/venv/lib/python3.9/site-packages/mlflow/gateway/providers/huggingface.py
Christian Mantha 2ca0b9ef7c star
2026-03-02 19:10:52 -05:00

116 lines
4.5 KiB
Python

import time
from typing import Any
from mlflow.gateway.config import HuggingFaceTextGenerationInferenceConfig, RouteConfig
from mlflow.gateway.exceptions import AIGatewayException
from mlflow.gateway.providers.base import BaseProvider
from mlflow.gateway.providers.utils import (
rename_payload_keys,
send_request,
)
from mlflow.gateway.schemas import completions
class HFTextGenerationInferenceServerProvider(BaseProvider):
NAME = "Hugging Face Text Generation Inference"
CONFIG_TYPE = HuggingFaceTextGenerationInferenceConfig
def __init__(self, config: RouteConfig) -> None:
super().__init__(config)
if config.model.config is None or not isinstance(
config.model.config, HuggingFaceTextGenerationInferenceConfig
):
raise TypeError(f"Unexpected config type {config.model.config}")
self.huggingface_config: HuggingFaceTextGenerationInferenceConfig = config.model.config
self.headers = {"Content-Type": "application/json"}
async def _request(self, path: str, payload: dict[str, Any]) -> dict[str, Any]:
return await send_request(
headers=self.headers,
base_url=self.huggingface_config.hf_server_url,
path=path,
payload=payload,
)
async def completions(self, payload: completions.RequestPayload) -> completions.ResponsePayload:
from fastapi.encoders import jsonable_encoder
payload = jsonable_encoder(payload, exclude_none=True)
self.check_for_model_field(payload)
key_mapping = {
"max_tokens": "max_new_tokens",
}
for k1, k2 in key_mapping.items():
if k2 in payload:
raise AIGatewayException(
status_code=422, detail=f"Invalid parameter {k2}. Use {k1} instead."
)
# HF TGI does not support generating multiple candidates.
n = payload.pop("n", 1)
if n != 1:
raise AIGatewayException(
status_code=422,
detail="'n' must be '1' for the Text Generation Inference provider."
f"Received value: '{n}'.",
)
prompt = payload.pop("prompt")
parameters = rename_payload_keys(payload, key_mapping)
# The range of HF TGI's temperature is 0-100, but ours is 0-2, so we multiply
# by 50
payload["temperature"] = 50 * payload["temperature"]
# HF TGI does not support 0 temperature
parameters["temperature"] = max(payload["temperature"], 1e-3)
parameters["details"] = True
parameters["decoder_input_details"] = True
final_payload = {"inputs": prompt, "parameters": parameters}
resp = await self._request(
"generate",
final_payload,
)
# Example Response:
# Documentation: https://huggingface.github.io/text-generation-inference/#/Text%20Generation%20Inference/compat_generate
# {'details': {'best_of_sequences': [{'finish_reason': 'length',
# 'generated_text': 'test',
# 'generated_tokens': 1,
# 'prefill': [{'id': 0, 'logprob': -0.34, 'text': 'test'}],
# 'seed': 42,
# 'tokens': [{'id': 0, 'logprob': -0.34, 'special': False, 'text': 'test'}],
# 'top_tokens': [[{'id': 0,
# 'logprob': -0.34,
# 'special': False,
# 'text': 'test'}]]}],
# 'finish_reason': 'length',
# 'generated_tokens': 1,
# 'prefill': [{'id': 0, 'logprob': -0.34, 'text': 'test'}],
# 'seed': 42,
# 'tokens': [{'id': 0, 'logprob': -0.34, 'special': False, 'text': 'test'}],
# 'top_tokens': [[{'id': 0,
# 'logprob': -0.34,
# 'special': False,
# 'text': 'test'}]]},
# 'generated_text': 'test'}
output_tokens = resp["details"]["generated_tokens"]
input_tokens = len(resp["details"]["prefill"])
return completions.ResponsePayload(
created=int(time.time()),
object="text_completion",
model=self.config.model.name,
choices=[
completions.Choice(
index=0,
text=resp["generated_text"],
finish_reason=resp["details"]["finish_reason"],
)
],
usage=completions.CompletionsUsage(
prompt_tokens=input_tokens,
completion_tokens=output_tokens,
total_tokens=input_tokens + output_tokens,
),
)