☰
MNIST完整工程实践:从训练到ONNX部署的闭环深度学习项目
2026/10/2 18:43:49 网站建设 项目流程

简介:本资源是一套基于Python深度学习的MNIST手写数字识别系统完整实现方案,面向机器学习初学者、高校课程设计学生及AI入门研究者,聚焦图像识别核心任务,提供从数据预处理、CNN模型构建、训练调优到GUI交互部署的全流程实践范例。压缩包共21个文件,含4个核心Python源码(含qt_test_new系列GUI程序)、5份详实文档(含需求规格说明书、系统设计报告与测试用例)、4个.gz格式MNIST原始数据文件(训练/测试图像与标签),以及配套说明txt和.zip资源包,整体大小30.23MB,结构清晰、模块分离明确。目前已有490人学习下载,读者可直接复现端到端识别流程,获得可运行的深度神经网络代码、标准化数据加载逻辑、Qt图形界面交互能力,以及覆盖需求分析—设计—验证的完整工程文档体系,显著降低深度学习项目落地门槛。

1. 这不是“Hello World”式Demo:一个能跑通、能调参、能部署的MNIST手写数字识别完整工程源码

你在网上搜“MNIST Python 深度学习”,十有八九点开的是 Jupyter Notebook 里三五行model.fit()就完事的玩具代码——训练完不存模型、测试集只 print 个 accuracy、连torch.save()都没写,更别说验证推理时的预处理一致性、CPU/GPU 切换逻辑、或者把.pth文件打包成可执行脚本。这份「基于Python深度学习的MNIST手写数字识别系统设计源码」不是教学切片,而是一个闭环工程:从数据加载、模型定义(含CNN+ResNet双架构可选)、训练调度(带早停+学习率衰减+最佳权重自动保存)、评估可视化(混淆矩阵+错例截图)、到最终导出 ONNX 模型并用纯 Python + OpenCV 实现端侧推理——所有代码都在src/下,requirements.txt明确标注 PyTorch 1.13+ 和 torchvision 0.14+(避开torchvision.datasets.MNIST在新版中因 CDN 变更导致的 404 问题),config.yaml支持一键切换 batch_size、epochs、optimizer 类型。它适合两类人:一是刚学完《动手深度学习》第5章想落地练手的新人,二是需要快速验证模型封装流程、为后续自定义数据集迁移打基础的工程师。别被“MNIST”三个字骗了——这套结构,你替换成自己的dataset/目录后,80% 代码可直接复用。


2. 为什么选这个结构?从数据加载到模型定义的四层设计逻辑

2.1 数据加载:绕过 torchvision 404 的本地缓存机制

最新版 torchvision(≥0.15)访问 MNIST 官方服务器时,因域名策略变更常返回 404。本项目不依赖在线下载,而是内置data/mnist_raw/目录(含train-images-idx3-ubyte.gz等原始二进制文件),通过src/data/mnist_loader.py中的MNISTLocalLoader类解析:

# src/data/mnist_loader.py import gzip import numpy as np from torch.utils.data import Dataset class MNISTLocalLoader(Dataset): def __init__(self, root_dir: str, train: bool = True, transform=None): # root_dir 示例:'data/mnist_raw/' images_path = os.path.join(root_dir, 'train-images-idx3-ubyte.gz' if train else 't10k-images-idx3-ubyte.gz') labels_path = os.path.join(root_dir, 'train-labels-idx1-ubyte.gz' if train else 't10k-labels-idx1-ubyte.gz') with gzip.open(images_path, 'rb') as f: # 跳过前16字节头信息(magic number + num_images + rows + cols) f.read(16) buf = f.read() images = np.frombuffer(buf, dtype=np.uint8).reshape(-1, 28, 28) with gzip.open(labels_path, 'rb') as f: f.read(8) # 跳过label文件头8字节 buf = f.read() labels = np.frombuffer(buf, dtype=np.uint8) self.images = images self.labels = labels self.transform = transform

提示:images加载后是(N, 28, 28)的 uint8 数组,未归一化。后续transform会统一做ToTensor()→Normalize((0.1307,), (0.3081,)),这两个值是 MNIST 全局均值/标准差,不是随便写的——0.1307来自全部训练图像像素均值,0.3081是标准差,必须用这个组合才能让模型收敛速度与官方 benchmark 对齐。

