"""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)