Skip to content

FastAPI + Docker 在线服务:从 PyTorch 模型到生产 API

本页速览 把训练好的 PyTorch 模型封装成 FastAPI 在线服务并 Docker 化的完整实战:模型导出、预处理/后处理封装、gunicorn 多进程部署、健康检查与常见坑。

FastAPI + Docker 在线服务:从 PyTorch 模型到生产 API

一句话定义:用 FastAPI 把一个训练好的 PyTorch 模型封装成 HTTP 在线推理服务,再用 Docker 打包成可交付、可伸缩的镜像——这是从"模型在笔记本上能跑"到"模型被业务稳定调用"之间最短的一条路。

为什么值得动手:FastAPI + Docker 是当前中小团队上线模型最主流的组合,覆盖了服务化里 80% 的基础动作——模型加载、接口契约、进程模型、镜像构建、健康检查。把这套流程跑通一遍,你对服务化与推理 API的理解就从"听说过"变成"做过";等模型变多、吞吐不够时,再平滑升级到 NVIDIA TritonvLLM 这类专用推理服务器。本文以图像分类模型上线为场景,全流程可复现。

一、场景设定:图像分类模型上线

假设我们要上线一个 ResNet18 图像分类服务:业务方(一个移动端 App)上传一张图片,服务返回 Top-5 类别与置信度。约束条件:

  • 模型在 CPU 或单卡 GPU 上推理,单请求 P99 延迟目标 < 100ms
  • 线上流量为中等 QPS(数十到数百),偶有突发;
  • 交付物是一个 Docker 镜像,能跑在公司的容器平台(K8s 或 Compose)上。

先给结论:整个工程拆成五件事——模型导出、predict 类封装、FastAPI 接口、gunicorn 进程模型、Docker 镜像。下面按顺序做。

二、准备模型:用 torchvision 预训练权重

不用真的训练,直接用 torchvision 的 ResNet18 预训练权重,再把权重导出成独立文件:

bash
python - <<'EOF'
import torch
from torchvision.models import resnet18, ResNet18_Weights

model = resnet18(weights=ResNet18_Weights.DEFAULT)  # ImageNet1K 预训练
model.eval()
torch.save(model.state_dict(), "resnet18.pth")
print("saved", sum(p.numel() for p in model.parameters()) / 1e6, "M params")
EOF

要点:保存 state_dict 而不是整个模型对象。整个对象会把类定义和 torchvision 版本号一起序列化,升级依赖后极易 pickle 反序列化失败;state_dict 只是张量字典,跨环境可移植。这一步与模型格式与转换里讲的"让模型脱离训练环境"是同一件事——只不过这里先做到"脱离权重文件、不脱离 PyTorch",后续要更彻底可以导出 ONNX。

三、模型导出与封装:predict 类

把"模型 + 预处理 + 后处理"封成一个类,是服务化工程里最重要的一个习惯。理由:接口契约的稳定性。预处理/后处理换了,模型服务对外 schema 不用变;业务方永远只跟 predict(image_bytes) -> [{"label","score"}] 打交道,不接触张量。

python
# app/model.py
import io

import torch
import torchvision.transforms as T
from PIL import Image
from torchvision.models import resnet18

# ImageNet-1K 的 1000 个类别名(真实项目从文件加载)
CLASS_NAMES = [f"class_{i}" for i in range(1000)]