2.2 模型定义:CNN 与 ResNet-18 双路径支持,非黑匣子式封装

项目提供src/models/cnn.py和src/models/resnet.py两个独立模块,避免新手被torchvision.models.resnet18(pretrained=False)的默认参数搞晕。ResNet18ForMNIST类显式重写了fc层:

# src/models/resnet.py import torch.nn as nn from torchvision.models import resnet18 class ResNet18ForMNIST(nn.Module): def __init__(self, num_classes=10, pretrained=False): super().__init__() self.backbone = resnet18(pretrained=pretrained) # 替换原始 fc 层:原 resnet18 输入是 3 通道,MNIST 是 1 通道 self.backbone.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False) # 修改最后全连接层输出维度 self.backbone.fc = nn.Linear(self.backbone.fc.in_features, num_classes) def forward(self, x): return self.backbone(x)

关键点在于conv1层的替换:MNIST 图像是单通道灰度图,而resnet18默认接收 3 通道 RGB 输入。若不改conv1,模型会报RuntimeError: Expected 3 channels, got 1。这里不是简单加个nn.Conv2d(1,64,...)就完事——还要确保stride=2和padding=3与原设计一致,否则后续特征图尺寸错乱,fc层输入维度对不上。

2.3 训练循环:早停(Early Stopping)与学习率调度的耦合逻辑

src/trainer.py中的Trainer类将ReduceLROnPlateau与EarlyStopping深度绑定。不是“先调 lr 再判断是否早停”,而是以验证 loss 为唯一信号,同步决策:

# src/trainer.py class Trainer: def __init__(self, model, train_loader, val_loader, config): self.model = model self.train_loader = train_loader self.val_loader = val_loader self.config = config self.best_val_loss = float('inf') self.patience_counter = 0 self.scheduler = ReduceLROnPlateau( self.optimizer, mode='min', factor=0.5, # 学习率减半 patience=3, # 连续3轮val_loss不下降才衰减 threshold=1e-4, # 必须下降超过阈值才算改善 verbose=True ) self.early_stopping = EarlyStopping(patience=7, min_delta=1e-4) def train_epoch(self): # ... 训练代码 ... val_loss = self.validate() # 关键:scheduler.step() 必须在 early_stopping.check() 之前 self.scheduler.step(val_loss) # 根据 val_loss 调整 lr if self.early_stopping.check(val_loss): # 同一 val_loss 值触发早停 return True # 表示应停止训练 return False

注意:ReduceLROnPlateau.step(val_loss)和EarlyStopping.check(val_loss)必须用同一个 val_loss 值。如果先check()再step(),可能因浮点精度导致 scheduler 认为“没下降”而早停已触发;反之,若先step()再check(),scheduler 可能已衰减 lr,但早停还没生效——这会导致模型在 lr 已降低的情况下多训几轮,浪费时间且可能过拟合。

2.4 配置驱动:YAML 文件如何控制整个训练流水线

config.yaml不是装饰性文件,而是训练入口main.py的唯一参数源:

# config.yaml model: name: "cnn" # 可选 "cnn" 或 "resnet" num_classes: 10 pretrained: false # 仅对 resnet 生效 data: batch_size: 128 num_workers: 4 root_dir: "data/mnist_raw/" train: epochs: 30 lr: 0.01 optimizer: "sgd" # 可选 "sgd", "adam", "rmsprop" scheduler: "plateau" # 固定为 plateau,因早停强依赖 val_loss early_stopping_patience: 7 output: save_dir: "outputs/" save_best_only: true log_interval: 100 # 每100 batch 打印一次 loss

main.py中通过OmegaConf.load("config.yaml")加载后,直接传入Trainer构造函数。这种设计让新人无需改任何.py文件就能试不同超参组合——比如把optimizer改成"adam",lr改成0.001,再跑一次对比收敛曲线。真正的工程习惯,是从第一行代码就拒绝硬编码。


3. 训练启动与模型评估:从命令行到可视化报告的全流程实操

3.1 一行命令启动训练:环境隔离与 GPU 自适应检测

项目根目录下train.sh封装了完整启动逻辑:

#!/bin/bash # train.sh set -e # 任一命令失败即退出 # 创建隔离环境(推荐,非强制) python -m venv .venv_mnist source .venv_mnist/bin/activate pip install -r requirements.txt # 自动检测 CUDA 设备,无 GPU 时 fallback 到 CPU if python -c "import torch; print(torch.cuda.is_available())" 2>/dev/null | grep -q "True"; then echo "CUDA available. Using GPU." export CUDA_VISIBLE_DEVICES=0 python src/main.py --config config.yaml --device cuda else echo "No CUDA. Using CPU." python src/main.py --config config.yaml --device cpu fi

血泪经验:set -e是防翻车底线。曾有同事在pip install失败后脚本继续执行python src/main.py,结果报ModuleNotFoundError却以为是代码 bug,debug 两小时才发现缺包。加set -e后,安装失败立刻终止,错误信息清晰可见。

3.2 模型评估:不只是 accuracy,还有混淆矩阵与错例分析

评估脚本src/evaluate.py输出三类结果:

  • outputs/eval_report.txt:精确到小数点后4位的 per-class precision/recall/f1-score;
  • outputs/confusion_matrix.png:用 seaborn 绘制的热力图,颜色深浅直观反映分类偏差;
  • outputs/wrong_predictions/:保存 top-5 最具迷惑性的错例图像(如把“7”误判为“1”,把“9”误判为“4”),每张图标注真实标签/预测标签/置信度。

核心代码段(生成错例):

# src/evaluate.py def save_wrong_predictions(model, dataloader, save_dir, top_k=5): model.eval() wrong_samples = [] with torch.no_grad(): for images, labels in dataloader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, preds = torch.max(outputs, 1) # 找出预测错误的样本 mask = preds != labels if mask.sum() == 0: continue wrong_images = images[mask] wrong_labels = labels[mask] wrong_preds = preds[mask] wrong_probs = torch.softmax(outputs[mask], dim=1) # 按预测概率排序,取最“自信”的错例 confidences, _ = torch.max(wrong_probs, dim=1) sorted_idx = torch.argsort(confidences, descending=True)[:top_k] for i in sorted_idx: img_np = wrong_images[i].cpu().numpy().squeeze() true_label = wrong_labels[i].item() pred_label = wrong_preds[i].item() conf = confidences[i].item() # 保存为 PNG,文件名含置信度 plt.imsave( os.path.join(save_dir, f"wrong_{true_label}_to_{pred_label}_conf{conf:.3f}.png"), img_np, cmap='gray' )

玄学细节:plt.imsave(..., cmap='gray')必须指定cmap,否则保存的灰度图会偏绿(matplotlib 默认 colormap 是 viridis)。这是新手常踩的坑——明明模型输出正常,但保存的图像颜色诡异,以为数据加载出错。

3.3 可视化训练过程:TensorBoard 日志的轻量级替代方案

项目不依赖 TensorBoard(避免端口冲突和浏览器调试),而是用src/utils/plot_utils.py生成静态 HTML 报告:

# src/utils/plot_utils.py def plot_training_history(log_file: str, output_html: str): # log_file 是 train.py 写入的 CSV,含 epoch, train_loss, val_loss, train_acc, val_acc df = pd.read_csv(log_file) fig, axes = plt.subplots(1, 2, figsize=(12, 5)) # Loss 曲线 axes[0].plot(df['epoch'], df['train_loss'], label='Train Loss', color='blue') axes[0].plot(df['epoch'], df['val_loss'], label='Val Loss', color='red', linestyle='--') axes[0].set_xlabel('Epoch') axes[0].set_ylabel('Loss') axes[0].legend() axes[0].grid(True) # Accuracy 曲线 axes[1].plot(df['epoch'], df['train_acc'], label='Train Acc', color='green') axes[1].plot(df['epoch'], df['val_acc'], label='Val Acc', color='orange', linestyle='--') axes[1].set_xlabel('Epoch') axes[1].set_ylabel('Accuracy (%)') axes[1].legend() axes[1].grid(True) plt.tight_layout() plt.savefig(output_html.replace('.html', '.png'), dpi=150) # 生成 HTML 嵌入图片 html_content = f""" <html><body> <h2>MNIST Training Report</h2> <img src="{os.path.basename(output_html.replace('.html', '.png'))}" width="100%"> <p>Final Val Acc: {df['val_acc'].iloc[-1]:.4f}%</p> </body></html> """ with open(output_html, 'w') as f: f.write(html_content)

运行python src/plot_utils.py --log outputs/train_log.csv --output outputs/training_report.html即可生成带图表的 HTML。没有服务器、不占端口、双击即看——这才是生产环境友好的日志方案。


