🔥 别再死磕枯燥理论了!AI 时代,拿实战作品说话才是硬道理!(原创不易哈,希望可以帮到还有些许学习劲儿的同学们) 🔥
【进阶版还在创作中,耗费精力中......】
入门实践工程十:PyTorch 模型量化与 ONNX 导出~附:安装依赖库及工程源码
简介
以 MNIST CNN 模型为例,演示动态量化(int8)对模型体积与推理速度的影响,并导出为 ONNX 格式用 onnxruntime 加载推理,对比 PyTorch 原始 / 动态量化 / ONNX 三者的准确率与耗时,直观理解模型优化与部署。
工程详细介绍
核心思想: 模型优化与部署------量化把全连接层权重从 float32 压成 int8 以减小体积、加速 CPU 推理;ONNX 作为中间表示让模型脱离 PyTorch,跨框架用 onnxruntime 高效推理,是模型上线的常用路径。
实现方法:
- 数据 / 模型: 复用入门实践工程 04 的 MNIST CNN 权重(缺失则自动快速训练 300 batch)。
- 流程: (1) PyTorch 原始推理计时 + 准确率;(2) torch.quantization.quantize_dynamic 对 Linear 层做 int8 动态量化再推理;(3) torch.onnx.export 导出含动态 batch 的 ONNX,onnxruntime 加载推理;(4) 汇总三者准确率与耗时对比表。
- 输出: 三方对比表(准确率 / 平均耗时)+ mnist_cnn.onnx。
目录结构
10_quant_onnx/
├── main.py # 训练(可选) + 量化 + ONNX 导出 + 对比
├── requirements.txt
├── data/ # MNIST 数据(自动下载)
├── mnist_cnn.pth # PyTorch 权重(无则自动训练)
├── mnist_cnn.onnx # 导出的 ONNX 模型
└── mnist_cnn_quant.onnx
正确安装
bash
pip install torch torchvision onnxruntime numpy
python main.py
- 若本目录无 mnist_cnn.pth,会自动快速训练一个(仅 300 batch)。
运行方式
bash
pip install -r requirements.txt
python main.py
预期结果
末尾输出对比表:
方式 准确率 平均耗时(ms)
PyTorch(原始) 0.98xx xx.xx
PyTorch(动态量化) 0.98xx xx.xx
ONNXRuntime 0.98xx xx.xx
说明
- 若
04_mnist_cnn/mnist_cnn.pth不在本目录,会自动快速训练一个(仅 300 batch)。 - 动态量化只对全连接层生效(卷积层不变),适合以线性层为主的模型。
- ONNX 推理在 CPU 上通常比原生 PyTorch 更快,且便于跨平台部署。
扩展方向
- 静态量化 + 校准(torch.quantization.quantize)
- 模型蒸馏:用小模型模仿大模型输出
- TensorRT 加速 / TFLite 移动端部署
工程源码
python
"""
入门实践工程十:PyTorch 模型量化与 ONNX 导出
=====================================
以一个训练好的 MNIST CNN 模型为例(若没有则现场快速训练一个),
演示动态量化对模型体积与推理速度的影响,并导出为 ONNX 格式,
用 onnxruntime 加载推理,对比 PyTorch / 量化 / ONNX 三者的结果与耗时。
运行:
python main.py
依赖:torch, torchvision, onnxruntime, numpy
"""
import os
import time
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
DATA_DIR = os.path.join(BASE_DIR, "data")
PTH_PATH = os.path.join(BASE_DIR, "mnist_cnn.pth")
ONNX_PATH = os.path.join(BASE_DIR, "mnist_cnn.onnx")
ONNX_Q_PATH = os.path.join(BASE_DIR, "mnist_cnn_quant.onnx")
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 复用入门实践工程 04 的轻量 CNN(此处自带一份,保证工程独立可运行)
class MnistCNN(nn.Module):
def __init__(self, num_classes: int = 10):
super().__init__()
self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
self.fc1 = nn.Linear(32 * 7 * 7, 128)
self.fc2 = nn.Linear(128, num_classes)
def forward(self, x):
x = F.relu(self.conv1(x))
x = F.max_pool2d(x, 2)
x = F.relu(self.conv2(x))
x = F.max_pool2d(x, 2)
x = x.flatten(1)
x = F.relu(self.fc1(x))
return self.fc2(x)
def get_model():
"""读取已有权重,没有则快速训练一个。"""
model = MnistCNN().to(DEVICE)
if os.path.exists(PTH_PATH):
print(f"加载已有权重: {PTH_PATH}")
model.load_state_dict(torch.load(PTH_PATH, map_location=DEVICE))
return model
print("未检测到权重,开始快速训练(1 个 Epoch)...")
import torch.optim as optim
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
train_set = datasets.MNIST(DATA_DIR, train=True, download=True, transform=transform)
loader = DataLoader(train_set, batch_size=128, shuffle=True)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
model.train()
for i, (x, y) in enumerate(loader):
if i >= 300: # 仅取 300 个 batch,快速跑通
break
x, y = x.to(DEVICE), y.to(DEVICE)
optimizer.zero_grad()
loss = F.cross_entropy(model(x), y)
loss.backward()
optimizer.step()
if i % 100 == 0:
print(f" batch {i} loss={loss.item():.4f}")
torch.save(model.state_dict(), PTH_PATH)
print(f"权重已保存到: {PTH_PATH}")
return model
def get_test_batch(n: int = 64):
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,)),
])
test_set = datasets.MNIST(DATA_DIR, train=False, download=True, transform=transform)
loader = DataLoader(test_set, batch_size=n, shuffle=False)
x, y = next(iter(loader))
return x, y
# ---------------- PyTorch 推理计时 ----------------
@torch.no_grad()
def torch_infer(model, x, repeats: int = 50):
model.eval()
model.to(DEVICE)
x = x.to(DEVICE)
# warmup
model(x)
t0 = time.time()
for _ in range(repeats):
model(x)
torch.cuda.synchronize() if DEVICE.type == "cuda" else None
elapsed = (time.time() - t0) / repeats
return model(x).cpu(), elapsed
# ---------------- 动态量化 ----------------
def dynamic_quantize(model):
"""对全连接层做 int8 动态量化,卷积层不变。"""
qmodel = torch.quantization.quantize_dynamic(
model, {nn.Linear}, dtype=torch.qint8
)
qmodel.eval()
return qmodel
@torch.no_grad()
def quant_infer(model, x, repeats: int = 50):
# 动态量化模型只能在 CPU 运行
model = model.to("cpu")
x = x.to("cpu")
model(x) # warmup
t0 = time.time()
for _ in range(repeats):
model(x)
elapsed = (time.time() - t0) / repeats
return model(x), elapsed
# ---------------- ONNX 导出与推理 ----------------
def export_onnx(model, onnx_path):
model = model.to("cpu").eval()
dummy = torch.randn(1, 1, 28, 28)
torch.onnx.export(
model, dummy, onnx_path,
input_names=["input"], output_names=["logits"],
dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}},
opset_version=14,
)
print(f"ONNX 已导出: {onnx_path} ({os.path.getsize(onnx_path)/1024:.1f} KB)")
def onnx_infer(onnx_path, x, repeats: int = 50):
import onnxruntime as ort
sess = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"])
x_np = x.numpy().astype(np.float32)
# warmup
sess.run(None, {"input": x_np})
t0 = time.time()
for _ in range(repeats):
out = sess.run(None, {"input": x_np})
elapsed = (time.time() - t0) / repeats
return out[0], elapsed
def main():
torch.manual_seed(42)
model = get_model()
x, y = get_test_batch(64)
print("\n=== 1. PyTorch 原始模型推理 ===")
out, t_torch = torch_infer(model, x)
acc_torch = (out.argmax(1) == y).float().mean().item()
print(f"准确率={acc_torch:.4f} 平均耗时={t_torch*1000:.2f} ms")
print("\n=== 2. 动态量化(int8,仅全连接层) ===")
qmodel = dynamic_quantize(model)
qout, t_quant = quant_infer(qmodel, x)
acc_quant = (qout.argmax(1) == y).float().mean().item()
print(f"准确率={acc_quant:.4f} 平均耗时={t_quant*1000:.2f} ms")
print("\n=== 3. 导出 ONNX 并用 onnxruntime 推理 ===")
export_onnx(model, ONNX_PATH)
onnx_out, t_onnx = onnx_infer(ONNX_PATH, x)
acc_onnx = (onnx_out.argmax(1) == y.numpy()).mean()
print(f"准确率={acc_onnx:.4f} 平均耗时={t_onnx*1000:.2f} ms")
print("\n=== 汇总对比 ===")
print(f"{'方式':<20}{'准确率':<12}{'平均耗时(ms)':<14}")
print(f"{'PyTorch(原始)':<20}{acc_torch:<12.4f}{t_torch*1000:<14.2f}")
print(f"{'PyTorch(动态量化)':<20}{acc_quant:<12.4f}{t_quant*1000:<14.2f}")
print(f"{'ONNXRuntime':<20}{acc_onnx:<12.4f}{t_onnx*1000:<14.2f}")
print("\n说明:动态量化主要压缩全连接层权重为 int8,体积更小、CPU 推理往往更快;")
print("ONNX 在 CPU 上的推理速度通常也优于原生 PyTorch,便于跨平台部署。")
if __name__ == "__main__":
main()