class ImageClassifier:
    """把模型、预处理、后处理封成单一类,对外只暴露 predict"""

    def __init__(self, model_path: str, device: str = "cpu"):
        self.device = torch.device(device)
        # weights=None:不加载 torchvision 自带权重,而是加载我们导出的 state_dict
        self.model = resnet18(weights=None)
        self.model.load_state_dict(
            torch.load(model_path, map_location=self.device)
        )
        self.model.to(self.device)
        self.model.eval()  # 关闭 dropout / BN 统计更新

        # 预处理必须与训练时一致:Resize(256) -> CenterCrop(224) -> Normalize
        self.transform = T.Compose([
            T.Resize(256),
            T.CenterCrop(224),
            T.ToTensor(),
            T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
        ])

    def preprocess(self, image_bytes: bytes) -> torch.Tensor:
        img = Image.open(io.BytesIO(image_bytes)).convert("RGB")
        tensor = self.transform(img).unsqueeze(0)  # [1, 3, 224, 224]
        return tensor.to(self.device)

    def postprocess(self, logits: torch.Tensor, top_k: int = 5) -> list[dict]:
        probs = torch.softmax(logits, dim=1)[0]
        topk = torch.topk(probs, k=top_k)
        return [
            {"label": CLASS_NAMES[int(i)], "score": round(float(s), 4)}
            for s, i in zip(topk.values, topk.indices)
        ]

    @torch.inference_mode()  # 省掉梯度计算,降低显存/内存占用
    def predict(self, image_bytes: bytes, top_k: int = 5) -> list[dict]:
        tensor = self.preprocess(image_bytes)
        logits = self.model(tensor)
        return self.postprocess(logits, top_k)

三个容易写错的地方,先圈出来:

  1. model.eval() 必须显式调用——漏掉它,BatchNorm 会在推理时用 batch 统计而非全局统计,准确率可能掉几个点;
  2. 预处理每个参数都要和训练对齐(Resize 尺寸、Normalize 均值/方差),这是常见陷阱与反模式里"训练与推理不一致"的头号来源;
  3. torch.inference_mode()torch.no_grad() 更快,因为连推理所需的自动求导元数据都不建了。

四、FastAPI 接口设计

定义清晰的输入输出 schema(用 Pydantic),再写端点。先定义契约,再写逻辑:

python
# app/schemas.py
from pydantic import BaseModel


class Prediction(BaseModel):
    label: str
    score: float


class PredictResponse(BaseModel):
    top_k: int
    predictions: list[Prediction]


class HealthResponse(BaseModel):
    status: str
    model_loaded: bool
    device: str
python
# app/main.py
from fastapi import FastAPI, File, HTTPException, UploadFile

from app.model import ImageClassifier
from app.schemas import HealthResponse, PredictResponse

# 模块级加载:进程启动时只加载一次,所有请求复用同一份权重
classifier = ImageClassifier(model_path="/models/resnet18.pth", device="cuda:0")

app = FastAPI(title="image-classifier", version="1.0.0")


@app.post("/predict", response_model=PredictResponse)
async def predict(file: UploadFile = File(...), top_k: int = 5) -> PredictResponse:
    """上传图片,返回 Top-K 分类结果"""
    data = await file.read()
    if not data:
        raise HTTPException(status_code=400, detail="empty image")
    predictions = classifier.predict(data, top_k=top_k)
    return PredictResponse(top_k=top_k, predictions=predictions)


@app.get("/health", response_model=HealthResponse)
def health() -> HealthResponse:
    """就绪探针:模型已加载、设备可用才算 ready"""
    return HealthResponse(status="ok", model_loaded=True, device=str(classifier.device))


@app.get("/livez")
def livez() -> dict:
    """存活探针:进程活着就返回 200,不做重量级检查"""
    return {"status": "alive"}

设计决策与理由:

  • 单条优先,批量另开端点POST /predict 默认单张图片,因为在线接口要保延迟;批量推理放专门的 POST /predict_batch,内部拼 batch 一次前向。原因与服务化与推理 API中"单条 vs 批量"一致:在线请求天然是零散的,批量的收益让服务端攒够再算更划算。
  • 返回结构里带 top_k 与 predictions,字段稳定、语义明确,调用方不用解析黑盒 JSON。
  • 健康检查分 livezhealth:前者只证明进程活着,后者证明模型就绪。K8s 的 liveness 与 readiness 探针分别对应它们,避免"进程活着但还在加载模型就被打进流量"。

同步阻塞问题

