sarajane's picture
Create app.py
84bf033 verified
Raw
History Blame
5.21 kB
import os
import time
import uuid
from typing import Any, Dict, List, Optional
import torch
from fastapi import FastAPI, Header, HTTPException
from fastapi.responses import PlainTextResponse
from pydantic import BaseModel
from transformers import AutoModelForCausalLM, AutoTokenizer
MODEL_ID = os.environ.get(
"MODEL_ID",
"Polygl0t/Tucano2-qwen-0.5B-Instruct"
)
API_KEY = os.environ.get("API_KEY", "")
MAX_INPUT_TOKENS = int(os.environ.get("MAX_INPUT_TOKENS", "2048"))
DEFAULT_MAX_NEW_TOKENS = int(os.environ.get("DEFAULT_MAX_NEW_TOKENS", "80"))
tokenizer = None
model = None
app = FastAPI(title="Tucano OpenAI-Compatible API")
class Message(BaseModel):
role: str
content: str
class ChatCompletionRequest(BaseModel):
model: Optional[str] = None
messages: List[Message]
temperature: Optional[float] = 0.1
max_tokens: Optional[int] = DEFAULT_MAX_NEW_TOKENS
top_p: Optional[float] = 0.95
stream: Optional[bool] = False
def check_auth(authorization: Optional[str]) -> None:
if not API_KEY:
return
if not authorization:
raise HTTPException(status_code=401, detail="Missing Authorization header")
expected = f"Bearer {API_KEY}"
if authorization != expected:
raise HTTPException(status_code=403, detail="Invalid API key")
@app.on_event("startup")
def load_model() -> None:
global tokenizer, model
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID,
trust_remote_code=True
)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.float32,
low_cpu_mem_usage=True,
trust_remote_code=True
)
model.eval()
@app.get("/", response_class=PlainTextResponse)
def root() -> str:
return "Tucano OpenAI-compatible API is running."
@app.get("/health")
def health() -> Dict[str, str]:
return {
"status": "ok",
"model": MODEL_ID
}
@app.get("/v1/models")
def list_models(
authorization: Optional[str] = Header(default=None)
) -> Dict[str, Any]:
check_auth(authorization)
return {
"object": "list",
"data": [
{
"id": MODEL_ID,
"object": "model",
"created": 0,
"owned_by": "huggingface-space"
}
]
}
@app.post("/v1/chat/completions")
def chat_completions(
request: ChatCompletionRequest,
authorization: Optional[str] = Header(default=None)
) -> Dict[str, Any]:
check_auth(authorization)
if request.stream:
raise HTTPException(
status_code=400,
detail="Streaming is not supported by this pilot endpoint"
)
if not request.messages:
raise HTTPException(
status_code=400,
detail="At least one message is required"
)
messages = [
{
"role": m.role,
"content": m.content
}
for m in request.messages
]
try:
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
except Exception:
prompt = ""
for m in messages:
prompt += f"{m['role']}: {m['content']}\n"
prompt += "assistant:"
inputs = tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=MAX_INPUT_TOKENS
)
max_new_tokens = request.max_tokens or DEFAULT_MAX_NEW_TOKENS
generation_args = {
"input_ids": inputs["input_ids"],
"attention_mask": inputs.get("attention_mask"),
"max_new_tokens": max_new_tokens,
"pad_token_id": tokenizer.eos_token_id,
}
if request.temperature is not None and request.temperature > 0:
generation_args["do_sample"] = True
generation_args["temperature"] = request.temperature
generation_args["top_p"] = request.top_p or 0.95
else:
generation_args["do_sample"] = False
start_time = time.time()
with torch.no_grad():
output_ids = model.generate(**generation_args)
elapsed = time.time() - start_time
prompt_length = inputs["input_ids"].shape[-1]
generated_ids = output_ids[0][prompt_length:]
content = tokenizer.decode(
generated_ids,
skip_special_tokens=True
).strip()
prompt_tokens = int(prompt_length)
completion_tokens = int(generated_ids.shape[-1])
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": int(time.time()),
"model": request.model or MODEL_ID,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": content
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens
},
"pilot_metadata": {
"served_model": MODEL_ID,
"elapsed_seconds": round(elapsed, 3)
}
}