Download app.py from WildOjisan/embeddinggemma_300m_fastapi: direct link, hf CLI and curl.
- Browser
- Download file 8.73 kB
-
https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/d6e603b14fc2e95ed59fb700fca2c3e166505a4f/app.py
- Command line
-
hf download hf://spaces/WildOjisan/embeddinggemma_300m_fastapi@d6e603b14fc2e95ed59fb700fca2c3e166505a4f/app.py
-
curl -L -o app.py https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/d6e603b14fc2e95ed59fb700fca2c3e166505a4f/app.py
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 | |
| 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 /) | |
| 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 | |
| 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 | |
| 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๋ ๊ธฐ์กด ์ ์ง | |
| # ---------------------------------------------------- | |
| 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 | |
| def read_item(q: str | None = None): | |
| return {"q": q, "description": "This is a query test."} | |
| 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) |