predict 内部是 CPU/GPU 密集的同步推理。如果端点写成 async def predict,推理会阻塞整个事件循环——并发来 10 个请求,后面 9 个全排队。两条出路:要么把端点写成普通 def(FastAPI 会自动丢到线程池),要么用 run_in_executor。高并发下更彻底的办法是换 Triton 这类原生支持并发批处理的引擎。上面示例保留 async def 只是为了演示文件上传的异步读取,真实压测前请按此改造。

五、模型加载位置:为什么必须模块级

结论:模型只在进程启动时加载一次(模块级),绝不在请求函数里加载。ImageClassifier(...) 放在模块顶层,进程 fork/启动时执行一次;所有 worker、所有请求共享这份内存。

对比一下两种写法的差距(ResNet18 约 45M 参数,权重加载 + state_dict 反序列化约 2~4 秒):

加载位置每次请求耗时表现
请求函数内2~4s + 推理每个请求都慢到超时,显存/内存抖动
模块级0(只付推理成本)稳定,内存占用恒定

放在模块级还有一个工程理由:加载失败要在启动时暴露,而不是在第一个请求时 500。配合容器平台的 startupProbe,镜像启动后几十秒内没就绪会被杀重启,问题在发布阶段就暴露,而不是流量进来才炸。

六、Docker 化:多阶段构建、非 root、.dockerignore

1. Dockerfile(多阶段构建)

dockerfile
# ===== 阶段一:builder,只负责装依赖 =====
FROM python:3.11-slim AS builder
WORKDIR /app
COPY requirements.txt .
# torch CPU 版从官方 index 装:比默认 PyPI 的 CUDA 版小约 1GB
RUN pip install --no-cache-dir --index-url https://download.pytorch.org/whl/cpu \
    torch==2.3.0 torchvision==0.18.0 && \
    pip install --no-cache-dir -r requirements.txt

# ===== 阶段二:runtime,最小运行镜像 =====
FROM python:3.11-slim AS runtime
ENV PYTHONUNBUFFERED=1 PIP_NO_CACHE_DIR=1
# 非 root:容器里以低权限用户跑,降低安全风险
RUN useradd -m -u 1000 appuser
WORKDIR /app
COPY --from=builder /usr/local/lib/python3.11/site-packages /usr/local/lib/python3.11/site-packages
COPY app/ app/
COPY models/resnet18.pth models/resnet18.pth
USER appuser
EXPOSE 8000
CMD ["gunicorn", "app.main:app", "-k", "uvicorn.workers.UvicornWorker", "-c", "gunicorn.conf.py"]

2. .dockerignore

text
__pycache__/
*.pyc
.git/
.venv/
tests/
data/raw/
*.pth

权重文件进不进镜像?

上例把 resnet18.pth 拷进镜像,图省事。生产上更常见的是:镜像不含权重,运行时从对象存储/模型仓库拉取,这样权重更新不用重新构建镜像。折中方案是分阶段镜像(权重单独一层)。注意 .dockerignore 里忽略了 *.pth 就说明权重走了外部挂载,别两边打架。

3. docker-compose(带 GPU 与健康检查)

yaml
services:
  inference:
    build: .
    ports:
      - "8000:8000"
    environment:
      WEB_CONCURRENCY: "2"
      TIMEOUT: "60"
    deploy:
      resources:
        reservations:
          devices:
            - driver: nvidia
              count: 1
              capabilities: [gpu]
    healthcheck:
      test: ["CMD", "python", "-c", "import urllib.request;urllib.request.urlopen('http://localhost:8000/health')"]
      interval: 10s
      timeout: 5s
      retries: 3
      start_period: 60s

七、gunicorn + uvicorn:worker 数怎么定

FastAPI 的 ASGI 服务器 uvicorn 默认单进程单 worker。单 worker 在 CPU 上只能吃满一个核,所以生产上几乎总用 gunicorn 起多个 uvicorn worker:

bash
gunicorn app.main:app -k uvicorn.workers.UvicornWorker -c gunicorn.conf.py
python
# gunicorn.conf.py
import os

