Download app_bk_260112.py from WildOjisan/embeddinggemma_300m_fastapi: direct link, hf CLI and curl.
- Browser
- Download file 7.7 kB
-
https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/main/app_bk_260112.py
- Command line
-
hf download hf://spaces/WildOjisan/embeddinggemma_300m_fastapi/app_bk_260112.py
-
curl -L -o app_bk_260112.py https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/main/app_bk_260112.py
7.7 kB
| from contextlib import asynccontextmanager | |
| from fastapi import FastAPI, Depends, HTTPException | |
| from pydantic import BaseModel | |
| import uvicorn | |
| import asyncpg | |
| from sentence_transformers import SentenceTransformer | |
| import torch | |
| from typing import List, Dict, Any, Union | |
| from database_conn import connect_to_db, close_db_connection, get_db_connection | |
| from test_router import router as test_router | |
| import os # ํ๊ฒฝ ๋ณ์๋ฅผ ์ฝ๊ธฐ ์ํด os ๋ชจ๋ ์ถ๊ฐ | |
| HF_TOKEN = os.getenv("HUGGING_FACE_HUB_TOKEN") | |
| # Load the SentenceTransformer model once when the application starts | |
| # This is a synchronous operation, which is fine for app startup. | |
| try: | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| model = SentenceTransformer( | |
| "google/embeddinggemma-300m", | |
| device=device, | |
| token=HF_TOKEN | |
| ) | |
| # Set the model to evaluation mode | |
| model.eval() | |
| print("Model 'google/embeddinggemma-300m' loaded successfully.") | |
| except Exception as e: | |
| print(f"Error loading model: {e}") | |
| # In a real application, you might want to raise an exception or handle this more gracefully | |
| # For simplicity, we'll let the app potentially fail if the model can't load. | |
| model = None # Keep model as None if loading failed | |
| # 2. ๋ฐ์ดํฐ ์ ํจ์ฑ ๊ฒ์ฌ๋ฅผ ์ํ Pydantic ๋ชจ๋ธ ์ ์ | |
| # ํด๋ผ์ด์ธํธ๋ก๋ถํฐ ๋ฐ์ ์์ฒญ(Request) ๋ฐ์ดํฐ ๊ตฌ์กฐ๋ฅผ ์ ์ํฉ๋๋ค. | |
| class Item(BaseModel): | |
| name: str | |
| price: float | |
| is_offer: bool | None = None | |
| async def lifespan(app: FastAPI): | |
| try: | |
| await connect_to_db() | |
| except Exception as e: | |
| print(f"!!! [Startup] DB Connection FAILED: {e!r}") | |
| # DB ์ฐ๊ฒฐ ์คํจ ์์๋ ์๋ฒ๋ฅผ ๋์ธ์ง ์ฌ๋ถ๋ ์๋น์ค ์ ์ฑ ์ ๋ฐ๋ผ ๊ฒฐ์ | |
| # --- ์๋ฒ ์คํ ์์ (Yield) --- | |
| yield | |
| # --- ์๋ฒ ์ข ๋ฃ ์ (Shutdown Event) --- | |
| await close_db_connection() | |
| print(">>> [Shutdown] FastAPI Server graceful shutdown complete.") | |
| app = FastAPI( | |
| title="Gemma Embedding Service", | |
| description="Implements text embedding generation via REST API.", | |
| version="1.0.0", | |
| lifespan=lifespan # <-- ์ฌ๊ธฐ์์ lifespan ํจ์๋ฅผ ๋ฑ๋กํฉ๋๋ค. | |
| ) | |
| # --- API ์๋ํฌ์ธํธ ์ ์ --- | |
| # 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") | |
| class MakeTextEmbedding(BaseModel): | |
| """ | |
| Input structure for the similarity endpoint. | |
| """ | |
| 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"]=f"Model not loaded. Service is unavailable." | |
| return result | |
| # The 'with torch.no_grad():' block is essential for efficient inference | |
| with torch.no_grad(): | |
| # Encode the query (single vector) | |
| document_embeddings = model.encode_document(data.documents) | |
| 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"]=f"Model not loaded. Service is unavailable." | |
| return result | |
| # The 'with torch.no_grad():' block is essential for efficient inference | |
| with torch.no_grad(): | |
| # Encode the query (single vector) | |
| query_embeddings = model.encode_query(data.query) | |
| document_embeddings = model.encode_document(data.documents) | |
| 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 | |
| # ---------------------------------------------------- | |
| # 2. API ์๋ํฌ์ธํธ: SELECT NOW() (์ฉ ์ฟผ๋ฆฌ ์คํ) | |
| # ---------------------------------------------------- | |
| async def get_db_time( | |
| # get_db_connection์ ํตํด asyncpg.Connection ๊ฐ์ฒด๋ง ์ฃผ์ ๋ฐ์ต๋๋ค. | |
| conn: asyncpg.Connection = Depends(get_db_connection) | |
| ): | |
| """ | |
| DB์ ์ ์ํ์ฌ ํ์ฌ ์๊ฐ์ ์กฐํํ๋ ์ฉ SQL ์ฟผ๋ฆฌ(SELECT NOW())๋ฅผ ์คํํฉ๋๋ค. | |
| """ | |
| result = {"success": True, "data": None, "msg": ""} | |
| try: | |
| # **์ฌ๊ธฐ์ ์ฉ SQL ์ฟผ๋ฆฌ๋ฅผ ์ง์ ์์ฑํ๊ณ ์คํํฉ๋๋ค.** | |
| query = "SELECT NOW();" | |
| # ์ฟผ๋ฆฌ ์คํ (fetchval()์ ์ฟผ๋ฆฌ ๊ฒฐ๊ณผ์ ์ฒซ ํ, ์ฒซ ์ด์ ๊ฐ๋ง ๋ฐํํฉ๋๋ค.) | |
| # asyncpg๊ฐ DB ์๊ฐ์ Python datetime ๊ฐ์ฒด๋ก ๋ณํํด ์ค๋๋ค. | |
| records = await conn.fetch(query) | |
| data_list_of_dicts: List[Dict[str, Any]] = [ | |
| # dict(record)๋ฅผ ์ฌ์ฉํ์ฌ Record ๊ฐ์ฒด๋ฅผ ์ผ๋ฐ ํ์ด์ฌ ๋์ ๋๋ฆฌ๋ก ๋ณํ | |
| dict(record) for record in records | |
| ] | |
| # ๊ฒฐ๊ณผ๋ฅผ ๋ฌธ์์ด๋ก ๋ณํํ์ฌ JSON ์๋ต์ ๋ด์ต๋๋ค. | |
| result["data"] = data_list_of_dicts[0] | |
| except Exception as e: | |
| # ์ฟผ๋ฆฌ ์คํ ์ค ์ค๋ฅ๊ฐ ๋ฐ์ํ๋ฉด 500 ์๋ฌ๋ฅผ ๋ฐํ | |
| result["success"] = False | |
| result["msg"] = f"Database query error: {e!r}" | |
| return result | |
| # 4. ์ฟผ๋ฆฌ ํ๋ผ๋ฏธํฐ๋ฅผ ๋ฐ๋ ์๋ํฌ์ธํธ (GET /items/{item_id}) | |
| def read_item(q: str | None = None): | |
| """ | |
| ํน์ ID์ ์์ดํ ์ ๋ณด๋ฅผ ๊ฐ์ ธ์ต๋๋ค. | |
| :param item_id: ์์ดํ ์ ๊ณ ์ ID (๊ฒฝ๋ก ๋งค๊ฐ๋ณ์) | |
| :param q: ์ ํ์ ์ธ ๊ฒ์ ๋ฌธ์์ด (์ฟผ๋ฆฌ ๋งค๊ฐ๋ณ์) | |
| """ | |
| return {"q": q, "description": "This is a query test."} | |
| # 5. ์์ฒญ ๋ณธ๋ฌธ(Body)์ ๋ฐ๋ ์๋ํฌ์ธํธ (POST /items/) | |
| def create_item(item: Item): | |
| """ | |
| ์๋ก์ด ์์ดํ ์ ์์ฑํ๊ณ ์ ๋ณด๋ฅผ ๋ฐํํฉ๋๋ค. | |
| :param item: Item Pydantic ๋ชจ๋ธ์ ์ ์๋ ๋ฐ์ดํฐ ๊ตฌ์กฐ | |
| """ | |
| # ์ค์ ๋ฐ์ดํฐ๋ฒ ์ด์ค์ ์ ์ฅํ๋ ๋์ , ๊ฐ๋จํ ์ฒ๋ฆฌ๋ฅผ ์ํ | |
| if item.price > 100.0: | |
| item.name = f"Premium {item.name}" | |
| return {"message": "Item created successfully", "item_data": item} | |
| if __name__ == "__main__": | |
| # --reload ์ต์ ์ ์ถ๊ฐํ์ฌ ์ฝ๋๊ฐ ๋ณ๊ฒฝ๋ ๋๋ง๋ค ์๋ ์ฌ์์๋๊ฒ ์ค์ ํฉ๋๋ค. | |
| uvicorn.run("app:app", host="0.0.0.0", port=8000, reload=True) |