从训练到部署:手把手教你用PyTorch落地一个AI推理服务
将训练好的PyTorch模型部署为一个生产可用的推理服务,是把AI研究成果转化为实际价值的关键一步。本文将以一个完整的图像分类任务为例,手把手带你走完从模型训练、导出、服务化到容器化部署的全流程。
一、训练并保存模型
我们首先使用PyTorch训练一个简单的图像分类模型。本例中,我们使用CIFAR-10数据集训练一个基础的卷积神经网络(CNN)。
训练的核心是让模型学习从输入图像到正确类别标签的映射。训练完成后,我们需要将模型保存下来,以便在后续的部署环节中加载使用。
# train.py
import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader
# 1. 定义模型架构
class SimpleCNN(nn.Module):
def __init__(self):
super(SimpleCNN, self).__init__()
self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.pool = nn.MaxPool2d(2, 2)
self.fc1 = nn.Linear(64 * 8 * 8, 128)
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
x = self.pool(torch.relu(self.conv1(x)))
x = self.pool(torch.relu(self.conv2(x)))
x = x.view(-1, 64 * 8 * 8)
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
# 2. 数据预处理与加载
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
trainloader = DataLoader(trainset, batch_size=64, shuffle=True)
# 3. 初始化模型、损失函数和优化器
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = SimpleCNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 4. 训练循环
print("开始训练...")
for epoch in range(5): # 简单起见,训练5个epoch
running_loss = 0.0
for i, data in enumerate(trainloader, 0):
inputs, labels = data[0].to(device), data[1].to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
print(f'Epoch {epoch + 1}, Loss: {running_loss / len(trainloader):.4f}')
print("训练完成!")
# 5. 保存模型权重 (state_dict)
# 推荐只保存模型的参数(state_dict),而不是整个模型对象
# 这样更灵活,且与模型定义解耦
torch.save(model.state_dict(), 'model_weights.pth')
print("模型权重已保存为 model_weights.pth")
关于模型保存的最佳实践:官方推荐使用
torch.save(model.state_dict(), 'model_weights.pth')仅保存模型的参数(state_dict)。这是因为state_dict是一个包含所有可学习参数的Python字典,保存它比保存整个模型对象更轻量、更安全。加载时,你需要先实例化模型对象,然后通过load_state_dict方法加载参数。
二、模型导出:为部署做好准备
训练好的 .pth 权重文件依赖于原始的Python模型类定义,无法直接在C++、Java等其他环境中使用。为了提升部署的灵活性和性能,我们需要将模型导出为与Python环境解耦的格式。
PyTorch官方提供了两种主要的导出方案:
- TorchScript:PyTorch的原生静态图表示,可在C++等环境中高性能运行。
- ONNX:一种开放的模型格式,可实现跨框架(如TensorRT)的部署。
对于大多数Python后端服务,使用TorchScript是兼顾性能与易用性的好选择。TorchScript的导出主要有两种方式:
- Tracing(跟踪):通过跟踪一次模型的前向传播来记录计算图。这种方式简单直接,但不适用于包含复杂控制流(如if-else循环)的模型。
- Scripting(脚本化):通过解析模型的源代码来编译。它支持控制流,但对代码有更多限制。
以下代码使用 torch.jit.trace 方法导出我们训练好的模型。对于更复杂的模型,可能需要使用 torch.jit.script。
# export.py
import torch
import torchvision.transforms as transforms
from PIL import Image
from train import SimpleCNN # 导入模型定义
# 1. 加载模型权重
model = SimpleCNN()
# 加载 state_dict 到模型中
model.load_state_dict(torch.load('model_weights.pth', map_location=torch.device('cpu')))
model.eval() # 切换到评估模式
# 2. 创建示例输入 (用于跟踪)
# 输入维度应为 [batch_size, channels, height, width]
example_input = torch.rand(1, 3, 32, 32)
# 3. 使用跟踪(Tracing)导出为 TorchScript
traced_script_module = torch.jit.trace(model, example_input)
# 4. 保存 TorchScript 模型
traced_script_module.save("model.torchscript.pt")
print("TorchScript模型已保存为 model.torchscript.pt")
导出后的 model.torchscript.pt 文件是独立的,不再依赖于原始的 SimpleCNN 类定义。
三、构建推理服务:让模型“活”起来
模型导出后,我们需要将其封装成一个可提供服务的API。FastAPI因其高性能、自动生成API文档以及现代异步特性,已成为构建AI推理服务的首选框架之一。
3.1 核心代码 (app/main.py)
# app/main.py
import io
import torch
import torchvision.transforms as transforms
from PIL import Image
from fastapi import FastAPI, File, UploadFile, HTTPException
from fastapi.responses import JSONResponse
# --- 1. 定义数据预处理 ---
# 必须与训练时的预处理保持一致
transform = transforms.Compose([
transforms.Resize((32, 32)),
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
# CIFAR-10 的类别标签
CLASSES = ['airplane', 'automobile', 'bird', 'cat', 'deer',
'dog', 'frog', 'horse', 'ship', 'truck']
# --- 2. 加载 TorchScript 模型 ---
# 使用 torch.jit.load 加载导出的模型
model = torch.jit.load("model.torchscript.pt", map_location=torch.device('cpu'))
model.eval()
print("TorchScript模型加载成功!")
# --- 3. 定义 FastAPI 应用 ---
app = FastAPI(title="PyTorch 图像分类服务", version="1.0.0")
@app.get("/health")
async def health_check():
"""健康检查端点"""
return {"status": "ok"}
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
"""
接收上传的图片,返回预测结果
"""
# 验证文件类型
if not file.content_type.startswith("image/"):
raise HTTPException(status_code=400, detail="上传的文件不是图片")
try:
# 读取图片数据
contents = await file.read()
image = Image.open(io.BytesIO(contents)).convert('RGB')
# 预处理
image_tensor = transform(image).unsqueeze(0) # 增加batch维度
# 推理
with torch.no_grad(): # 禁用梯度计算
outputs = model(image_tensor)
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
# 解析结果
predicted_idx = torch.argmax(probabilities).item()
predicted_label = CLASSES[predicted_idx]
confidence = probabilities[predicted_idx].item()
return {
"predicted_label": predicted_label,
"confidence": confidence,
"predicted_index": predicted_idx
}
except Exception as e:
raise HTTPException(status_code=500, detail=f"推理失败: {str(e)}")
3.2 启动服务
为了在生产环境中获得更好的性能,建议使用 uvicorn 启动服务。
# 安装依赖
pip install fastapi uvicorn[standard] torch torchvision pillow python-multipart
# 启动服务 (监听所有网络接口)
uvicorn app.main:app --host 0.0.0.0 --port 8000
服务启动后,你可以在浏览器中访问 http://127.0.0.1:8000/docs 查看由FastAPI自动生成的交互式API文档,并进行在线测试。
3.3 测试API
你可以使用 curl 命令来测试服务是否正常工作。
curl -X POST "http://127.0.0.1:8000/predict" \
-H "accept: application/json" \
-H "Content-Type: multipart/form-data" \
-F "file=@path/to/your/test_image.jpg"
如果一切正常,你会得到一个类似下面的JSON响应:
{
"predicted_label": "cat",
"confidence": 0.92,
"predicted_index": 3
}
四、容器化部署:终结环境不一致
为了确保服务在任何环境中都能稳定运行,我们需要使用Docker将应用及其所有依赖打包成一个镜像。这是解决“在我机器上能跑”问题的标准方案。
4.1 编写 Dockerfile
我们使用官方的PyTorch镜像作为基础,这样可以确保CUDA等底层库的正确配置。
# Dockerfile
# 使用官方 PyTorch 镜像作为基础
FROM pytorch/pytorch:2.0.1-cuda11.7-cudnn8-runtime
# 设置工作目录
WORKDIR /app
# 先复制依赖文件,利用 Docker 缓存层
COPY requirements.txt .
RUN pip install --no-cache-dir -r requirements.txt
# 复制应用代码和模型文件
COPY app/ ./app/
COPY model.torchscript.pt .
# 暴露服务端口
EXPOSE 8000
# 启动命令
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
4.2 准备 requirements.txt
# requirements.txt
fastapi==0.104.1
uvicorn[standard]==0.24.0
torch==2.0.1
torchvision==0.15.2
pillow==10.1.0
python-multipart==0.0.6
4.3 构建与运行
# 1. 构建镜像
docker build -t pytorch-inference:latest .
# 2. 运行容器
docker run -d -p 8000:8000 --name pytorch-service pytorch-inference:latest
# 3. 查看容器日志
docker logs -f pytorch-service
现在,你的PyTorch推理服务已经在一个独立的Docker容器中运行,可以通过宿主机的 8000 端口进行访问。你也可以使用 docker-compose.yml 来编排更复杂的服务栈。
五、生产环境进阶优化
当服务需要应对更高并发和更严苛的延迟要求时,可以考虑以下优化方向:
- 模型量化(Quantization):将模型权重从FP32转换为INT8,可显著减小模型体积并提升推理速度。
- 批处理(Batching):在服务端积累多个请求,合并成一个Batch进行推理,可大幅提升GPU利用率。
- 使用专用推理服务器:对于超大规模模型或要求极高性能的场景,可考虑TorchServe或Triton Inference Server等专用框架。
总结
从训练到部署,我们完成了一个PyTorch模型的全生命周期管理:
- 模型训练:使用PyTorch训练了一个图像分类模型,并保存了其权重(
state_dict)。 - 模型导出:将模型转换为与Python环境解耦的TorchScript格式。
- 服务构建:使用FastAPI创建了一个高性能的REST API服务。
- 容器化部署:通过Docker将服务打包,实现了环境的一致性和可移植性。
掌握这套工作流,是成为一名合格AI工程师的关键一步。它将帮助你打通从算法到产品的“最后一公里”,让模型真正在业务场景中发挥价值。
更多推荐




所有评论(0)