将训练好的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官方提供了两种主要的导出方案:

  1. TorchScript:PyTorch的原生静态图表示,可在C++等环境中高性能运行。
  2. 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利用率。
  • 使用专用推理服务器:对于超大规模模型或要求极高性能的场景,可考虑TorchServeTriton Inference Server等专用框架。

总结

从训练到部署,我们完成了一个PyTorch模型的全生命周期管理:

  1. 模型训练:使用PyTorch训练了一个图像分类模型,并保存了其权重(state_dict)。
  2. 模型导出:将模型转换为与Python环境解耦的TorchScript格式。
  3. 服务构建:使用FastAPI创建了一个高性能的REST API服务。
  4. 容器化部署:通过Docker将服务打包,实现了环境的一致性和可移植性。

掌握这套工作流,是成为一名合格AI工程师的关键一步。它将帮助你打通从算法到产品的“最后一公里”,让模型真正在业务场景中发挥价值。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