Testowy / app.py
StrawberryJelly's picture
Update app.py
22c09f6 verified
Raw
History Blame Contribute Delete
5.07 kB
import os
import json
from fastapi import FastAPI, HTTPException, Request, Header
from fastapi.responses import JSONResponse, StreamingResponse
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import List, Optional
from huggingface_hub import hf_hub_download
from llama_cpp import Llama
app = FastAPI(title="KAI STUDIO")
# CORS żeby działało z przeglądarki nawet przy błędach
@app.middleware("http")
async def add_cors_headers(request: Request, call_next):
try:
response = await call_next(request)
except Exception as e:
response = JSONResponse({"detail": str(e)}, status_code=500)
response.headers["Access-Control-Allow-Origin"] = "*"
response.headers["Access-Control-Allow-Methods"] = "POST, GET, OPTIONS"
response.headers["Access-Control-Allow-Headers"] = "*"
return response
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Konfiguracja modelu
MODEL_REPO = "lmstudio-community/Meta-Llama-3.1-8B-Instruct-GGUF"
MODEL_FILE = "Meta-Llama-3.1-8B-Instruct-Q4_K_M.gguf"
MODEL_PATH = f"/data/{MODEL_FILE}"
API_KEY = os.environ.get("LLAMA_API_KEY") # opcjonalny
# System prompt wymuszający polski
SYSTEM_PL = {
"role": "system",
"content": "Jesteś pomocnym asystentem o imieniu KAI. Zawsze odpowiadaj w języku polskim. Odpowiadaj zwięźle i na temat. Nie używaj angielskiego chyba że użytkownik wyraźnie o to poprosi."
}
llm = None
@app.on_event("startup")
async def startup_event():
global llm
# 1. Pobierz model do /data jeśli go nie ma
if not os.path.exists(MODEL_PATH):
print(f"Model not found at {MODEL_PATH}. Downloading from HF Hub...")
try:
hf_hub_download(
repo_id=MODEL_REPO,
filename=MODEL_FILE,
local_dir="/data"
)
print("Download complete")
except Exception as e:
print(f"Download failed: {e}")
raise RuntimeError(f"Could not download model: {e}")
# 2. Walidacja rozmiaru
size_gb = os.path.getsize(MODEL_PATH) / 1024**3
print(f"Model file: {MODEL_PATH} - {size_gb:.2f} GB")
if size_gb < 4.0:
raise RuntimeError(f"Model file too small: {size_gb:.2f} GB. Likely corrupted.")
# 3. Załaduj model
print("Loading model into memory...")
llm = Llama(
model_path=MODEL_PATH,
n_ctx=2048,
n_gpu_layers=0,
verbose=True,
chat_format="llama-3"
)
print("Model ready")
# Modele requestów
class Message(BaseModel):
role: str
content: str
class ChatRequest(BaseModel):
messages: List[Message]
max_tokens: Optional[int] = 256
temperature: Optional[float] = 0.7
top_p: Optional[float] = 0.95
stream: Optional[bool] = False
def check_key(x_api_key: str = Header(None, alias="x-api-key")):
if API_KEY and x_api_key!= API_KEY:
raise HTTPException(status_code=401, detail="Invalid API key")
@app.post("/v1/chat/completions")
async def chat_completions(data: ChatRequest, x_api_key: str = Header(None, alias="x-api-key")):
check_key(x_api_key)
if llm is None:
raise HTTPException(status_code=503, detail="Model not loaded yet")
messages = [m.dict() for m in data.messages]
# Wstrzykujemy system prompt PL jeśli user go nie podał
if not messages or messages[0]["role"]!= "system":
messages = [SYSTEM_PL] + messages
try:
if data.stream:
def generate():
for chunk in llm.create_chat_completion(
messages=messages,
max_tokens=data.max_tokens,
temperature=data.temperature,
top_p=data.top_p,
stop=["<|eot_id|>", "<|end_of_text|>", "<|eom_id|>"],
stream=True
):
yield f"data: {json.dumps(chunk)}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(generate(), media_type="text/event-stream")
else:
res = llm.create_chat_completion(
messages=messages,
max_tokens=data.max_tokens,
temperature=data.temperature,
top_p=data.top_p,
stop=["<|eot_id|>", "<|end_of_text|>", "<|eom_id|>"]
)
return JSONResponse(content=res)
except Exception as e:
print(f"Inference error: {e}")
raise HTTPException(status_code=500, detail=str(e))
@app.get("/health")
async def health():
return {
"status": "ok",
"model_loaded": llm is not None,
"model_path": MODEL_PATH,
"model_exists": os.path.exists(MODEL_PATH),
"model_size_gb": round(os.path.getsize(MODEL_PATH) / 1024**3, 2) if os.path.exists(MODEL_PATH) else 0
}
@app.get("/")
async def root():
return {"message": "KAI STUDIO API", "docs": "/docs", "health": "/health"}