Spaces:
Sleeping
Sleeping
| """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} | |
| # ---- 수명주기: 시작 시 모델 로드 + 워밍업 ---------------------------------- | |
| 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=["*"], | |
| ) | |
| # ---- 엔드포인트 ------------------------------------------------------------ | |
| def root(): | |
| """루트 안내 (이 서버는 JSON API). 브라우저 테스트는 /docs 사용.""" | |
| return { | |
| "service": "P3 Fashion Segmentation API", | |
| "endpoints": {"health": "GET /health", "predict": "POST /predict (multipart 'file')"}, | |
| "interactive_docs": "/docs", | |
| } | |
| 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"], | |
| } | |
| 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) | |