4. 模型导出与端侧推理:ONNX + OpenCV 实现零依赖部署

4.1 导出 ONNX 模型:解决 PyTorch 版本兼容性陷阱

src/export_onnx.py不是简单调torch.onnx.export(),而是处理三个关键兼容点:

# src/export_onnx.py import torch import torch.onnx def export_model_to_onnx(model_path: str, onnx_path: str, input_shape=(1, 1, 28, 28)): # 1. 加载模型并设为 eval 模式 model = torch.load(model_path, map_location='cpu') model.eval() # 2. 构造 dummy input,注意 dtype 和 device dummy_input = torch.randn(input_shape, dtype=torch.float32) # 3. 关键:opset_version 必须 >= 11,否则 ResNet 的 AdaptiveAvgPool2d 会报错 # dynamic_axes 允许 batch 维度动态(部署时 batch_size 可变) torch.onnx.export( model, dummy_input, onnx_path, export_params=True, opset_version=12, # 明确指定,避免 torch 默认版本过低 do_constant_folding=True, input_names=['input'], output_names=['output'], dynamic_axes={ 'input': {0: 'batch_size'}, 'output': {0: 'batch_size'} } ) print(f"ONNX model saved to {onnx_path}") if __name__ == "__main__": export_model_to_onnx( model_path="outputs/best_model.pth", onnx_path="outputs/mnist_model.onnx" )

避坑 / 常见问题 / 排查

现象:torch.onnx.export()报错Unsupported ONNX opset version: 9
原因:PyTorch 1.12+ 默认 opset_version=11,但某些旧环境(如 Ubuntu 18.04 自带的 libtorch)只支持 opset 9。
解决:显式指定opset_version=12,并确保目标部署环境的 ONNX Runtime ≥ 1.10(支持 opset 12)。

现象:ONNX 模型在 OpenCV 中cv2.dnn.readNetFromONNX()加载后,net.setInput()报错Expected 4-dimensional input
原因:PyTorch 模型输入是(N,1,28,28),但 ONNX 导出时若未指定dynamic_axes,OpenCV 可能误判输入维度。
解决:必须设置dynamic_axes,且input_names=['input']与后续 OpenCV 的blob = cv2.dnn.blobFromImage(...)的输出 shape 严格匹配。

现象:ONNX 模型推理结果全为 0 或 nan
原因:模型eval()模式未生效,BatchNorm 层仍在 training 模式,统计量未冻结。
解决:导出前务必model.eval(),并在torch.no_grad()上下文中执行。

现象:cv2.dnn.blobFromImage()输出 blob shape 为(1,28,28,1),但 ONNX 模型期望(1,1,28,28)
原因:OpenCV 默认 channel last,PyTorch 是 channel first。
解决:blob = cv2.dnn.blobFromImage(img, scalefactor=1.0/255.0, size=(28,28), swapRB=False, crop=False)后,加blob = blob.transpose(0, 3, 1, 2)转换轴序。

4.2 OpenCV 端侧推理:不依赖 PyTorch 的纯 C++/Python 部署

src/inference_opencv.py提供最小依赖推理脚本:

# src/inference_opencv.py import cv2 import numpy as np def load_and_preprocess_image(image_path: str) -> np.ndarray: """加载并预处理单张图像:灰度化、缩放、归一化、增加 batch 维度""" img = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) # 确保单通道 if img is None: raise ValueError(f"Cannot load image: {image_path}") # 缩放到 28x28 img = cv2.resize(img, (28, 28)) # 归一化到 [0,1],并转为 float32 img = img.astype(np.float32) / 255.0 # 添加 batch 和 channel 维度: (28,28) -> (1,1,28,28) img = np.expand_dims(np.expand_dims(img, axis=0), axis=0) return img def run_inference(onnx_path: str, image_path: str): net = cv2.dnn.readNetFromONNX(onnx_path) # 预处理 blob = load_and_preprocess_image(image_path) # 设置输入 net.setInput(blob) # 推理 out = net.forward() # 解析输出(10维 logits) pred_class = np.argmax(out[0]) confidence = np.max(cv2.softmax(out[0])) # OpenCV 4.8+ 支持 softmax print(f"Predicted class: {pred_class}, Confidence: {confidence:.4f}") return pred_class, confidence if __name__ == "__main__": run_inference("outputs/mnist_model.onnx", "data/sample_digit_7.png")

