ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch CNN图像识别实战:从数据预处理到模型训练评估全流程

PyTorch CNN图像识别实战:从数据预处理到模型训练评估全流程 这次我们不看推理框架不折腾 Stable Diffusion直接回到深度学习里最常被问起的一条完整链路CNN PyTorch 做图像识别。很多刚入门的朋友手里有显卡也装了 PyTorch但真正要自己处理一套图像数据集、搭建卷积网络、训练到能评估效果时反而不知道从哪里下手。这篇文章就把“数据集处理 → CNN 模型搭建 → 训练 → 评估”全流程拆开讲清楚代码可以直接复制改到自己项目里跑。CNN卷积神经网络是图像识别最基础也最成熟的方案。PyTorch 则是目前研究和工程落地都绕不开的深度学习框架生态完整调试方便。两者搭配非常适合做图像分类、目标识别、风格迁移、OCR 预处理等视觉任务。本文的重点不是堆概念而是给出一套能落地的操作流程从原始图片文件夹开始做完数据清洗、统一尺寸、数据增强、数据集划分再用 PyTorch 搭建 CNN 模型完成训练和评估最后输出混淆矩阵、准确率、召回率等指标。全文涉及的核心内容包括PyTorch 环境准备、torchvision.datasets.ImageFolder数据集加载方式、DataLoader批量读取、CNN 卷积层/池化层/全连接层设计、训练循环、模型保存与加载、评估脚本、显存与 CPU 占用观察、常见报错排查。适合正在做课程设计、毕业设计或者想把手里的图片数据跑成一个分类模型的工程师阅读。1. 核心能力速览在用代码之前先把这套流程的关键信息做个概览方便你判断自己的电脑能不能跑、要准备哪些东西。能力项说明适用任务图像分类、图像识别、简单目标识别框架依赖PyTorch、torchvision、NumPy、Matplotlib、Pillow支持系统Windows 10/11、Ubuntu 18.04、macOSCPU 可跑硬件要求CPU 可完成训练速度慢推荐 NVIDIA 显卡显存 4G 起CUDA 支持可选没有 NVIDIA 显卡也能用 CPU 完成全流程数据集格式按类别分文件夹的图片目录结构数据增强随机翻转、旋转、裁剪、归一化训练方式自定义简单 CNN 或 torchvision 预训练模型微调评估指标Accuracy、Loss、Precision、Recall、Confusion Matrix模型导出PyTorchstate_dict可转 ONNX 部署API 支持本文不涉及服务化 API但训练后的模型可接入 FastAPI/Flask批量任务支持批量图片推理DataLoader 本身即批处理这里先说清楚本文给的代码不是某个一键包而是标准 PyTorch 工程化写法。你需要自己准备数据集自己执行训练脚本适合理解原理并在此基础上扩展。2. 适用场景与使用边界CNN PyTorch 这套组合最常用的场景有几类图像分类比如猫狗识别、花卉分类、产品瑕疵分类、垃圾图片分类。简单目标区域识别配合 OpenCV 做轮廓检测后对裁剪区域做分类。迁移学习用 ImageNet 预训练权重微调适配自己的小数据集。教学和实验验证卷积核、池化层、Dropout、BatchNorm 的效果。但也有明显边界。如果你的任务是密集目标检测比如检测图中每个人的位置单纯 CNN 分类模型不够需要 Faster R-CNN、YOLO 或 SSD。如果要做像素级分割需要 U-Net、DeepLab 等网络结构。CNN 分类模型解决的是“这张图属于哪个类别”的问题不是“图中每个物体在哪里”。使用边界方面有三点必须提醒数据版权与授权不要随意爬取他人图片做商用识别模型尤其是人脸、车牌、商品图等数据。涉及人脸识别、人物属性识别时必须遵守相关法律法规并获得数据授权。数据偏见训练集类别不平衡会导致模型对少数类别识别率低。项目上线前要检查每个类别样本量。测试集独立评估必须用训练时未参与梯度更新的测试集。很多人直接在训练集上算准确率得到一个虚高的分数这不具备参考价值。3. 环境准备与前置条件3.1 安装 Python 与虚拟环境建议使用 Anaconda 管理 Python 环境避免多个项目依赖冲突。Windows 下也可以用venv但 Anaconda 在切换 CUDA 版本时更方便。# 创建 Python 3.10 环境 conda create -n cnn-pytorch python3.10 -y conda activate cnn-pytorch3.2 安装 PyTorchCPU 版本安装最简单先跑通流程再考虑 GPU# CPU 版 pip install torch torchvision torchaudio如果你有 NVIDIA 显卡先到官网查看自己的 CUDA 版本再安装对应版本。这里不写死具体 CUDA 版本号因为 PyTorch 安装源会随版本更新最稳妥的方法是在 PyTorch 官方首页选择本机对应的命令。以 CUDA 12.x 为例安装指令格式通常如下# GPU 版具体 CUDA 版本号请以官网为准 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121安装完成后验证方式import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU mode)如果torch.cuda.is_available()返回False说明 PyTorch 安装成了 CPU 版或者显卡驱动与 CUDA 版本不匹配。后面第 8 节会专门讲排查。3.3 安装依赖库除了 PyTorch还需要图像处理、数值计算和可视化库。pip install numpy matplotlib pillow scikit-learn tqdmscikit-learn用来生成混淆矩阵、计算分类报告tqdm用来显示训练进度条非常实用。3.4 数据集准备本文以图像分类项目为例。数据集目录结构要求如下data/ ├── train/ │ ├── cat/ │ │ ├── cat_001.jpg │ │ ├── cat_002.jpg │ │ └── ... │ └── dog/ │ ├── dog_001.jpg │ └── ... └── test/ ├── cat/ └── dog/每个子文件夹的名字就是类别标签。ImageFolder会自动根据子文件夹名生成标签索引。如果你有自己的数据按照这个目录结构整理即可。4. 数据集处理与增强数据集处理是整个项目中最容易被低估的环节。很多模型训练效果差不是因为网络结构不行而是数据没处理好。4.1 数据清洗与基本检查拿到图片后先检查三件事图片是否损坏能不能正常打开。图片尺寸是否统一是否需要裁剪或缩放。类别样本量是否均衡。先用一段脚本快速检查import os from PIL import Image from collections import Counter data_root data/train class_counter Counter() broken_images [] for class_name in os.listdir(data_root): class_dir os.path.join(data_root, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): img_path os.path.join(class_dir, img_name) class_counter[class_name] 1 try: img Image.open(img_path) img.verify() except Exception as e: broken_images.append(img_path) print(类别样本数:, class_counter) print(损坏图片数量:, len(broken_images)) print(损坏图片示例:, broken_images[:5])如果类别严重不平衡比如猫有 5000 张、狗只有 300 张建议做数据扩充或类别加权。4.2 数据预处理流程PyTorch 官方推荐的预处理流程是统一尺寸 → 数据增强 → 转 Tensor → 归一化。from torchvision import transforms # 训练集包含数据增强增强泛化能力 train_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 测试集不做随机增强只做尺寸归一化和归一化 test_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这里的mean和std是 ImageNet 数据集的统计值。如果训练自己的小数据集可以使用这份通用归一化参数也可以重新计算数据集的均值和标准差。最终效果一般差别不大。4.3 使用 ImageFolder 加载数据from torchvision import datasets train_dataset datasets.ImageFolder(rootdata/train, transformtrain_transforms) test_dataset datasets.ImageFolder(rootdata/test, transformtest_transforms) print(训练集类别:, train_dataset.classes) print(训练集样本数:, len(train_dataset)) print(测试集样本数:, len(test_dataset))ImageFolder会按照文件夹名排序生成标签比如cat的索引为 0dog的索引为 1。4.4 数据集划分如果你只有一个总文件夹没有单独划分 train/test可以先用torch.utils.data.random_split按比例切分或者用shutil移动文件。推荐用random_split简单且随机性可控from torch.utils.data import random_split total_len len(full_dataset) train_len int(total_len * 0.8) val_len int(total_len * 0.1) test_len total_len - train_len - val_len train_dataset, val_dataset, test_dataset random_split( full_dataset, [train_len, val_len, test_len] )注意random_split切分后每个子数据集仍使用原始图像的索引如果你在切分前已经对数据做了统一预处理没问题。但如果原始数据集图片尺寸差异很大建议先统一目录格式再切分避免训练图片有黑边。4.5 DataLoader 批量加载from torch.utils.data import DataLoader batch_size 32 train_loader DataLoader( datasettrain_dataset, batch_sizebatch_size, shuffleTrue, num_workers0, drop_lastTrue ) test_loader DataLoader( datasettest_dataset, batch_sizebatch_size, shuffleFalse, num_workers0 )shuffleTrue是训练集必须配置的避免模型学习到排列顺序。num_workers在 Windows 下容易触发多进程问题为了方便调试可以先设置为 0。如果你的机器内存足够Linux 下可以设置成 4 或 8 加速数据读取。5. CNN 模型搭建CNN 模型可以自己从零搭建也可以使用 torchvision 里现成的预训练模型。前者适合理解原理后者适合工程落地时快速拿到更高准确率。5.1 自定义 CNN从零开始以经典的卷积块结构为例一个简单的 CNN 分类模型包含卷积层 → 激活函数 → 池化层 → Dropout → 全连接层。import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self, num_classes2): super().__init__() self.features nn.Sequential( # 第一块3 通道输入16 个卷积核 nn.Conv2d(in_channels3, out_channels16, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), # 第二块 nn.Conv2d(in_channels16, out_channels32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), # 第三块 nn.Conv2d(in_channels32, out_channels64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), ) self.classifier nn.Sequential( nn.Dropout(p0.3), nn.Linear(in_features64 * 28 * 28, out_features128), nn.ReLU(inplaceTrue), nn.Linear(in_features128, out_featuresnum_classes) ) def forward(self, x): x self.features(x) x x.view(x.size(0), -1) x self.classifier(x) return x这段代码的输入图片尺寸是 224×224经过三次 MaxPool2d 池化特征图尺寸变为 224 / 8 28因此全连接层的输入维度是 64 × 28 × 28。如果你换了输入尺寸这里要同步改。5.2 使用预训练模型迁移学习如果你的数据集有几千张图从头训练 CNN 效率不高。更常见的是加载 ImageNet 预训练权重只替换最后的全连接层import torchvision.models as models def create_resnet18(num_classes2): model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT) in_features model.fc.in_features model.fc nn.Linear(in_featuresin_features, out_featuresnum_classes) return model迁移学习的优势很明显预训练模型已经学到了大量底层视觉特征即使你的数据集只有每类几百张也能得到不错的效果。如果你的显卡显存不大resnet18是性价比很高的选择。5.3 模型与设备初始化import torch device torch.device(cuda if torch.cuda.is_available() else cpu) model create_resnet18(num_classes2) model model.to(device) print(模型加载完成运行设备:, device)到这里数据集和模型都准备好了接下来是训练。6. 模型训练全流程6.1 定义损失函数与优化器import torch.nn as nn import torch.optim as optim criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) scheduler optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1)CrossEntropyLoss是分类任务的标准损失函数内部已经包含 Softmax。Adam优化器适合 CV 任务收敛速度比 SGD 更快但精度上传统 SGD 配合调节后的学习率有时更好。可以先从 Adam 起步后期再换 SGD 调优。StepLR每 10 个 epoch 学习率缩小 10%帮助模型收敛到更优区域。6.2 一个完整的训练循环from tqdm import tqdm num_epochs 30 for epoch in range(num_epochs): model.train() running_loss 0.0 correct 0 total 0 loop tqdm(train_loader, descfEpoch [{epoch1}/{num_epochs}]) for images, labels in loop: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() * images.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() loop.set_postfix(lossloss.item()) epoch_loss running_loss / total epoch_acc correct / total print(fEpoch {epoch1}: Loss {epoch_loss:.4f}, Acc {epoch_acc:.4f}) scheduler.step()每次迭代的关键步骤是取一批图片和标签。前向传播计算输出和损失。optimizer.zero_grad()清空上一轮梯度。loss.backward()反向传播。optimizer.step()更新参数。6.3 验证集评估每个 epoch 结束后在验证集上计算准确率以便判断模型是否过拟合def evaluate(model, dataloader, device): model.eval() correct 0 total 0 total_loss 0.0 with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) total_loss loss.item() * images.size(0) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() avg_loss total_loss / total accuracy correct / total return avg_loss, accuracy在训练循环中调用val_loss, val_acc evaluate(model, val_loader, device) print(f验证集损失: {val_loss:.4f}, 验证集准确率: {val_acc:.4f})6.4 模型保存与加载训练完成后保存模型的state_dict。这是推荐做法比整个保存模型更省空间也更容易版本管理torch.save(model.state_dict(), model_cnn.pth)后续加载model create_resnet18(num_classes2) model.load_state_dict(torch.load(model_cnn.pth, map_locationdevice)) model.to(device) model.eval()注意load_state_dict必须保证模型的网络结构和你保存时完全一致否则会报 key 不匹配错误。如果后续修改了网络结构必须重新保存。7. 模型评估与效果验证7.1 测试集评估test_loss, test_acc evaluate(model, test_loader, device) print(f测试集损失: {test_loss:.4f}) print(f测试集准确率: {test_acc:.4f})7.2 混淆矩阵准确率有时会骗人。比如一个二分类问题上类别不平衡A 类占 90%那全猜 A 也有 90% 准确率。所以要输出混淆矩阵import numpy as np import matplotlib.pyplot as plt import torch from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay all_labels [] all_preds [] model.eval() with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) all_labels.extend(labels.cpu().numpy()) all_preds.extend(predicted.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) disp ConfusionMatrixDisplay(confusion_matrixcm, display_labelstrain_dataset.classes) disp.plot(cmapBlues) plt.title(Confusion Matrix) plt.savefig(confusion_matrix.png, dpi200) plt.show()从混淆矩阵可以看出模型具体把哪个类别识别错了。比如猫被误判成狗的数量很多说明这两个类别之间存在视觉混淆需要补充更多差异化的训练数据或调整网络容量。7.3 单张图片推理测试训练完模型之后最直观的验证方法是随便拿一张不在训练集里的图片来做推理from PIL import Image def predict_image(image_path, model, class_names, transform, device): img Image.open(image_path).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) model.eval() with torch.no_grad(): outputs model(img_tensor) probs torch.softmax(outputs, dim1) confidence, predicted_idx torch.max(probs, 1) class_name class_names[predicted_idx.item()] confidence confidence.item() return class_name, confidence class_names train_dataset.classes result, conf predict_image(data/test/cat/test_001.jpg, model, class_names, test_transforms, device) print(f预测结果: {result}, 置信度: {conf:.4f})7.4 批量图片推理如果你的需求是批量预测几百张图片可以直接用DataLoader遍历全部图片并保存结果到 CSVimport csv results [] model.eval() with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) for idx in range(images.size(0)): results.append({ true_label: class_names[labels[idx].item()], pred_label: class_names[predicted[idx].item()] }) with open(predict_results.csv, w, newline, encodingutf-8) as f: writer csv.DictWriter(f, fieldnames[true_label, pred_label]) writer.writeheader() writer.writerows(results)8. 接口 API 调用示例虽然本文的训练脚本没有内置 Web 服务但训练好的模型要接到业务系统里一般是通过 FastAPI 或 Flask 包装成 HTTP 接口。这里给出一个通用模板你只需替换模型加载路径和预处理逻辑即可。import io import torch from PIL import Image from torchvision import transforms from fastapi import FastAPI, UploadFile, File app FastAPI() device torch.device(cuda if torch.cuda.is_available() else cpu) model create_resnet18(num_classes2) model.load_state_dict(torch.load(model_cnn.pth, map_locationdevice)) model.to(device) model.eval() class_names [cat, dog] transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) app.post(/predict) async def predict(file: UploadFile File(...)): img_bytes await file.read() img Image.open(io.BytesIO(img_bytes)).convert(RGB) img_tensor transform(img).unsqueeze(0).to(device) with torch.no_grad(): outputs model(img_tensor) probs torch.softmax(outputs, dim1) confidence, pred_idx torch.max(probs, 1) return { class: class_names[pred_idx.item()], confidence: round(confidence.item(), 4) }启动接口服务uvicorn api_server:app --host 0.0.0.0 --port 8000用 curl 测试curl -X POST http://127.0.0.1:8000/predict \ -F filedata/test/cat/test_001.jpg返回示例{ class: cat, confidence: 0.9821 }接口服务可以自己控制但要注意不要直接暴露到公网默认绑定127.0.0.1通过内网或反向代理访问更安全。9. 资源占用与性能观察9.1 显存占用怎么看训练过程中显存占用主要来自三个方面模型参数、优化器状态、中间激活值。其中中间激活值最占显存和 Batch Size、输入分辨率成正比。用nvidia-smi实时观察nvidia-smi -l 1-l 1表示每秒刷新一次。训练时可以看到python进程的显存占用。如果显存不够优先降低batch_size比如从 32 降到 16 或 8。9.2 CPU 推理和 GPU 推理的差异GPU 推理通常比 CPU 快很多但小模型和单张图片推理时GPU 的优势不一定明显因为数据传输和初始化有开销。如果你的使用场景是批量离线预测GPU 优势明显如果是单张图片即时响应CPU 也够用。CPU 推理时可以给模型加torch.no_grad()和model.eval()减少内存占用和冗余计算。9.3 如何降低显存占用降低batch_size比如从 64 降到 32、16。降低输入图片尺寸比如从 224×224 降到 160×160。使用torch.cuda.amp混合精度训练。混合精度简单用法from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() with autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()注意混合精度只在支持 CUDA 的环境中有效。如果你用的显卡不支持会直接报错。9.4 训练速度不理想怎么办先判断瓶颈在哪里如果 GPU 利用率低但 CPU 跑满说明数据加载太慢增大num_workers。如果 GPU 利用率高但速度仍然慢考虑换更小的网络结构或使用预训练模型冻结前几层。如果显存不足导致 OOM降低batch_size是第一步。10. 常见问题与排查方法问题现象可能原因排查方式解决方案torch.cuda.is_available()返回 FalseGPU 驱动不对或 PyTorch 为 CPU 版检查nvidia-smi输出检查 PyTorch 版本重装对应 CUDA 版本的 PyTorch运行时报CUDA out of memoryBatch Size 过大或输入分辨率过高查看nvidia-smi显存占用降低 batch_size降低图片分辨率FileNotFoundError: No such file or directory数据集路径不对检查当前工作目录和数据集路径用绝对路径或cd到项目根目录训练 Loss 不下降学习率过高或过低、数据标签错误打印前 10 个 batch 的 loss调低学习率到 0.0001检查 DataLoader 的标签验证集准确率远低于训练集过拟合对比训练集和验证集准确率差距增加 Dropout、减小模型容量、增加数据增强加载模型报 key 不匹配网络结构不一致打印state_dict的 key确保模型结构相同重新保存Windows 下 DataLoader 卡住num_workers多进程问题观察进程是否卡死设置num_workers0打开图片报PIL.UnidentifiedImageError图片损坏或不是标准图片格式检查文件后缀和文件头清洗数据删除损坏图片11. 最佳实践与使用建议11.1 先小规模跑通再全量训练第一次训练时不要直接上全部数据和 100 个 epoch。使用 10% 的数据、10 个 epoch 跑一遍确认代码没有 bug再全量训练。这样可以节省大量排错时间。11.2 固定随机种子为了让实验可复现训练前固定随机种子import random import numpy as np def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)11.3 训练日志与模型版本管理建议每次训练保存一份日志文件名带上日期和关键参数logs/ ├── 20250120_resnet18_lr0.001_bs32/ │ ├── train.log │ ├── model_best.pth │ └── confusion_matrix.pngmodel_best.pth可以在每个 epoch 后比较验证集准确率验证集最好时保存一次best_acc 0.0 for epoch in range(num_epochs): # 训练代码... val_loss, val_acc evaluate(model, val_loader, device) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), model_best.pth)11.4 数据和模型文件分目录管理推荐的项目目录结构project/ ├── data/ │ ├── train/ │ └── test/ ├── models/ │ ├── model_best.pth │ └── model_cnn.pth ├── outputs/ │ ├── confusion_matrix.png │ └── predict_results.csv ├── scripts/ │ ├── train.py │ ├── evaluate.py │ └── predict.py └── README.md11.5 商用前必须做的合规检查训练数据是否有合法来源和授权。是否涉及人脸、车牌等敏感信息。模型预测结果是否会被用于自动决策是否会造成偏见或误判。这些都是模型上线前要回答的问题不是可选项。涉及人脸识别、声音属性识别等高敏场景时更要严格评估数据来源和算法公平性。12. 总结与下一步这篇文章从数据目录整理开始完整走了一遍 CNN PyTorch 图像识别的工程流程数据集检查、预处理、数据增强、ImageFolder加载、DataLoader 批量读取、自定义 CNN 搭建、预训练模型迁移学习、训练循环、验证集评估、模型保存、混淆矩阵分析和接口封装。如果你是本项目的新手建议按顺序做三件事准备一个小的分类数据集每类 50 张图先跑通训练脚本。把预训练模型从resnet18换成resnet34或mobilenet_v3_small观察准确率和显存占用变化。把评估脚本里的混淆矩阵和分类报告跑通真正理解模型在每个类别上的表现。最容易踩的坑有两个一是数据集路径和标签不匹配二是显卡驱动和 PyTorch 版本不匹配导致 CUDA 不可用。前者在数据检查阶段解决后者在安装验证阶段排查。下一步可以继续扩展的方向包括模型转为 ONNX 部署、加入 TensorRT 加速、引入更复杂的数据增强策略MixUp、CutMix、使用torchvision.models中更先进的网络结构做模型对比实验。建议收藏备用后面训练自己的数据集时可以直接拿这套代码改。
返回列表