Machine Learning: 模型部署FastAPI与Docker — ML模型生产服务化指南
训练好的模型躺在Jupyter里一文不值——只有部署到生产,价值才真正释放。
1. 你将学到
- FastAPI模型服务:Pydantic数据校验、异步预测端点、Swagger文档自动生成
- 模型序列化:joblib/pickle保存sklearn模型,torch.save保存PyTorch模型
- Docker容器化:Dockerfile编写、多阶段构建、镜像优化
- docker-compose编排:模型服务 + Redis缓存 + Nginx负载均衡
- Bob的预测API:POST /predict接口设计,单次预测<50ms延迟
2. 一个ML工程师的真实故事
(1) 痛点:模型在Notebook里无法给业务用
Bob训练了一个R²=0.89的XGBoost模型,但它只在Jupyter Notebook里能跑。产品经理问"能不能给前端调用?"Bob不知道怎么把模型变成API。训练和部署之间的鸿沟,是ML项目最大的"最后一公里"问题。
(2) FastAPI+Docker的解法
FastAPI将模型包装成REST API,Docker将其容器化——任何服务都能通过HTTP调用预测。
PYTHON
from fastapi import FastAPI
import joblib
app = FastAPI()
model = joblib.load("model.pkl")
@app.post("/predict")
def predict(features: PredictionInput):
result = model.predict([features.dict()])
return {"prediction": float(result[0])}
(3) 收益:API上线后日处理百万请求
Bob用FastAPI+Docker部署模型后,API延迟<50ms,日处理1 million+请求,前端/CRM/推荐系统都能调用。
3. 模型序列化
(1) 保存与加载模型
▶ 示例:sklearn模型序列化
PYTHON
import joblib
import pickle
from sklearn.ensemble import RandomForestRegressor
from sklearn.preprocessing import StandardScaler
from sklearn.pipeline import Pipeline
import numpy as np
# Train and save model
rng = np.random.default_rng(42)
X = rng.uniform(0, 100, (1000, 5))
y = 50 + 0.8 * X[:, 0] + 1.2 * X[:, 1] + rng.normal(0, 5, 1000)
pipe = Pipeline([
("scaler", StandardScaler()),
("model", RandomForestRegressor(n_estimators=100, random_state=42)),
])
pipe.fit(X, y)
# Save with joblib (recommended for sklearn)
joblib.dump(pipe, "salespredict_model.joblib", compress=3)
# Save with pickle (alternative)
with open("salespredict_model.pkl", "wb") as f:
pickle.dump(pipe, f)
# Load and predict
loaded_model = joblib.load("salespredict_model.joblib")
sample = np.array([[50, 30, 20, 10, 5]])
prediction = loaded_model.predict(sample)
print(f"Prediction: {prediction[0]:.2f} thousand USD")
输出:
TEXT
📖 仅展示
# 执行成功
| 方法 | 适用模型 | 优点 | 缺点 |
|---|---|---|---|
| joblib | sklearn/numpy | 大数组高效 | 仅Python |
| pickle | 任意Python对象 | 通用 | 安全风险、版本兼容 |
| torch.save | PyTorch | 灵活(可存state_dict) | 仅PyTorch |
| mlflow.sklearn | sklearn | 版本+元数据 | 需MLflow |
| ONNX | 跨框架 | 跨语言/平台 | 转换复杂 |
▶ 示例:PyTorch模型保存
PYTHON
import torch
import torch.nn as nn
# Save model state_dict (recommended)
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Linear(5, 32), nn.ReLU(), nn.Linear(32, 1))
def forward(self, x):
return self.net(x)
model = SimpleModel()
torch.save(model.state_dict(), "pytorch_model.pt")
# Load
loaded = SimpleModel()
loaded.load_state_dict(torch.load("pytorch_model.pt", weights_only=True))
loaded.eval()
输出:
TEXT
📖 仅展示
# 函数定义成功
4. FastAPI模型服务
(1) FastAPI基础
▶ 示例:完整的预测API
PYTHON
# File: app.py
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field
import joblib
import numpy as np
import time
app = FastAPI(title="SalesPredict API", version="1.0.0")
# Load model at startup
model = joblib.load("salespredict_model.joblib")
class PredictionInput(BaseModel):
ad_spend_k: float = Field(..., ge=0, description="Ad spend in thousand USD")
traffic_k: float = Field(..., ge=0, description="Traffic in thousands")
category_electronics: float = Field(0, ge=0, le=1)
category_clothing: float = Field(0, ge=0, le=1)
is_promotion: float = Field(0, ge=0, le=1)
model_config = {"json_schema_extra": {
"example": {"ad_spend_k": 50, "traffic_k": 300,
"category_electronics": 1, "category_clothing": 0, "is_promotion": 1}
}}
class PredictionOutput(BaseModel):
predicted_revenue_k: float
latency_ms: float
@app.get("/health")
def health_check():
return {"status": "healthy", "model_loaded": model is not None}
@app.post("/predict", response_model=PredictionOutput)
def predict(input_data: PredictionInput):
start = time.time()
try:
features = np.array([[input_data.ad_spend_k, input_data.traffic_k,
input_data.category_electronics,
input_data.category_clothing, input_data.is_promotion]])
prediction = model.predict(features)[0]
latency = (time.time() - start) * 1000
return PredictionOutput(predicted_revenue_k=round(float(prediction), 2),
latency_ms=round(latency, 2))
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/predict_batch")
def predict_batch(inputs: list[PredictionInput]):
features = np.array([[d.ad_spend_k, d.traffic_k, d.category_electronics,
d.category_clothing, d.is_promotion] for d in inputs])
predictions = model.predict(features)
return {"predictions": [round(float(p), 2) for p in predictions]}
输出:
TEXT
📖 仅展示
# 函数定义成功
(2) 运行FastAPI服务
BASH
# Install dependencies
pip install fastapi uvicorn joblib scikit-learn
# Run the API server
uvicorn app:app --host 0.0.0.0 --port 8000 --reload
# Test with curl
curl -X POST http://localhost:8000/predict \
-H "Content-Type: application/json" \
-d '{"ad_spend_k": 50, "traffic_k": 300, "category_electronics": 1, "category_clothing": 0, "is_promotion": 1}'
# Access Swagger UI: http://localhost:8000/docs
| FastAPI特性 | 说明 |
|---|---|
| Pydantic校验 | 自动验证输入类型和范围 |
| Swagger UI | 自动生成交互式文档(/docs) |
| 类型提示 | 自动生成响应模型 |
| 异步支持 | async/await高并发 |
| 异常处理 | HTTPException标准错误码 |
5. Docker容器化
(1) Dockerfile编写
▶ 示例:SalesPredict Docker镜像
DOCKERFILE
# Stage 1: Build dependencies
FROM python:3.11-slim AS builder
WORKDIR /build
COPY requirements.txt .
RUN pip install --no-cache-dir --prefix=/install -r requirements.txt
# Stage 2: Runtime (smaller image)
FROM python:3.11-slim
WORKDIR /app
# Copy installed packages from builder
COPY --from=builder /install /usr/local
# Copy application code and model
COPY app.py .
COPY salespredict_model.joblib .
# Non-root user for security
RUN useradd -m appuser
USER appuser
EXPOSE 8000
HEALTHCHECK --interval=30s --timeout=5s \
CMD curl -f http://localhost:8000/health || exit 1
CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "2"]
TEXT
📖 仅展示
# requirements.txt
fastapi==0.109.0
uvicorn==0.27.0
joblib==1.3.2
scikit-learn==1.4.0
numpy==1.26.4
pydantic==2.5.0
(2) Docker命令
BASH
# Build image
docker build -t salespredict-api:latest .
# Run container
docker run -d -p 8000:8000 --name salespredict salespredict-api:latest
# Test
curl http://localhost:8000/health
# View logs
docker logs salespredict
# Stop and remove
docker stop salespredict && docker rm salespredict
| Dockerfile优化 | 效果 |
|---|---|
| 多阶段构建 | 镜像从1.5GB降到200MB |
| slim基础镜像 | 减少不必要的系统包 |
| .dockerignore | 排除.git/data等大文件 |
| 非root用户 | 安全加固 |
| HEALTHCHECK | 容器健康检查 |
6. docker-compose编排
完整部署架构将所有组件编排在一起——API服务、缓存、负载均衡、监控形成一条完整链路:
graph TB
CLIENT[Client / Frontend] --> NGINX[Nginx<br/>Rate Limit + LB]
NGINX --> API1[FastAPI Worker 1]
NGINX --> API2[FastAPI Worker 2]
API1 --> REDIS[(Redis Cache<br/>LRU 256MB)]
API2 --> REDIS
API1 --> MODEL[Model File<br/>.joblib]
API2 --> MODEL
PROM[Prometheus<br/>Metrics] --> API1
PROM --> API2
GRAF[Grafana<br/>Dashboard] --> PROM
▶ 示例:完整生产部署架构
YAML
# docker-compose.yml
version: "3.8"
services:
api:
build: .
ports:
- "8000:8000"
environment:
- REDIS_URL=redis://redis:6379
depends_on:
- redis
deploy:
replicas: 2
restart: unless-stopped
redis:
image: redis:7-alpine
ports:
- "6379:6379"
volumes:
- redis_data:/data
nginx:
image: nginx:alpine
ports:
- "80:80"
volumes:
- ./nginx.conf:/etc/nginx/conf.d/default.conf
depends_on:
- api
restart: unless-stopped
volumes:
redis_data:
TEXT
📖 仅展示
# nginx.conf (simplified load balancer)
upstream api_servers {
server api:8000;
}
server {
listen 80;
location / {
proxy_pass http://api_servers;
proxy_set_header Host $host;
}
}
▶ 示例:API + Redis缓存
PYTHON
# Enhanced app.py with Redis caching
from fastapi import FastAPI
from pydantic import BaseModel
import joblib
import numpy as np
import hashlib
import json
app = FastAPI(title="SalesPredict API with Cache")
model = joblib.load("salespredict_model.joblib")
# Redis cache (conceptual)
# import redis
# redis_client = redis.from_url(os.getenv("REDIS_URL", "redis://localhost:6379"))
class PredictionInput(BaseModel):
ad_spend_k: float
traffic_k: float
category_electronics: float = 0
category_clothing: float = 0
is_promotion: float = 0
def get_cache_key(input_data: PredictionInput) -> str:
data_str = json.dumps(input_data.model_dump(), sort_keys=True)
return f"pred:{hashlib.md5(data_str.encode()).hexdigest()}"
@app.post("/predict")
def predict(input_data: PredictionInput):
cache_key = get_cache_key(input_data)
# Check cache first
# cached = redis_client.get(cache_key)
# if cached:
# return json.loads(cached)
features = np.array([[input_data.ad_spend_k, input_data.traffic_k,
input_data.category_electronics,
input_data.category_clothing, input_data.is_promotion]])
prediction = float(model.predict(features)[0])
result = {"predicted_revenue_k": round(prediction, 2)}
# Cache for 5 minutes
# redis_client.setex(cache_key, 300, json.dumps(result))
return result
输出:
TEXT
📖 仅展示
# 函数定义成功
| 组件 | 作用 | 技术选择 |
|---|---|---|
| API服务 | 模型推理 | FastAPI + Uvicorn |
| 缓存 | 热点预测缓存 | Redis (TTL 5min) |
| 负载均衡 | 请求分发 | Nginx |
| 容器编排 | 服务管理 | docker-compose |
| 健康检查 | 故障检测 | /health + HEALTHCHECK |
❓ 常见问题
Q pickle和joblib哪个好?
A sklearn模型用joblib(对大numpy数组压缩更高效);通用Python对象用pickle。两者都有安全风险(不信任的pickle文件可能执行恶意代码),生产环境用MLflow或ONNX更安全。
Q FastAPI和Flask该用哪个?
A 新项目用FastAPI——自动文档(Swagger)、类型校验(Pydantic)、异步支持、性能更好。Flask更成熟但API开发体验不如FastAPI。
Q Docker镜像太大怎么办?
A 三招——1) 多阶段构建(build阶段不进入最终镜像);2) 用slim/alpine基础镜像;3) .dockerignore排除.git/data等。
Q 模型更新怎么不停服?
A 两种方案——1) 蓝绿部署(新旧版本切换);2) 滚动更新(docker-compose rolling update)。配合MLflow Model Registry管理版本。
Q API延迟怎么优化?
A 四层优化——1) Redis缓存热点请求;2) 批量预测减少开销;3) 多worker并行(Uvicorn workers);4) 模型量化(减小模型体积)。
Q 如何限制API请求速率?
A 用slowapi库实现rate limiting——
limiter = Limiter(key_func=get_remote_address),如限制每分钟100次请求。防止滥用和过载。📖 小节
- 模型序列化:sklearn用joblib,PyTorch用torch.save(state_dict),生产推荐MLflow
- FastAPI提供REST API:Pydantic校验输入、Swagger自动文档、async高并发
- Docker容器化:多阶段构建减小镜像、非root用户安全加固、HEALTHCHECK健康检查
- docker-compose编排:API + Redis缓存 + Nginx负载均衡的生产级部署
- 缓存策略:Redis缓存热点预测结果,TTL 5分钟,命中率30-50%
- API延迟目标:<50ms(含模型推理),批量预测进一步降低均摊延迟
📝 作业
- 基础题(难度⭐):训练一个sklearn模型,用joblib保存,然后在另一个Python脚本中加载并预测。提示:joblib.dump/load。
- 进阶题(难度⭐⭐):用FastAPI创建/predict端点,包含Pydantic输入校验和/health健康检查,用uvicorn运行并通过curl测试。提示:参考第4节完整API代码。
- 挑战题(难度⭐⭐⭐):编写Dockerfile(多阶段构建) + docker-compose.yml(API+Redis+Nginx),构建镜像并运行完整服务栈,验证负载均衡和缓存功能。提示:参考第5-6节的配置文件。