外观
FastAPI + Docker 在线服务:从 PyTorch 模型到生产 API
一句话定义:用 FastAPI 把一个训练好的 PyTorch 模型封装成 HTTP 在线推理服务,再用 Docker 打包成可交付、可伸缩的镜像——这是从"模型在笔记本上能跑"到"模型被业务稳定调用"之间最短的一条路。
为什么值得动手:FastAPI + Docker 是当前中小团队上线模型最主流的组合,覆盖了服务化里 80% 的基础动作——模型加载、接口契约、进程模型、镜像构建、健康检查。把这套流程跑通一遍,你对服务化与推理 API的理解就从"听说过"变成"做过";等模型变多、吞吐不够时,再平滑升级到 NVIDIA Triton 或 vLLM 这类专用推理服务器。本文以图像分类模型上线为场景,全流程可复现。
一、场景设定:图像分类模型上线
假设我们要上线一个 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)三个容易写错的地方,先圈出来:
model.eval()必须显式调用——漏掉它,BatchNorm 会在推理时用 batch 统计而非全局统计,准确率可能掉几个点;- 预处理每个参数都要和训练对齐(Resize 尺寸、Normalize 均值/方差),这是常见陷阱与反模式里"训练与推理不一致"的头号来源;
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: strpython
# 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。
- 健康检查分
livez与health:前者只证明进程活着,后者证明模型就绪。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.pypython
# 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: 10startupProbe 不可省
多 worker 时冷启动总时间 ≈ worker 数 × 单 worker 加载时间。2 个 worker 各加载 4 秒,加依赖初始化,Pod 可能在 30~60 秒后才就绪。没有 startupProbe,readinessProbe 一开始就按 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 | 慢请求返回 502 | timeout 设为 P99 延迟的 3~5 倍 |
漏 model.eval() | 线上准确率莫名下降 | 推理前显式 eval |
| 权重打进镜像 | 镜像几个 GB,CI 慢 | .dockerignore 排除 + 运行时挂载 |
| 多 worker 冷启动 | Pod 反复重启 | 加 startupProbe,见第八节 |
延伸阅读
- 服务化与推理 API —— 接口设计、单条 vs 批量、协议选择的理论依据,本文是其落地案例
- 压测与容量规划 —— 把本文的 worker 数、QPS 目标用压测数据闭环验证
- 模型格式与转换 —— 下一阶段的优化起点:PyTorch → ONNX → TensorRT
- 常见陷阱与反模式 —— 服务化踩坑清单的完整版,本文只列了高频项
- NVIDIA Triton 多模型服务 —— 单模型 FastAPI 扛不住高吞吐/多模型时的升级路径
- 推理系统总体架构解剖 —— 把本文的单个服务放回完整推理系统里看它的位置
参考资料
- FastAPI 官方文档:https://fastapi.tiangolo.com/
- Uvicorn 官方文档:https://www.uvicorn.org/
- Gunicorn 官方文档:https://docs.gunicorn.org/en/stable/
- PyTorch 官方文档:https://pytorch.org/docs/stable/
- Docker 官方文档:https://docs.docker.com/