포스트

반려동물 품종 분류기 만들기 (7) 모델 export와 FastAPI 서빙

학습한 모델을 TorchScript로 export하고 FastAPI로 이미지 업로드 예측 API를 만듭니다. 서빙에서 가장 흔한 사고인 전처리 불일치를 확인하고 시리즈를 마무리합니다.

반려동물 품종 분류기 만들기 (7) 모델 export와 FastAPI 서빙

TorchScript로 export

학습 코드 없이 모델을 싣기 위해 TorchScript로 변환합니다. state_dict만 저장하면 로드하는 쪽에 모델 클래스 정의가 필요하지만, TorchScript는 구조와 가중치가 한 파일에 들어가서 서빙 코드가 torchvision 없이도 돌아갑니다.

1
2
3
4
model.eval()
example = torch.randn(1, 3, 224, 224)
traced = torch.jit.trace(model, example)
traced.save("outputs/finetune/model_ts.pt")

model.eval()을 반드시 trace 전에 호출해야 합니다. PyTorch 기초 5편에서 다뤘듯 BatchNorm과 Dropout은 train과 eval 모드의 동작이 다른데, trace는 호출 시점의 모드를 그대로 굳힙니다. train 모드로 trace하면 배포된 모델이 배치 통계를 계속 갱신하는 잘못된 상태가 됩니다.

전처리 불일치라는 함정

서빙에서 가장 흔한 사고는 모델이 아니라 전처리에서 납니다. 학습 때 적용한 eval 전처리(Resize 256, CenterCrop 224, ImageNet normalize)를 API 쪽에서 하나라도 빠뜨리면 정확도가 조용히 무너집니다.

실제로 확인해봤습니다. normalize를 빼고 test 셋을 다시 평가하면 accuracy가 92.4%에서 8.1%로 떨어집니다. 에러는 한 줄도 나지 않습니다. 입력이 모델이 학습한 분포와 다를 뿐이라서, 코드는 멀쩡히 돌고 결과만 틀립니다. 그래서 전처리를 API 코드에 다시 작성하지 않고, 학습 코드의 eval_tf를 그대로 import해서 씁니다.

FastAPI 예측 API

이미지를 업로드받아 top-3 예측을 돌려주는 엔드포인트입니다.

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
import io
import torch
from fastapi import FastAPI, UploadFile
from PIL import Image

from src.data import eval_tf, CLASSES

app = FastAPI()
model = torch.jit.load("outputs/finetune/model_ts.pt")
model.eval()

@app.post("/predict")
async def predict(file: UploadFile):
    image = Image.open(io.BytesIO(await file.read())).convert("RGB")
    x = eval_tf(image).unsqueeze(0)
    with torch.no_grad():
        probs = torch.softmax(model(x), dim=1)[0]
    top = torch.topk(probs, k=3)
    return {
        "predictions": [
            {"breed": CLASSES[i], "prob": round(p.item(), 4)}
            for p, i in zip(top.values, top.indices)
        ]
    }
  • convert("RGB")는 PNG의 알파 채널이나 흑백 이미지가 들어와도 3채널로 통일합니다
  • 추론은 torch.no_grad() 안에서 합니다. gradient 기록이 없어 메모리와 시간이 줄어듭니다
  • unsqueeze(0)으로 배치 차원을 추가합니다. 모델은 항상 (N, 3, 224, 224)를 받습니다

실행과 테스트는 두 줄입니다.

1
2
uv run uvicorn src.serve:app --port 8000
curl -X POST -F "file=@cat.jpg" http://localhost:8000/predict
1
2
3
{"predictions": [{"breed": "Russian Blue", "prob": 0.8734},
                 {"breed": "British Shorthair", "prob": 0.0912},
                 {"breed": "Korat", "prob": 0.0119}]}

M1 Mac의 CPU에서 요청당 응답 시간은 이미지 디코딩과 전처리를 포함해 약 60ms입니다. 이 규모의 모델은 GPU 없이도 실시간 서빙이 가능합니다.

시리즈를 마치며

일곱 편의 결과를 요약하면 다음과 같습니다.

  • 3천 장 규모의 fine-grained 분류에서 scratch CNN은 41.3%가 한계였습니다
  • 사전학습 ResNet-34의 분류층만 학습해도 86.7%, 전체를 fine-tuning하면 93.1%(test 92.4%)에 도달했습니다
  • 오답은 무작위가 아니라 실제로 닮은 품종 쌍에 몰려 있었고, Grad-CAM으로 모델이 동물의 얼굴을 근거로 판단한다는 것을 확인했습니다
  • TorchScript와 FastAPI로 모델을 학습 코드에서 분리해 API로 만들었습니다

이후 확장으로는 Docker 이미지로 묶어 배포 환경을 고정하는 것, ONNX Runtime으로 추론을 더 줄이는 것, 그리고 NYC 택시 프로젝트에서 다룬 모니터링과 재학습 루프를 이 프로젝트에 붙이는 것이 남아 있습니다.

이 기사는 저작권자의 CC BY 4.0 라이센스를 따릅니다.