WildOjisan's picture
.
0f1d138
Raw History Blame
8.73 kB
from contextlib import asynccontextmanager
from fastapi import FastAPI, Depends, HTTPException
from pydantic import BaseModel
import uvicorn
import asyncpg
#import torch
from typing import List, Dict, Any, Union
import os
import numpy as np
# ONNX ๋ฐ HuggingFace ๊ด€๋ จ ์ž„ํฌํŠธ
from huggingface_hub import hf_hub_download
import onnxruntime as ort
from transformers import AutoTokenizer
from cnn_router import router as cnn_router, init_vision_model
from database_conn import connect_to_db, close_db_connection, get_db_connection
from test_router import router as test_router
HF_TOKEN = os.getenv("HUGGING_FACE_HUB_TOKEN")
# --- Wrapper Class ์ •์˜ ---
# ๊ธฐ์กด SentenceTransformer์™€ ๋™์ผํ•œ ๋ฉ”์„œ๋“œ(encode_document, encode_query, similarity)๋ฅผ ์ œ๊ณต
class OnnxGemmaWrapper:
def __init__(self, model_id, token=None):
print(f"Loading ONNX model: {model_id}...")
self.tokenizer = AutoTokenizer.from_pretrained(model_id, token=token)
# ONNX ๋ชจ๋ธ ๋ฐ ๊ฐ€์ค‘์น˜ ๋‹ค์šด๋กœ๋“œ
model_path = hf_hub_download(model_id, subfolder="onnx", filename="model.onnx", token=token)
hf_hub_download(model_id, subfolder="onnx", filename="model.onnx_data", token=token)
# ์ถ”๋ก  ์„ธ์…˜ ์ƒ์„ฑ (GPU ์‚ฌ์šฉ ๊ฐ€๋Šฅ ์‹œ CUDAProvider ์‚ฌ์šฉ, ์—†์œผ๋ฉด CPU)
available_providers = ort.get_available_providers()
if 'CUDAExecutionProvider' in available_providers:
print("CUDA detected. Using GPU.")
providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
else:
print("CUDA not detected. Using CPU.")
providers = ['CPUExecutionProvider']
self.session = ort.InferenceSession(model_path, providers=providers)
# Prefix ์ •์˜
self.prefixes = {
"query": "task: search result | query: ",
"document": "title: none | text: ",
}
print("ONNX Model loaded successfully.")
def _run_inference(self, texts: List[str]):
inputs = self.tokenizer(texts, padding=True, truncation=True, return_tensors="np")
# ONNX Runtime ์‹คํ–‰ (output[0]: last_hidden_state, output[1]: pooler_output or sentence_embedding)
# EmbeddingGemma ONNX ๋ชจ๋ธ์€ ๋ณดํ†ต ๋‘ ๋ฒˆ์งธ ๋ฆฌํ„ด๊ฐ’์ด sentence embedding์ž…๋‹ˆ๋‹ค.
outputs = self.session.run(None, dict(inputs))
# outputs[1]์ด (Batch, 768) ํ˜•ํƒœ์˜ ์ž„๋ฒ ๋”ฉ
return outputs[1]
def encode_document(self, documents: List[str]) -> np.ndarray:
# ๋ฌธ์„œ์šฉ Prefix ์ถ”๊ฐ€
prefixed_docs = [self.prefixes["document"] + doc for doc in documents]
return self._run_inference(prefixed_docs)
def encode_query(self, query: str) -> np.ndarray:
# ์ฟผ๋ฆฌ์šฉ Prefix ์ถ”๊ฐ€ (๋‹จ์ผ ์ฟผ๋ฆฌ๋„ ๋ฆฌ์ŠคํŠธ๋กœ ์ฒ˜๋ฆฌ)
prefixed_query = [self.prefixes["query"] + query]
# ๊ฒฐ๊ณผ๋Š” (1, 768) ํ˜•ํƒœ์ด๋ฏ€๋กœ ์ฒซ ๋ฒˆ์งธ ์š”์†Œ๋ฅผ ๋ฐ˜ํ™˜ํ•˜์—ฌ (768,)๋กœ ๋งž์ถœ ์ˆ˜๋„ ์žˆ์œผ๋‚˜,
# ๊ธฐ์กด ๋กœ์ง๊ณผ์˜ ํ˜ธํ™˜์„ฑ์„ ์œ„ํ•ด ๋ฐฐ์น˜ ์ฐจ์›์„ ์œ ์ง€ํ•˜๊ฑฐ๋‚˜ ํ•„์š” ์‹œ ์กฐ์ •.
# ์—ฌ๊ธฐ์„œ๋Š” (1, 768) ํ˜•ํƒœ๋กœ ๋ฐ˜ํ™˜ํ•ฉ๋‹ˆ๋‹ค.
return self._run_inference(prefixed_query)[0]
def similarity(self, query_emb: np.ndarray, doc_embs: np.ndarray) -> np.ndarray:
# ์ฝ”์‚ฌ์ธ ์œ ์‚ฌ๋„ ๊ณ„์‚ฐ (Dot Product)
# query_emb: (768,) ๋˜๋Š” (1, 768)
# doc_embs: (N, 768)
# ์ฐจ์› ๋งž์ถ”๊ธฐ (query_emb๊ฐ€ 1์ฐจ์›์ด๋ฉด 2์ฐจ์›์œผ๋กœ ๋ณ€ํ™˜)
if query_emb.ndim == 1:
query_emb = query_emb.reshape(1, -1)
# Dot Product ์ˆ˜ํ–‰ (@ ์—ฐ์‚ฐ์ž)
scores = query_emb @ doc_embs.T
# ๊ฒฐ๊ณผ๊ฐ€ (1, N) ํ˜•ํƒœ์ด๋ฏ€๋กœ 1์ฐจ์› ๋ฐฐ์—ด (N,)์œผ๋กœ ๋ณ€ํ™˜ํ•˜์—ฌ ๋ฐ˜ํ™˜
return scores.flatten()
# ์ „์—ญ ๋ณ€์ˆ˜ ์ดˆ๊ธฐํ™”
model = None
@asynccontextmanager
async def lifespan(app: FastAPI):
global model
try:
await connect_to_db()
except Exception as e:
print(f"!!! [Startup] DB Connection FAILED: {e!r}")
# --- ๋ชจ๋ธ ๋กœ๋“œ ---
try:
model = OnnxGemmaWrapper(
model_id="onnx-community/embeddinggemma-300m-ONNX",
token=HF_TOKEN
)
except Exception as e:
print(f"Error loading ONNX model: {e}")
model = None
# --- Vision Model ๋กœ๋“œ ---
try:
init_vision_model()
except Exception as e:
print(f"Error loading Vision model: {e}")
yield
await close_db_connection()
print(">>> [Shutdown] FastAPI Server graceful shutdown complete.")
app = FastAPI(
title="Gemma Embedding Service (ONNX)",
description="Implements text embedding generation via REST API using ONNX Runtime.",
version="1.0.0",
lifespan=lifespan
)
# 3. ๋ฃจํŠธ ์—”๋“œํฌ์ธํŠธ (GET /)
@app.get("/")
def read_root():
result={"success":True,"data":None,"msg":""}
try:
result["data"]="ok"
return result
except Exception as e:
result["success"] = False
result["msg"]=f"server error. {e!r}"
return result
app.include_router(test_router, prefix="/api/test")
from cnn_router import router as cnn_router
app.include_router(cnn_router, prefix="/api/cnn")
class Item(BaseModel):
name: str
price: float
is_offer: bool | None = None
class MakeTextEmbedding(BaseModel):
query: str
documents: List[str]
class EmbeddingOutput(BaseModel):
success: bool
msg: str
data: Union[List[List[float]], None] = None
@app.post("/make_text_embedding", summary="Calculate semantic similarity and find the best match")
async def calculate_similarity(data: MakeTextEmbedding):
result={"success":True,"data":None,"msg":""}
try:
if model is None:
result["success"] = False
result["msg"]="Model not loaded. Service is unavailable."
return result
# ONNX Runtime์€ ๋‚ด๋ถ€์ ์œผ๋กœ ์ตœ์ ํ™”๋˜์–ด ์žˆ์œผ๋ฏ€๋กœ torch.no_grad() ๋ถˆํ•„์š”ํ•˜์ง€๋งŒ,
# ๊ธฐ์กด ํ๋ฆ„์ƒ ๊ทธ๋ƒฅ ๋‘ฌ๋„ ์ƒ๊ด€์—†๊ฑฐ๋‚˜ ์ œ๊ฑฐํ•ด๋„ ๋ฉ๋‹ˆ๋‹ค. ์—ฌ๊ธฐ์„  ์ œ๊ฑฐํ•ฉ๋‹ˆ๋‹ค.
# Encode documents
document_embeddings = model.encode_document(data.documents)
# numpy array -> list ๋ณ€ํ™˜
embeddings_list = document_embeddings.tolist()
result["data"] = embeddings_list
result["msg"] = f"document_embeddings.shape:{document_embeddings.shape}"
return result
except Exception as e:
result["success"] = False
result["msg"] = f"server error. {e!r}"
return result
@app.post("/string_distance_compare", summary="Calculate semantic similarity and find the best match")
async def string_distance_compare(data: MakeTextEmbedding):
result={"success":True,"data":None,"msg":""}
try:
if model is None:
result["success"] = False
result["msg"]="Model not loaded. Service is unavailable."
return result
# Encode query and documents
query_embeddings = model.encode_query(data.query)
document_embeddings = model.encode_document(data.documents)
# Calculate similarity
similarities = model.similarity(query_embeddings, document_embeddings)
result["data"] = similarities.tolist()
result["msg"] = f"query_embeddings.shape: {query_embeddings.shape}, document_embeddings.shape: {document_embeddings.shape}"
return result
except Exception as e:
result["success"] = False
result["msg"] = f"server error. {e!r}"
return result
# ----------------------------------------------------
# DB ๊ด€๋ จ ์—”๋“œํฌ์ธํŠธ ๋ฐ ๊ธฐํƒ€ API๋Š” ๊ธฐ์กด ์œ ์ง€
# ----------------------------------------------------
@app.get("/time", response_model=Dict[str, Any])
async def get_db_time(conn: asyncpg.Connection = Depends(get_db_connection)):
result = {"success": True, "data": None, "msg": ""}
try:
query = "SELECT NOW();"
records = await conn.fetch(query)
data_list_of_dicts: List[Dict[str, Any]] = [dict(record) for record in records]
result["data"] = data_list_of_dicts[0]
except Exception as e:
result["success"] = False
result["msg"] = f"Database query error: {e!r}"
return result
@app.get("/items")
def read_item(q: str | None = None):
return {"q": q, "description": "This is a query test."}
@app.post("/items/")
def create_item(item: Item):
if item.price > 100.0:
item.name = f"Premium {item.name}"
return {"message": "Item created successfully", "item_data": item}
if __name__ == "__main__":
uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True)