yeonjung77's picture
Upload folder using huggingface_hub
3c5e54b verified
Raw
History Blame Contribute Delete
5.39 kB
"""P3 추론 서버 -- FastAPI (Phase 5 배포)
학습한 2-Stage 모델을 HTTP API 로 제공한다. 서버 시작 시 Pipeline 을 한 번 로드하고
워밍업 추론을 돌려, 이후 요청은 cold start 없이 빠르게 처리한다.
엔드포인트:
GET /health 상태 + 모델 로드 여부
POST /predict 이미지 파일(multipart) → garment 트리 JSON (postprocess 스키마)
[실행]
cd serving
uvicorn app:app --host 0.0.0.0 --port 8000
# 또는: python app.py
[환경변수] (없으면 프로젝트 기본 경로 사용)
P3_STAGE1_CKPT, P3_STAGE2_CKPT, P3_ANNO_PATH, P3_DEVICE
"""
from __future__ import annotations
import io
import os
import sys
import logging
import traceback
from contextlib import asynccontextmanager
from fastapi import FastAPI, File, UploadFile, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from PIL import Image
# 상위 폴더(pipeline.py 위치)를 import 경로에 추가
PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, PROJECT_DIR)
from pipeline import Pipeline # noqa: E402
# ---- 설정 ------------------------------------------------------------------
DEFAULT_DATA_DIR = "/home/ellin/github/fashionproject/003_datasets/P3_fashionpedia"
STAGE1_CKPT = os.environ.get("P3_STAGE1_CKPT", os.path.join(PROJECT_DIR, "outputs_stage1", "stage1_best.pth"))
STAGE2_CKPT = os.environ.get("P3_STAGE2_CKPT", os.path.join(PROJECT_DIR, "outputs_stage2", "stage2_best.pth"))
ANNO_PATH = os.environ.get("P3_ANNO_PATH", os.path.join(DEFAULT_DATA_DIR, "instances_attributes_train2020.json"))
DEVICE = os.environ.get("P3_DEVICE", "cuda")
MAX_FILE_BYTES = 10 * 1024 * 1024 # 10MB
ALLOWED_ORIGINS = [
"https://yeonjung77.github.io",
"http://localhost:7860", "http://localhost:8000", "http://localhost:3000",
]
logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(levelname)s - %(message)s")
logger = logging.getLogger("p3-serving")
# 전역 상태 (서버 1개당 Pipeline 1개)
STATE: dict = {"pipeline": None, "error": None}
# ---- 수명주기: 시작 시 모델 로드 + 워밍업 ----------------------------------
@asynccontextmanager
async def lifespan(app: FastAPI):
"""서버 시작 시 Pipeline 로드 + 워밍업, 종료 시 정리."""
try:
import torch
device = DEVICE if (DEVICE != "cuda" or torch.cuda.is_available()) else "cpu"
logger.info(f"Pipeline 로딩... (device={device})")
pipe = Pipeline(STAGE1_CKPT, STAGE2_CKPT, device=device, anno_path=ANNO_PATH)
pipe.warmup() # 두 stage 모두 데움 (cold start 제거)
STATE["pipeline"] = pipe
logger.info("Pipeline 준비 완료 (워밍업 포함)")
except Exception as e: # 로드 실패해도 서버는 떠서 /health 로 알림
STATE["error"] = str(e)
logger.error(f"Pipeline 로드 실패: {e}\n{traceback.format_exc()}")
yield
STATE["pipeline"] = None
app = FastAPI(title="P3 Fashion Segmentation API", version="1.0", lifespan=lifespan)
app.add_middleware(
CORSMiddleware, allow_origins=ALLOWED_ORIGINS,
allow_methods=["GET", "POST"], allow_headers=["*"],
)
# ---- 엔드포인트 ------------------------------------------------------------
@app.get("/")
def root():
"""루트 안내 (이 서버는 JSON API). 브라우저 테스트는 /docs 사용."""
return {
"service": "P3 Fashion Segmentation API",
"endpoints": {"health": "GET /health", "predict": "POST /predict (multipart 'file')"},
"interactive_docs": "/docs",
}
@app.get("/health")
def health():
"""서버/모델 상태. 모델 로드 성공 시 stage1/stage2 모두 True."""
loaded = STATE["pipeline"] is not None
return {
"status": "ok" if loaded else "degraded",
"stage1_loaded": loaded,
"stage2_loaded": loaded,
"error": STATE["error"],
}
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
"""이미지 1장 → garment 트리 JSON. (multipart 업로드)"""
pipe = STATE["pipeline"]
if pipe is None:
raise HTTPException(status_code=503, detail="모델이 아직 로드되지 않았습니다.")
# 1) 크기 검증
content = await file.read()
if len(content) > MAX_FILE_BYTES:
raise HTTPException(status_code=413, detail=f"파일이 너무 큽니다 (최대 {MAX_FILE_BYTES // (1024*1024)}MB).")
if not content:
raise HTTPException(status_code=400, detail="빈 파일입니다.")
# 2) 이미지 형식 검증 (PIL 로 실제 디코딩)
try:
image = Image.open(io.BytesIO(content)).convert("RGB")
except Exception:
raise HTTPException(status_code=400, detail="유효한 이미지 파일이 아닙니다.")
# 3) 추론 (내부 오류는 stack trace 로깅 + 사용자에겐 간단 메시지)
try:
return JSONResponse(content=pipe.predict(image))
except HTTPException:
raise
except Exception:
logger.error(f"추론 실패 (file={file.filename}):\n{traceback.format_exc()}")
raise HTTPException(status_code=500, detail="추론 중 오류가 발생했습니다.")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)