ARTICLE DETAIL

资讯详情

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

PyTorch实战CIFAR-10图像分类:从Kaggle数据到ResNet提交全流程

PyTorch实战CIFAR-10图像分类:从Kaggle数据到ResNet提交全流程 简介面向深度学习初学者的实战Kaggle图像分类资源包围绕CIFAR-10数据集使用PyTorch实现完整竞赛流程。压缩包共1017个文件大小仅2.34MB1006张PNG图像是CIFAR-10数据样例可直接观察各类别特征4个Python脚本与2个Jupyter Notebook实现数据加载、模型训练与预测提交3个CSV文件包含标签和提交记录另有少量pyc缓存整体流程完整且代码组织清晰。已有180人学习适合希望在真实竞赛环境中从代码入手理解图像分类的读者。资源内含两版可运行的Notebook对照checkpoint可查看代码演进与提交差异训练脚本与submission.csv提交样例能帮助入门者快速对齐数据格式补全“训练-预测-提交”的完整链路建议首次使用时按Notebook顺序执行从而逐步掌握数据加载、模型构建与预测输出的关键写法。整体内容紧凑可直接结合Kaggle赛题页面边练边学是实践CIFAR-10图像分类任务的便捷素材。1. 从零到第一份提交CIFAR-10 图像分类为什么值得用 PyTorch 打一轮先把结论放这儿用 PyTorch 打一轮 Kaggle 的 CIFAR-10 图像分类就是那个叫 CIFAR-10 - Object Recognition in Images 的经典赛大概是入门阶段投入产出比最高的实战路线。图只有 32×32一张普通消费级显卡就训得动5 万张训练数据也远没到要上分布式、调显存的地步但一个完整竞赛流程该踩的坑它几乎全都有——Kaggle 注册、数据集下载、标签映射、数据增强、训练循环、提交格式一条龙走完你对「用框架做项目」的理解会和只跑 MNIST 完全不一样。这适合刚学完 PyTorch 基础框架、想上一个有排行榜反馈的实战项目的人也适合想在简历里放一个可量化结果的一线开发。我会顺着一条可复现的路径讲数据怎么进到模型、训练循环怎么写、参数为什么这样设、提交文件怎么对以及实际踩过的坑。照着做能跑通想改参数也知道改哪里。2. 数据闭环Kaggle 注册、下载与 32×32 图片的 Dataset 封装2.1 Kaggle 注册与数据下载验证码、CLI 和两种姿势先把数据长什么样说清楚。这个比赛里train 目录下有 5 万张 32×32 的 PNG按 id 命名trainLabels.csv 里是 id 到类别名的映射类别一共 10 个airplane、automobile、bird、cat、deer、dog、frog、horse、ship、truck。test 目录下有 30 万张无标注图片提交时按 id 顺序给出预测的类别字符串。这和 torchvision 内置的那份 CIFAR-10 不一样——内置版已经按 5 万/1 万切好且都带标签而竞赛版要自己处理一整条「无标注推断」链路。想拿到排行榜分数就得用 Kaggle 竞赛官网的那份数据。注册这一步我见过太多人卡在验证码上页面加载出来了验证码区域却是空白的点提交就报「captcha must be filled out」。这多半不是 Kaggle 的问题而是浏览器广告拦截扩展把验证码所在的 iframe 一起拦掉了。常见做法是关掉广告拦截扩展、改用普通窗口而不是隐身模式、换一个浏览器再试还不行就用手机浏览器注册一次注册通过之后就不再需要页面验证码了。账号搞定后推荐用 CLI 下载而不是在浏览器里等 zip# 1. 在当前 Python 环境里装 Kaggle 官方命令行工具 pip install kaggle # 2. 配置 API 密钥kaggle.com 个人页 - Settings - API - Create New Token mkdir -p ~/.kaggle mv ~/Downloads/kaggle.json ~/.kaggle/ chmod 600 ~/.kaggle/kaggle.json # 3. 先打开比赛页面点 Join Competition 并接受规则否则下面这步会 403 kaggle competitions download -c cifar-10 unzip -q cifar-10.zip -d cifar-10/说明一下-c cifar-10是比赛标识符去比赛页看 URL 最后一段就是这个值别自己改。chmod 600是让 CLI 不抱怨存放权限Windows 用户如果用了 Git Bash 同样需要这一步。如果你习惯在 Kaggle Notebook 里直接干活那可以改用 kagglehub 接口跑下载路径逻辑会不太一样但背后的数据是一样的。我一般推荐本地 CLI原因是测试集有 30 万张图解压后占几个 GB本地存一份后面反复做验证和推断都方便。环境上顺便说一下我用 Anaconda 建独立环境避免污染 baseconda create -n cifar python3.10 -y conda activate cifar pip install torch torchvision pandas pillow参数说明Python 3.10 以上和当前版本 PyTorch 的兼容性都很好不必追最新torch 和 torchvision 必须配套安装直接一条pip install torch torchvision默认拉取匹配版本最不容易出错。GPU 版别靠 pip 默认源猜去 PyTorch 官网复制对应 CUDA 版本的安装命令装完再验证。2.2 目录结构与数据校验训练前先数清文件解压完成后先别急着写模型把目录整理成下面这样后面所有脚本的路径都不会乱data/cifar-10/ ├── train/ # 50000 张 ├── test/ # 300000 张 ├── trainLabels.csv └── sampleSubmission.csv然后做一次文件计数这是整个项目里最便宜的一道保险import os train_dir data/cifar-10/train test_dir data/cifar-10/test print(train 图片数:, len(os.listdir(train_dir))) # 期望 50000 print(test 图片数:, len(os.listdir(test_dir))) # 期望 300000如果 test 目录数出来不是 30 万多半是解压时漏了文件或者下载的版本不对。这个检查必须在训练前做因为等训练完才发现测试集少了一万张提交时平台直接报行数不符整个训练周期就白费了。注意test 目录必须正好 30 万张多一张少一张都先别进训练环节先搞清楚数据哪来的。顺带看sampleSubmission.csv的格式通常两列id、labelid 不带.png后缀label 是类别字符串。后面生成提交文件时严格按这个格式对齐别自己发明列名。2.3 写一个 Dataset 类把 CSV 和图片文件合成模型的输入现在把「一张一张 PNG 一张标签表」变成一个 Dataset。我一般这么写import pandas as pd from PIL import Image from torch.utils.data import Dataset class KaggleCIFAR10(Dataset): def __init__(self, csv_path, img_dir, transformNone): self.df pd.read_csv(csv_path) # 两列: id, label self.img_dir img_dir self.transform transform self.label_list sorted(self.df[label].unique()) # 10 个类别 self.label2idx {name: i for i, name in enumerate(self.label_list)} def __len__(self): return len(self.df) def __getitem__(self, idx): row self.df.iloc[idx] img Image.open(f{self.img_dir}/{row[id]}.png).convert(RGB) label self.label2idx[row[label]] if self.transform: img self.transform(img) return img, label逻辑说明label2idx从整张表取唯一类别名并排序保证 10 个类稳定映射到 09且每次运行顺序一致。这里有个隐藏收益——训练前打印一次self.label_list和 torchvision 内置 CIFAR-10 的类别顺序对比一下。理论上两边都按字母序天然一致但这个动作花十秒能省掉后面「标签错位」一整个排查周期。Image.open读出来是 PIL 对象convert(RGB)是为防止个别图片是灰度模式导致通道数不一致。transform 放在__getitem__里执行而不是提前一次性做完因为增强必须每次采样随机执行提前做完就失去意义了。容易忽略的一点测试集没有 csv只有图片文件所以推断时不能用这个 Dataset。我会单独写一个只读图片的 Dataset或者直接包一层返回img_id和图像这个到第 4 章推断部分再展开。3. 模型与训练PyTorch 实战 CIFAR-10 的 ResNet-18 基线3.1 模型选型为什么 ResNet-18 是性价比之王CIFAR-10 的图只有 32×32这个尺寸决定了选模型的第一原则别把 ImageNet 那套 224×224 输入的模型直接硬搬。常见做法是用 ResNet 系把第一个 7×7 卷积换成 3×3、去掉开头的大池化让网络在 32×32 上正常工作。torchvision 里的resnet18(num_classes10)虽然默认按 224 设计但在 CIFAR-10 上配合增强直接跑也能到 88%90%想再进一步换 ResNet-34 或宽残差WideResNet都行代价是训练时间变长。有人会问现在最新的图像分类模型不是 Swin、ViT 这些 Transformer 系吗确实能在 CIFAR-10 上跑到 95% 以上但训练轮数、学习率、增强策略的调参复杂度都明显上升。5 万张 32×32 的小图CNN 的归纳偏置占尽便宜Transformer 靠数据量堆出来的优势在这里发挥不出来。我的建议很直接第一版永远用 ResNet 这类卷积网络做基线先把流程跑通、把提交分数拿到手再考虑要不要换更贵的模型。基线的作用是给你一个「正常水平」的参照之后所有改动都跟它比这比一上来就追大模型有意义得多。3.2 最小训练脚本一屏能看完的训练循环模型、数据、优化器凑齐之后训练循环其实很短。下面这版先用 torchvision 内置数据做说明流程更少真实比赛里把 2.3 节的 Dataset 换进来其余完全不用动import torch import torch.nn as nn import torchvision from torchvision import datasets, transforms from torch.utils.data import DataLoader transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_val transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) train_ds datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) val_ds datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_val) train_loader DataLoader(train_ds, batch_size128, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size256, shuffleFalse, num_workers4) model torchvision.models.resnet18(num_classes10) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.1, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max60) model model.cuda() for epoch in range(60): model.train() running_loss 0.0 for x, y in train_loader: x, y x.cuda(), y.cuda() optimizer.zero_grad() loss criterion(model(x), y) loss.backward() optimizer.step() running_loss loss.item() scheduler.step() if (epoch 1) % 10 0: print(fepoch {epoch1}, loss {running_loss / len(train_loader):.4f})参数说明SGD 配 lr0.1、momentum0.9、weight_decay5e-4 是 CIFAR-10 上被验证过很多次的组合weight_decay 对每层参数做 L2 正则能明显压过拟合CosineAnnealingLR 让学习率在 60 个 epoch 里从 0.1 余弦降到接近 0前期大步快走、后期小幅精修比固定学习率稳定多 12 个点。batch_size128 在单卡上足够4GB 显存就能跑num_workers 按 CPU 核数一半设太小会饿着 GPU太大反而因为进程调度变慢。参数整理成表会更直观参数取值作用lr0.1配合余弦退火的大初始学习率momentum0.9SGD 动量加快收敛weight_decay5e-4L2 正则防止过拟合batch_size12832×32 小图单卡友好T_max60必须等于 epoch 数否则余弦曲线走不完另外提一句如果发现前 5 个 epoch 的 loss 不降先查两个地方——归一化 stats 是否匹配、输入是否真的过了 ToTensor。这两处错了换什么模型都救不回来。3.3 数据增强参数CIFAR-10 的四个默认增强CIFAR-10 的增强有非常成熟的配方我的默认配置长这样from torchvision import transforms train_tf transforms.Compose([ transforms.RandomCrop(32, padding4), # 四周补 4px再随机裁回 32×32 transforms.RandomHorizontalFlip(p0.5), # 一半概率水平翻转 transforms.ToTensor(), # 从 0~255 的 uint8 变成 0~1 的 float32 transforms.Normalize((0.4914, 0.4822, 0.4465), # CIFAR-10 全量 RGB 均值 (0.2470, 0.2435, 0.2616)), # CIFAR-10 全量 RGB 标准差 transforms.RandomErasing(p0.5, scale(0.02, 0.33)), # Cutout 变体随机擦除一块 ])每个参数都有讲究padding4的意思是把 32×32 的图四周各补 4 像素成 40×40再随机裁 32×32 回来等价于让模型见到轻微平移的图这是这个数据集上收益最高的增强之一翻转概率 0.5 是对称性标准设置不要随便调大Normalize 用的是 CIFAR-10 自己的统计量不是 ImageNet 那组——用 ImageNet 的均值去归一化 01 范围的数据输入分布会整体偏移轻则收敛慢重则 loss 卡住。RandomErasing 是后期加的类似 Cutout强迫模型不依赖单一特征块但如果训练轮数少于 40建议先关掉它因为它在数据少时会引入额外噪声反而掉点。4. 提交与涨点生成 submission.csv 的三个稳定动作4.1 推断与提交文件30 万张图怎么在本地跑完训练完成后要把 test 目录 30 万张图全部过一遍模型生成提交文件。这里和训练时有两处关键差异没有标签、文件量大所以不能复用训练 Dataset推断循环也要写成块状别一张一张读import os import pandas as pd import torch from PIL import Image model.eval() test_dir data/cifar-10/test test_files sorted(os.listdir(test_dir)) # 必须排序保证 id 顺序稳定 label_names [airplane, automobile, bird, cat, deer, dog, frog, horse, ship, truck] # 与训练 label2idx 对应 pred_ids, pred_labels [], [] with torch.no_grad(): for i in range(0, len(test_files), 256): batch [] for f in test_files[i:i 256]: img Image.open(os.path.join(test_dir, f)).convert(RGB) batch.append(transform_val(img)) x torch.stack(batch).cuda() out model(x) # 水平翻转做一次 TTA两次前向取平均稳定涨 0.2~0.4 个点 out (out model(torch.flip(x, dims[3]))) / 2 probs torch.softmax(out, dim1) preds probs.argmax(dim1).cpu().tolist() for f, p in zip(test_files[i:i 256], preds): pred_ids.append(f.split(.)[0]) pred_labels.append(label_names[p]) sub pd.DataFrame({id: pred_ids, label: pred_labels}) sub.to_csv(submission.csv, indexFalse) print(sub.shape) # 期望 (300000, 2)逻辑说明sorted保证 test 文件按 id 稳定排列不然每次生成的提交文件顺序都会变torch.flip(x, dims[3])是水平翻转把原图和翻转图的输出概率取平均这就是测试时增强TTA对 ResNet 这类模型通常能稳定涨一点argmax拿到的序号必须用和训练时同一个label2idx的反向映射转成类别字符串这里是最容易出错的一环。注意一个问题30 万张图在本地逐张推断可能要几十分钟到几小时。如果觉得慢常见做法是把模型转 ONNX再用推理引擎跑 FP16或者把 batch 调大到 512 减少 Python 循环开销。但我的建议是——第一个提交先用最简单的版本拿分数确认流程没问题再考虑优化速度。一个能打的慢提交比一个跑飞快的废提交有用得多。4.2 三个稳定涨点动作标签平滑、余弦退火与模型平均基线跑到 90% 上下之后下面三个动作是我每次必做的都不动网络结构风险低、收益稳。第一个是标签平滑。CrossEntropyLoss 默认把正确类当 1、其他当 0模型容易过于自信。标签平滑把「1」改成「1-ε」并把 ε 均匀分给其他类能明显改善泛化class LabelSmoothingCE(nn.Module): def __init__(self, smoothing0.1): super().__init__() self.smoothing smoothing def forward(self, logits, target): log_probs torch.log_softmax(logits, dim-1) n logits.size(-1) true_dist torch.full_like(log_probs, self.smoothing / (n - 1)) true_dist.scatter_(1, target.unsqueeze(1), 1.0 - self.smoothing) return (-true_dist * log_probs).sum(dim-1).mean()smoothing0.1 是 CIFAR-10 上常见的取值太大0.3 以上会让训练 loss 下不去太小0.01等于没加。第二个是余弦退火脚本里已经有了但这里要强调T_max必须等于 epoch 数而不是随手填。如果 T_max60 而只训 30 个 epoch学习率只走完半条余弦曲线相当于没退火完反过来 T_max 大于 epoch 数学习率降不到最低点。想多涨一点训完一整轮后把学习率重置再训一轮这就是 warm restart 思路对 CIFAR-10 效果很稳。第三个是模型平均最稳的涨点方式。不用真的训多个模型把同一份训练里不同 epoch 的 checkpoint 拿出来对输出概率取平均即可models [model_epoch_40, model_epoch_50, model_epoch_60] final_prob sum(torch.softmax(m(x), dim1) for m in models) / len(models)参数说明参与平均的 checkpoint 之间差异要够大才有收益一般每隔 10 个 epoch 存一个挑验证集准确率最高的 35 个全用同一训练阶段的会没变化。这三个动作叠一起从 90% 推到 93%94% 是可以预期的。5. 避坑手册Kaggle 图像分类最容易翻车的 5 个环节到提交为止完整链路都走通了但真实比赛里大部分人不是输给模型而是输给边界环节。按我的经历排序注册、下载、环境、推理、训练每个环节都有一个高频翻车点先看总表环节典型现象一句话原因注册验证码空白报 captcha must be filled out广告拦截扩展拦了验证码 iframe下载CLI 报 403 Forbidden没在比赛页接受规则环境torch.cuda.is_available() 为 Falsewheel 不带对应 GPU 后端推理本地 92%提交只有 40 多标签映射顺序错位训练loss 不降或变 nan归一化用错或没做 ToTensor下面每条按「现象 → 原因 → 解决」拆开讲。5.1 账号与下载验证码消失和 403坑 1Kaggle 注册时验证码不显示点提交就报「captcha must be filled out」。现象页面上验证码区域一片空白或一直转圈手机浏览器能显示但网页端不行。原因验证码是第三方 iframe 嵌入的浏览器的广告拦截扩展或隐私模式把它当广告拦掉了少数情况是浏览器自动填充干扰了勾选框。解决逐个关掉浏览器扩展优先广告拦截类用普通窗口而非隐身模式还不行就换手机浏览器完成注册。注册只需要一次之后用 API token 操作不再需要页面验证码。坑 2命令行下载kaggle competitions download -c cifar-10报 403 Forbidden。现象API key 配好了kaggle competitions list能列比赛但 download cifar-10 就 403。原因这个比赛需要先接受规则。CLI 按你账号的权限下载你从没在浏览器里点过 Join Competition / 接受规则服务端就认为无权访问数据。解决浏览器打开比赛页点 Join Competition勾选并接受规则必要时重新生成一次 API token。之后重跑下载命令即可。5.2 环境与装包GPU 识别和版本错位坑 3torch.cuda.is_available()返回 False训练悄悄跑在 CPU 上。现象装完 PyTorch 后一切正常train 也跑得动但一个 epoch 要好几分钟GPU 占用始终为 0另一类是 AMD 卡比如 7900XTX在 WSL2 里装官方预编译包torch 直接回退 CPU 或报 kernel 不匹配。原因绝大多数是装了 CPU 版 wheelpip 在某些源上默认把 torch 解析成不带 CUDA 的版本AMD 卡在 WSL2 下需要 ROCm 后端官方 CUDA wheel 识别不了。解决装完先跑一段诊断python -c import torch; print(torch.__version__, torch.cuda.is_available())。版本号里能看到cu118、cu121这类后缀才说明是 GPU 版AMD/WSL2 环境按官方 ROCm 安装指引装配套 wheelROCm 和 WSL 内核版本有对应关系装完要重启 WSL 才生效。我的做法是本地环境不确定时先把训练脚本在 Kaggle Notebook 的免费 GPU 上跑通拿分数本地再慢慢折腾 GPU这样不耽误比赛进度。5.3 数据与推理标签错位和归一化用错坑 4本地验证集 92%提交到排行榜只剩百分之四十几。现象训练过程一切正常验证准确率不低推断脚本也跑完生成了 submission.csv提交后分数低得离谱甚至接近随机猜。原因标签映射错位。训练时类别字符串变成数字推断时数字再变回字符串两处用了两个不一样的顺序列表——比如训练用sorted(label_list)推断时手写了另一种顺序。模型没学错是输出对错了标签。解决把训练label2idx的列表序列化保存推断时原样复用不要手敲提交前用 sampleSubmission 的前 50 行跑一次预测肉眼核对类别字符串和图片内容对不对得上。这个 30 秒的核对能救回一整次提交窗口。坑 5训练 loss 停在 2.3 附近不动或者第一轮就变 nan。现象训练开始后 loss 非常平稳地不下降或者某个 epoch 直接变 nan之后再也回不来。原因归一化参数用错。最常见是把 ImageNet 的 (0.485, 0.456, 0.406) 当 mean 用到 CIFAR-10 上输入分布整体偏移模型学不到东西另一类是 transform 里忘写 ToTensor把 0255 的 uint8 图直接送到卷积层数值范围过大把梯度炸掉。解决用 CIFAR-10 自己的统计量 (0.4914, 0.4822, 0.4465) 和 (0.2470, 0.2435, 0.2616)并保证 ToTensor 在 Normalize 前面。排查时从 train_loader 拿一个 batch 打印 x.min() 和 x.max()有效输入范围应该是 -3 到 3 之间如果打印出来是 0255ToTensor 一定漏了。6. 提交通关的最后一道校验本地验证集与 10 分钟核对6.1 先切一个本地验证集最后一个实用技巧不是调模型而是「别拿测试集当验证集用」。比赛版数据没有官方验证集很多人直接在 train 上训完就对着 test 推断然后赌一次提交。我的做法是先切出一份本地验证集import pandas as pd from sklearn.model_selection import train_test_split df pd.read_csv(data/cifar-10/trainLabels.csv) train_df, val_df train_test_split( df, test_size5000, stratifydf[label], random_state42 ) train_df.to_csv(train_split.csv, indexFalse) val_df.to_csv(val_split.csv, indexFalse)用stratify按类别等比例切保证 10 个类在验证集里各 500 张。这 5000 张就是你的「本地排行榜」每个模型改动都拿它打分比盲提交快得多。注意切完训练集变成 4.5 万张训练脚本里的 csv 路径要指向train_split.csv别再用原始的 trainLabels.csv否则验证集就漏进训练里了。6.2 提交前的三个核对提交文件生成之后我固定会做三件事顺序不能乱第一sub.shape必须等于 (300000, 2)多一行少一行都是废文件第二id 列和sampleSubmission.csv的 id 前 10 个比对确认没有.png后缀差异第三label 列做一次value_counts()10 个类都要出现且数量别出现极端值——比如某一个类超过一半那基本是标签错位了。这三步加一起不到 10 分钟但能挡掉我犯过的几乎所有低级错误。打了这几年比赛我养成的习惯是宁可错过一个提交窗口也不交一份格式存疑的文件每轮改动先跑本地验证再用同一份代码生成提交。这套流程笨但稳。CIFAR-10 这场打完之后你会发现这些坑在别的图像分类比赛里几乎原样重演到时候照着这份排查单走一遍就行。希望帮到你。本文还有配套的精品资源点击获取
返回列表