入门实践工程十:PyTorch 模型量化与 ONNX 导出~附:安装依赖库及工程源码

🔥 别再死磕枯燥理论了!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 移动端部署

工程源码

main.py

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()
相关推荐
一直都在5722 小时前
LangChain4j精讲
开发语言·人工智能
八号当铺2 小时前
使用 Figma Agent Kit:插件 + MCP + 还原 Skill,打通本地设计协作
前端·人工智能·ai编程
今天AI了吗2 小时前
时序大模型 TimechoAI 实战:从数据接入到智能时序分析全链路指南
人工智能
AI服务老曹2 小时前
多路摄像头AI分析完整流程:硬件选型与GPU/NPU算力估算指南
人工智能
触底反弹2 小时前
🚀 浏览器里跑 1.5B 参数大模型?我用 WebGPU + DeepSeek 做到了
人工智能·面试·typescript
pearbing2 小时前
AI搜索流量密码:8个核心GEO优化打法,拉高品牌曝光优先级
人工智能·geo
品牌测评2 小时前
Token Plan平台分享|七条算力订阅路径拆解
大数据·人工智能·架构
大模型码小白3 小时前
【AI】一文讲清 RAG:从大模型局限到企业级知识库落地流程
人工智能·深度学习·学习
MomentYY3 小时前
RAG 图检索&多跳推理:有些答案需要“顺藤摸瓜”
人工智能·agent·ai编程