WildOjisan's picture
.
cee7d2b
Raw History Blame Contribute Delete
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 = ""
@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)}"
}