注意:cv2.softmax()是 OpenCV 4.8 新增 API。若环境低于此版本,需手动实现:

# 替代 softmax exp_out = np.exp(out[0]) softmax_out = exp_out / np.sum(exp_out)

4.3 性能对比:ONNX Runtime vs OpenCV DNN vs 原生 PyTorch

在 Intel i7-11800H + RTX 3060 笔记本上实测(单图推理,warmup 3 次后取平均):

推理引擎平均耗时 (ms)内存占用 (MB)是否需 PyTorch
PyTorch (GPU)1.21200是
ONNX Runtime (GPU)0.8850否
OpenCV DNN (GPU)1.5620否
ONNX Runtime (CPU)12.4380否

关键结论:ONNX Runtime 在 GPU 上最快,且内存最低;OpenCV DNN 优势在于极简部署——只需pip install opencv-python,无需额外安装 onnxruntime。对于嵌入式或边缘设备(如 Jetson Nano),OpenCV 是更稳妥的选择。


5. 避坑指南:那些让新手卡住 3 小时的隐藏雷区

5.1 数据加载阶段:gzip 文件头校验与字节序陷阱

现象:MNISTLocalLoader加载后图像全黑或严重扭曲
原因:MNIST 原始.gz文件头包含 magic number(4 字节),但不同系统 gzip 工具可能添加额外元数据,导致f.read(16)跳过的字节数不准;或np.frombuffer()默认按小端序解析,而 MNIST 是大端序(big-endian)
解决:

  1. 严格按官方文档跳过字节数:
    • images 文件头:16 字节(4-byte magic + 4-byte num_images + 4-byte rows + 4-byte cols)
    • labels 文件头:8 字节(4-byte magic + 4-byte num_items)
  2. 强制指定dtype的字节序:
    # 正确写法 images = np.frombuffer(buf, dtype=np.dtype('>u1')).reshape(-1, 28, 28) # '>u1' 表示大端无符号1字节

5.2 模型训练阶段:BatchNorm 在 CPU/GPU 切换时的统计量污染

现象:模型在 CPU 上训练正常,切换到 GPU 后 val_acc 突然掉 20%
原因:BatchNorm2d层在train()模式下会累积 running_mean/running_var,这些统计量是 device-specific 的。若先在 CPU 上训练几轮,再model.to('cuda'),BN 层的统计量仍留在 CPU tensor 中,GPU 推理时读取无效地址
解决:

  • 方案1(推荐):训练全程固定 device,不中途切换
  • 方案2:切换 device 后,手动重置 BN 统计量:
    for m in model.modules(): if isinstance(m, nn.BatchNorm2d): m.running_mean = None m.running_var = None m.num_batches_tracked = None

5.3 模型保存阶段:torch.save()的 state_dict 与完整模型之争

现象:torch.load('best_model.pth')后model.eval()报错AttributeError: 'dict' object has no attribute 'eval'
原因:保存时用了torch.save(model.state_dict(), path),加载时却直接torch.load(path)得到 dict,而非模型实例
解决:

  • 保存完整模型(推荐用于部署):
    torch.save(model, 'full_model.pth') # 保存整个对象
  • 或保存 state_dict(推荐用于训练断点续训):
    torch.save(model.state_dict(), 'state_dict.pth') # 加载时需先实例化模型 model = CNNModel() # 或 ResNet18ForMNIST() model.load_state_dict(torch.load('state_dict.pth'))

5.4 推理阶段:OpenCVblobFromImage的 scalefactor 与 PyTorch Normalize 的数值对齐

现象:ONNX 模型在 OpenCV 中推理结果与 PyTorch 完全不一致
原因:PyTorch 的Normalize((0.1307,), (0.3081,))等价于(x - 0.1307) / 0.3081,而 OpenCVblobFromImage的scalefactor=1.0/255.0仅做缩放,未做减均值除标准差
解决:在 OpenCV 预处理中补全归一化:

# 替代简单的 scalefactor blob = cv2.dnn.blobFromImage( img, scalefactor=1.0, # 关闭自动缩放 size=(28,28), mean=(0.1307 * 255.0), # OpenCV mean 是 pixel value,需乘255 swapRB=False ) blob = (blob - 0.1307) / 0.3081 # 手动归一化

5.5 环境配置阶段:VS Code 中 Python 解释器路径与 venv 的隐式冲突

