Download cnn_router.py from WildOjisan/embeddinggemma_300m_fastapi: direct link, hf CLI and curl.
- Browser
- Download file 3.38 kB
-
https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/main/cnn_router.py
- Command line
-
hf download hf://spaces/WildOjisan/embeddinggemma_300m_fastapi/cnn_router.py
-
curl -L -o cnn_router.py https://huggingface.co/spaces/WildOjisan/embeddinggemma_300m_fastapi/resolve/main/cnn_router.py
3.38 kB
| from fastapi import APIRouter, File, UploadFile | |
| from pydantic import BaseModel | |
| from typing import Any, Optional, List | |
| import io | |
| import torch | |
| import torchvision.models as models | |
| import torchvision.transforms as transforms | |
| import torch.nn.functional as F | |
| from PIL import Image | |
| router = APIRouter( | |
| tags=["cnn"], | |
| responses={404: {"description": "Not found"}}, | |
| ) | |
| # --- Global Vision Model --- | |
| vision_model = None | |
| preprocess = None | |
| def init_vision_model(): | |
| global vision_model, preprocess | |
| try: | |
| print("Loading EfficientNet-B1 model...") | |
| # 1. ๋ชจ๋ธ ๋ก๋ | |
| # weights='IMAGENET1K_V1' ๋ฑ์ผ๋ก ์ฌ์ ํ์ต๋ ๊ฐ์ค์น ์ฌ์ฉ | |
| vision_model = models.efficientnet_b1(weights='IMAGENET1K_V1') | |
| # 2. ๋ง์ง๋ง ๋ถ๋ฅ ๋ ์ด์ด(classifier) ์ ๊ฑฐ | |
| vision_model.classifier = torch.nn.Identity() | |
| vision_model.eval() # ์ถ๋ก ๋ชจ๋๋ก ๋ณ๊ฒฝ | |
| # 3. ์ด๋ฏธ์ง ์ ์ฒ๋ฆฌ | |
| preprocess = transforms.Compose([ | |
| transforms.Resize(255), | |
| transforms.CenterCrop(240), # B1 ๊ถ์ฅ ์ฌ์ด์ฆ | |
| transforms.ToTensor(), | |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), | |
| ]) | |
| print("EfficientNet-B1 model loaded successfully.") | |
| except Exception as e: | |
| print(f"Error loading Vision model: {e}") | |
| vision_model = None | |
| preprocess = None | |
| class APIResponse(BaseModel): | |
| success: bool | |
| data: Optional[Any] = None | |
| msg: str = "" | |
| async def health_check(): | |
| return { | |
| "success": True, | |
| "data": None, | |
| "msg": "cnn_router is active" | |
| } | |
| async def extract_features(files: List[UploadFile] = File(...)): | |
| result_data = [] | |
| if vision_model is None or preprocess is None: | |
| return { | |
| "success": False, | |
| "data": None, | |
| "msg": "Vision model is not loaded." | |
| } | |
| try: | |
| for file in files: | |
| # Read file content | |
| contents = await file.read() | |
| img = Image.open(io.BytesIO(contents)).convert('RGB') | |
| # Preprocess | |
| input_tensor = preprocess(img).unsqueeze(0) | |
| # Inference | |
| with torch.no_grad(): | |
| raw_embedding = vision_model(input_tensor) | |
| # [์ถ๊ฐ ์์ ] L2 ์ ๊ทํ (Normalize) | |
| # p=2๋ L2 norm์ ์๋ฏธ, dim=1์ ์ฐจ์ ๋ฐฉํฅ | |
| normalized_embedding = F.normalize(raw_embedding, p=2, dim=1) | |
| # Convert to list | |
| embedding_list = normalized_embedding.squeeze().tolist() | |
| result_data.append({ | |
| "key": file.filename, | |
| "embedding": embedding_list | |
| }) | |
| msg = f"Processed {len(files)} images." | |
| if result_data: | |
| embedding_dim = len(result_data[0]["embedding"]) | |
| msg += f" \n Embedding Shape: (1, {embedding_dim})" | |
| return { | |
| "success": True, | |
| "data": result_data, | |
| "msg": msg | |
| } | |
| except Exception as e: | |
| return { | |
| "success": False, | |
| "data": None, | |
| "msg": f"Error processing images: {str(e)}" | |
| } | |