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 = "" @router.get("/", response_model=APIResponse) async def health_check(): return { "success": True, "data": None, "msg": "cnn_router is active" } @router.post("/extract_features", response_model=APIResponse) 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)}" }