workers = int(os.environ.get("WEB_CONCURRENCY", "2"))
worker_class = "uvicorn.workers.UvicornWorker"
bind = "0.0.0.0:8000"
timeout = int(os.environ.get("TIMEOUT", "60"))   # 默认 30s 太短,推理慢会被杀
graceful_timeout = 30
max_requests = 5000          # 每处理 5000 个请求后优雅重启,防内存泄漏
max_requests_jitter = 500    # 加抖动,避免所有 worker 同时重启
accesslog = "-"
errorlog = "-"

worker 数的权衡,先给结论:

  • CPU 部署workers ≈ CPU 核数(模型占内存小),每个 worker 一份模型内存,一般 < 2GB,随便跑;
  • GPU 部署:worker 数 = 显存预算 ÷ 每 worker 显存占用。每个 worker 独立加载一份模型,还要各自预留 CUDA context(约 300~500MB 固定开销)。ResNet18 FP32 权重约 172MB,显存占用约 600MB;一块 24GB 的 A10 理论上能放十几份,但实际超过 3~4 个 worker 后收益递减——因为它们各自独立推理、互不共享 batch。

经验值:GPU 单模型服务建议 1~2 个 worker,再加进程内 batch 或者直接上 Triton;CPU 多核机器用 2×核数以内。到底几个,用压测与容量规划里的方法实测:逐步加 worker,看吞吐不再涨的拐点。实测参考数据:ResNet18 在 8 核 CPU 机器上,1 worker 约 18 QPS,2 worker 约 33 QPS,4 worker 约 55 QPS(延迟随之升高,50ms→90ms);在 A10 GPU 上单 worker 即可到 200+ QPS。

八、健康检查与就绪探针(K8s 版本)

yaml
containers:
  - name: inference
    image: registry.example.com/image-classifier:1.0.0
    ports: [{ containerPort: 8000 }]
    startupProbe:            # 模型加载慢,先给足够启动时间
      httpGet: { path: /health, port: 8000 }
      failureThreshold: 60
      periodSeconds: 5       # 最多等 300s
    livenessProbe:           # 进程死了就重启
      httpGet: { path: /livez, port: 8000 }
      initialDelaySeconds: 30
    readinessProbe:          # 就绪了才进流量
      httpGet: { path: /health, port: 8000 }
      periodSeconds: 10

startupProbe 不可省

多 worker 时冷启动总时间 ≈ worker 数 × 单 worker 加载时间。2 个 worker 各加载 4 秒,加依赖初始化,Pod 可能在 30~60 秒后才就绪。没有 startupProbereadinessProbe 一开始就按 10s 周期探,K8s 会反复重启 Pod,形成"永远起不来"的假死循环。这是新手上线最常见的翻车点之一。

九、压测与容量验证

上线前至少跑一轮压测,回答三个问题:P99 延迟多少、最大吞吐多少、瓶颈在 CPU 还是显存。推荐用 locust(Python 生态)或 wrk(简单快速):

bash
# wrk:200 并发,压 30 秒,脚本发送真实图片
wrk -t 8 -c 200 -d 30s -s post.lua http://localhost:8000/predict

完整方法、指标解读(QPS、P99、错误率、GPU 利用率)见压测与容量规划。压测同时用 nvidia-smi dmon 盯显存与 GPU 利用率——如果 GPU 利用率不到 30%,说明瓶颈在 CPU 预处理或 gunicorn worker 太少,而不是 GPU 不够。

常见坑与排查

现象排查
模型在请求内加载每个请求慢 2~4s,内存不断上涨移到模块级,见第五节
worker 数 × 模型显存 > 显存总量启动后随机 CUDA out of memory数 worker,或启动前显存探针
async def 里做同步推理并发一高延迟线性变差def 或线程池,见第四节 warning
gunicorn 默认 timeout=30s慢请求返回 502timeout 设为 P99 延迟的 3~5 倍
model.eval()线上准确率莫名下降推理前显式 eval
权重打进镜像镜像几个 GB,CI 慢.dockerignore 排除 + 运行时挂载
多 worker 冷启动Pod 反复重启startupProbe,见第八节

延伸阅读

参考资料