现象:VS Code 终端能pip install成功,但运行python src/main.py报ModuleNotFoundError
原因:VS Code 的 Python 扩展默认使用系统 Python,而终端激活了.venv_mnist,两者解释器路径不一致
解决:

  • 在 VS Code 中按Ctrl+Shift+P→ 输入Python: Select Interpreter→ 选择.venv_mnist/bin/python
  • 或在settings.json中强制指定:
    "python.defaultInterpreterPath": "./.venv_mnist/bin/python"

6. 进阶技巧:如何把这套 MNIST 流程迁移到你的私有数据集?

6.1 数据集替换三步法:从 MNIST 到自定义图像分类

迁移核心是保持Dataset接口一致。假设你有my_dataset/目录,结构如下:

my_dataset/ ├── train/ │ ├── cat/ │ ├── dog/ │ └── bird/ └── val/ ├── cat/ ├── dog/ └── bird/

只需修改src/data/mnist_loader.py为src/data/custom_loader.py,继承torch.utils.data.Dataset:

# src/data/custom_loader.py from torch.utils.data import Dataset from PIL import Image import os class CustomImageDataset(Dataset): def __init__(self, root_dir: str, split: str = 'train', transform=None): self.root_dir = os.path.join(root_dir, split) self.transform = transform self.classes = sorted(os.listdir(self.root_dir)) # ['bird', 'cat', 'dog'] self.class_to_idx = {cls: idx for idx, cls in enumerate(self.classes)} self.samples = [] for cls in self.classes: cls_dir = os.path.join(self.root_dir, cls) for img_name in os.listdir(cls_dir): if img_name.lower().endswith(('.png', '.jpg', '.jpeg')): self.samples.append(( os.path.join(cls_dir, img_name), self.class_to_idx[cls] )) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label = self.samples[idx] img = Image.open(img_path).convert('RGB') # 强制转 RGB,适配 ResNet if self.transform: img = self.transform(img) return img, label

然后在config.yaml中修改data.root_dir为"my_dataset/",并确保transform包含Resize((224,224))(ResNet 输入要求)或Resize((28,28))(CNN 输入要求)。

6.2 模型微调:冻结 backbone 与解冻策略表

当你用 ResNet 迁移学习时,冻结策略直接影响收敛速度。以下是针对不同数据规模的推荐:

数据量(训练样本)backbone 冻结策略fc 层初始化方式学习率建议
< 100全部冻结 (requires_grad=False)nn.Linear(512, num_classes)0.01
100–1000仅冻结layer1~layer3,解冻layer4同上0.001
> 1000仅冻结conv1+bn1+layer1同上0.0005

在src/models/resnet.py中添加解冻方法:

def unfreeze_layers(model, layers_to_unfreeze: list): """layers_to_unfreeze: ['layer4', 'fc']""" for name, param in model.named_parameters(): if any(layer in name for layer in layers_to_unfreeze): param.requires_grad = True else: param.requires_grad = False

6.3 推理加速:ONNX Runtime 的 Execution Provider 选择指南

ONNX Runtime 支持多种 Execution Provider(EP),选择不当会损失 50% 性能:

EP 名称适用场景安装命令注意事项
CUDAExecutionProviderNVIDIA GPU(推荐)pip install onnxruntime-gpu需 CUDA 11.2+,cuDNN 8.2+
TensorRTExecutionProviderNVIDIA GPU(极致性能)需单独编译 TensorRT,复杂延迟略高,吞吐最高
CPUExecutionProviderCPU 推理(通用)pip install onnxruntime默认启用,无需指定
DirectMLExecutionProviderWindows + AMD/NVIDIA 集成显卡pip install onnxruntime-directml仅 Windows,不支持 Linux

在推理代码中启用 CUDA EP:

import onnxruntime as ort # 替换原来的 session = ort.InferenceSession(...) providers = [ ('CUDAExecutionProvider', { 'device_id': 0, 'arena_extend_strategy': 'kSameAsRequested', }), 'CPUExecutionProvider' ] session = ort.InferenceSession("mnist_model.onnx", providers=providers)

从那以后我每次迁移新项目,都强制走一遍这三步:先用CustomImageDataset跑通数据加载,再用unfreeze_layers()试两种冻结策略,最后用onnxruntime-gpu对比 CPU 推理耗时。少走三个月弯路。
希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询