CNN 真的更抗遮挡吗?把训练增强和测试鲁棒性分开验证
本文基于 2026 年 9 月 25 日在本机完成的真实实验整理,使用 AI 辅助编写实验脚本与文稿。硬件为 Apple M3 Pro、36GB 内存,使用 Conda
zhcho环境与 PyTorch MPS。四组对照、三个训练种子共 12 次训练均已完成,本文数字来自保存的预测结果。
在上一篇《卷积神经网络的引入4——局部扰动与空间结构破坏下的鲁棒性验证》中,我通过随机擦除来观察 MLP 和 CNN 的差异。
重新检查代码后,我发现一个需要补上的实验环节:随机擦除加在了训练集上,测试集却仍然是干净图片。
因此,上次那段示例直接回答的是:
经过遮挡增强训练后,模型在干净图片上表现如何?
而“CNN 是否真的更抗遮挡”,还需要回答:
同一张测试图片真正被挡住一部分时,模型还能答对多少?相比没遮挡时,又退化了多少?
这次我把这两个问题拆开,重新做了完整对照。结果有一个值得注意的地方:正常训练的 CNN 在遮挡后仍然比 MLP 准,但相对退化也更严重;加入遮挡训练后,这个 CNN 才表现出明显更好的抗遮挡能力。
一、实验目标与事前猜想
这次不追求刷新 CIFAR-10 的准确率,而是验证四个具体猜想。
| 猜想 | 需要观察什么 |
|---|---|
| H1:正常训练的 CNN 在遮挡测试中仍然比 MLP 准 | 同一批遮挡测试图片上的绝对准确率;同时检查它是否真的下降更少 |
| H2:遮挡增强能改善模型在遮挡图片上的表现 | 同一模型、同一初始化种子,增强组减去普通组的准确率差 |
| H3:遮挡增强可能牺牲干净图片准确率 | 干净测试集上的增强组与普通组差异 |
| H4:同样遮挡面积,不同位置会影响结果 | 随机位置与固定中心位置的对照 |
这些猜想和实验配置在正式训练之前写入 EXPERIMENT_PLAN.md。正式测试结束后,没有根据结果继续调超参、挑种子或挑最好的 epoch。
这里首先区分三个指标:
遮挡后准确率 = 遮挡图片预测正确数 / 测试图片数
下降百分点 = 干净准确率(%)-遮挡准确率(%)
相对准确率保留比例 = 遮挡准确率 / 干净准确率 × 100%
例如,一个模型从 80% 降到 55%,另一个从 50% 降到 40%。前者遮挡后仍更准,但下降了 25 个百分点;后者只下降了 10 个百分点。“遮挡后更准”和“遇到遮挡退化更少”不是同一个结论。
准确率保留比例也是两个总体准确率的比值,并不是“干净图片原本答对的样本中,有多少在遮挡后仍答对”的条件概率。它帮助观察相对退化,但不能独自定义模型的全部鲁棒性。
二、实验策略:训练方式与测试条件分开
1. 四个训练组
| 模型 | 正常训练 plain |
遮挡增强训练 erase |
|---|---|---|
| MLP | MLP / 正常训练 | MLP / 遮挡增强 |
| CNN | CNN / 正常训练 | CNN / 遮挡增强 |
四个组都测试干净图片和相同的遮挡图片。这样就能分别比较模型结构差异、训练增强效果,以及两者共同作用后的表现。
每个组用 17、29、43 三个随机种子从头训练。同一模型、同一种子的普通组与增强组初始权重完全相同,样本打乱顺序也相同;只改变是否加入遮挡。保存的初始权重哈希已核对一致。
2. 数据划分
继续使用 CIFAR-10 官方数据,不使用合成图片或预训练权重。
| 用途 | 数量 | 每类数量 | 用法 |
|---|---|---|---|
| 训练集 | 45,000 | 4,500 | 更新模型参数 |
| 验证集 | 5,000 | 500 | 记录干净图片学习曲线 |
| 官方测试集 | 10,000 | 1,000 | 所有正式训练结束后统一评估 |
训练和验证从官方 50,000 张训练图片中,按固定种子 20260925 分层划分。两个集合没有交集。每组都取第 10 个 epoch 的模型,不使用测试集选模型,也不根据验证集选择最佳 checkpoint。
3. 训练遮挡怎样设置
每张训练图片以 50% 的概率被遮挡一次:
- 方块边长从
6—16像素的整数中均匀采样。 - 方块完全位于图片内部,位置随机。
- 被选中的图片,其遮挡面积约为
3.52%—25%。 - 输入先从
[0,255]归一化到[-1,1],再用0填充遮挡区域。
这里的 0 对应原图中的中灰色 RGB=(0.5,0.5,0.5),不是黑色。遮挡的是整张图片的面积,不是标注物体的面积。
这是一种便于控制面积和位置的方形遮挡增强。它与 Random Erasing 和 Cutout 的思路相关,但不等同于 torchvision.transforms.RandomErasing 的默认矩形长宽比分布,也不是对这两篇论文的完整复现。
为了减少额外变量,本次没有同时使用随机裁剪、翻转或颜色增强。
4. 测试遮挡怎样设置
| 测试条件 | 方块边长 | 实际遮挡面积 |
|---|---|---|
| 干净图片 | 0 | 0% |
| 随机位置 | 8 | 6.25% |
| 随机位置 | 12 | 14.0625% |
| 随机位置 | 16 | 25% |
| 随机位置 | 20 | 39.0625% |
| 固定中心位置 | 16 | 25% |
随机位置测试使用 101、202、303 三组固定遮挡种子。相同测试图片、相同边长、相同遮挡种子,所有模型看到的位置完全一致。
每次训练先对三个随机遮挡测试副本取平均,再对三个训练种子求均值和样本标准差。不能把“3 次训练 × 3 个测试遮挡副本”当成 9 次独立训练。
20×20 的测试方块比训练时的最大方块更大,可以观察较强遮挡下的表现,但仍然属于相同的灰色方块遮挡家族。

图中选取四个类别在官方测试集中首次出现的图片,没有按模型成败挑图。随机遮挡采用种子 101,最后一列为固定中心遮挡。位置和边长变化可能同时影响保留下来的语义信息,因此不会预设每一张图片的难度都随面积严格递增。
三、本机环境、模型与训练步骤
1. 为什么选择这个 Conda 环境
本机检查到两个已有环境:d2l 的 Python 3.9 / PyTorch 2.0.1,以及 zhcho 的 Python 3.10 / PyTorch 2.8.0。本次选择后者,并实际完成 MPS 卷积前向、反向和小规模流程检查。
| 项目 | 实测配置 |
|---|---|
| 芯片 / 内存 | Apple M3 Pro / 36GB |
| 系统 | macOS 26.6.2,arm64 |
| Conda 环境 | zhcho |
| Python | 3.10.19 |
| PyTorch / NumPy | 2.8.0 / 2.2.6 |
| matplotlib / pandas | 3.10.7 / 2.2.3 |
| 计算设备 | MPS;未切换为 CPU 训练 |
| CPU 线程数 | 4 |
已有 CIFAR-10 位于 /Users/zhcho/Documents/study/jupyter/data。脚本先检查官方批文件的 MD5,再加载数据。本次无需重新下载,也没有改动原有 Conda 环境。
训练代码只依赖 PyTorch 与 NumPy,直接读取经过校验的官方数据批,不依赖 torchvision。绘图与统计另外使用 matplotlib 和 pandas。
2. 控制模型规模
| 模型 | 简化结构 | 可训练参数 |
|---|---|---|
| MLP | 展平 → 3072→128 → ReLU → 128→10 | 394,634 |
| CNN | 3 个卷积块(32/64/128 通道)→ 4×4 特征 → 144 维隐藏层 → 10 类 | 390,202 |
两者参数量相差约 1.1%,比直接拿一个很大的 MLP 与很小的 CNN 对比更容易解释。但参数量相近不代表计算量相同,CNN 还包含 BatchNorm、池化和更深的计算路径。因此,本实验比较的是这两个具体模型,不能把差异全部归因于“卷积”这一个因素。
3. 固定训练预算
epoch = 10
batch_size = 256
optimizer = AdamW
initial learning rate= 0.001
weight_decay = 0.0001
learning-rate policy = cosine decay to 0.0001
training seeds = 17, 29, 43
checkpoint = fixed final epoch
每轮记录损失、训练准确率、干净验证准确率和耗时。训练增强组的训练准确率是在更难的增强输入上计算的,不能仅凭它低于普通组就认定模型更差。
4. 在本机运行
下载实验复现包(代码、图表、日志、原始预测及统计结果,约 1.8 MB)。受博客园单文件 10 MB 限制,附件省略 12 个模型权重文件,运行脚本即可重新生成。文末也附有完整训练代码。解压并进入目录后运行:
cd /path/to/cnn-occlusion # 替换为复现包的实际解压目录
/Users/zhcho/miniconda3/bin/conda run --no-capture-output -n zhcho \
python run_experiment.py \
--data-root /Users/zhcho/Documents/study/jupyter/data \
--out results-reproduce
results-reproduce 是新目录,避免覆盖这次已经保存的实测结果。如果要先检查流程,可以加上:
--epochs 1 --seeds 17 --train-limit 1000 --val-limit 500 --test-limit 1000
小规模检查须使用独立输出目录。它用于确认代码、设备和数据流程,不参与本文结论。
本次实际执行的是同一脚本的默认正式配置,使用 zhcho 环境中 Python 的绝对路径启动。含数据加载、12 次训练、验证和全部测试,/usr/bin/time -l 记录总耗时 245.96 秒,约 4 分 6 秒。
| 训练组 | 每次 10 epoch 的平均纯训练时间 | 平均训练+验证时间 |
|---|---|---|
| MLP / 正常训练 | 3.40 秒 | 3.56 秒 |
| MLP / 遮挡增强 | 6.44 秒 | 6.58 秒 |
| CNN / 正常训练 | 29.39 秒 | 30.28 秒 |
| CNN / 遮挡增强 | 31.46 秒 | 32.35 秒 |
以上时间来自本次设备和实现:数据一次性预装入计算设备内存,训练计时使用 MPS 同步。不同机器、首次图编译、系统负载和实现方式都会改变耗时,不应将它当成跨设备性能基准。
macOS 记录该进程最大常驻内存约 1.87 GiB、峰值 memory footprint 约 2.62 GiB;这些是进程口径,不是整台机器或独立 GPU 显存的总占用。
四、核心实验代码:把训练随机性和测试随机性分离
真正决定实验是否公平的,不只是模型定义,还包括随机数、遮挡与统计方式。
1. 使用独立随机数流
seed_all(seed) # 每个模型/种子配对重新初始化
model = MLP() if name == "MLP" else CNN()
order_rng = np.random.default_rng(seed + 10000)
augmentation_rng = np.random.default_rng(seed + 20000)
for epoch in range(epochs):
order = order_rng.permutation(len(train_x))
# 数据顺序不受“增强分支多消耗了随机数”影响
如果所有随机操作共用一个不断前进的随机数流,增强组多抽取几次随机数,就可能连后续样本顺序都改变。这里把样本顺序和增强参数分开。
2. 面积精确、位置固定的测试遮挡
下面是便于阅读的单张图版本;完整脚本使用语义一致的批量向量化实现,并逐项检查面积与复现性:
def occlude_one(image, side, top, left):
"""image: [C, H, W],已归一化到 [-1, 1]。"""
result = image.clone()
result[:, top:top + side, left:left + side] = 0.0
return result
def fixed_positions(count, side, mask_seed):
rng = np.random.default_rng(mask_seed)
top = (rng.random(count) * (33 - side)).astype(np.int64)
left = (rng.random(count) * (33 - side)).astype(np.int64)
return top, left
位置在整批测试图片上生成,再按 batch 切片;不会因为 batch_size 或模型不同而悄悄换一组遮挡。所有方块都完整落在 32×32 图片内,边长 16 就确实遮掉 256 个像素位置。
3. 测试不再依赖训练 transform
# 普通模型和增强模型都要分别经过这些评估条件
clean_prediction = predict(model, test_x, batch_size)
for side in [8, 12, 16, 20]:
for mask_seed in [101, 202, 303]:
occluded_prediction = predict(
model, test_x, batch_size,
side=side, mask_seed=mask_seed,
)
也就是说,训练时是否加遮挡 与 测试时是否加遮挡 是两个独立开关。完整脚本会保存各条件下每张图片的预测,而不只保存一个最终百分比。
五、实验结果:CNN 更准,但不一定退化更少
下表单位为百分比,格式为 三个训练种子的均值 ± 样本标准差。随机遮挡先在每个训练种子内部对三个固定遮挡副本取平均。标准差不是 95% 置信区间。
| 训练组 | 干净测试 | 遮挡 6.25% | 遮挡 14.06% | 遮挡 25% | 遮挡 39.06% |
|---|---|---|---|---|---|
| MLP / 正常训练 | 53.05 ± 0.35 | 50.74 ± 0.22 | 47.54 ± 0.09 | 42.32 ± 0.34 | 34.93 ± 0.41 |
| MLP / 遮挡增强 | 52.78 ± 0.04 | 51.19 ± 0.09 | 48.91 ± 0.09 | 45.00 ± 0.14 | 38.22 ± 0.41 |
| CNN / 正常训练 | 78.88 ± 0.20 | 74.61 ± 0.12 | 66.53 ± 0.59 | 52.93 ± 1.95 | 39.25 ± 2.41 |
| CNN / 遮挡增强 | 78.90 ± 0.11 | 76.85 ± 0.25 | 74.11 ± 0.24 | 68.73 ± 0.24 | 59.02 ± 0.62 |

蓝色为 MLP,橙色为 CNN;实线为正常训练,虚线为遮挡增强训练。阴影表示三个训练种子间的样本标准差。
1. 只比较遮挡后的正确率,CNN 确实更高
在 25% 随机遮挡下,正常训练 CNN 的平均准确率是 52.93%,MLP 是 42.32%,前者高 10.61 个百分点。在本次另外三档随机遮挡中,CNN 普通组的平均准确率也更高。
这支持 H1 关于“遮挡后绝对准确率”的部分。
2. 但普通 CNN 的退化更严重
把干净图片的基线也放进来,结论就不同了:
| 训练组 | 干净准确率 | 25% 遮挡准确率 | 下降(百分点) | 相对准确率保留比例 |
|---|---|---|---|---|
| MLP / 正常训练 | 53.05% | 42.32% | 10.73 | 79.78% |
| MLP / 遮挡增强 | 52.78% | 45.00% | 7.79 | 85.25% |
| CNN / 正常训练 | 78.88% | 52.93% | 25.94 | 67.11% |
| CNN / 遮挡增强 | 78.90% | 68.73% | 10.17 | 87.11% |
普通 CNN 从 78.88% 降到 52.93%,下降 25.94 个百分点;普通 MLP 从 53.05% 降到 42.32%,下降 10.73 个百分点。
普通 CNN 的相对准确率保留比例为 67.11%,普通 MLP 为 79.78%。因此,不能把“普通 CNN 遮挡后仍更准”直接改写为“普通 CNN 天然退化更少”。
反过来也不能说 MLP 就是更好的遮挡识别模型:它干净和遮挡后的绝对准确率都更低。较低起点会影响退化指标的解释,所以本文同时展示绝对准确率、下降百分点和相对比值,而不单看一个数字。
3. 遮挡增强对本次 CNN 的帮助更明显
在 25% 随机遮挡下:
- CNN:52.93% → 68.73%,提高 15.80 个百分点。
- MLP:42.32% → 45.00%,提高 2.67 个百分点。
CNN 的下降幅度也由 25.94 缩小到 10.17 个百分点,相对准确率保留比例由 67.11% 提高到 87.11%。
为了确认这个结论不是只靠某个训练种子,进一步检查配对结果:
| 模型 | 训练种子 | 正常训练 | 遮挡增强训练 | 增益(百分点) |
|---|---|---|---|---|
| CNN | 17 | 50.79% | 68.46% | +17.67 |
| CNN | 29 | 54.61% | 68.92% | +14.31 |
| CNN | 43 | 53.40% | 68.82% | +15.42 |
| MLP | 17 | 42.68% | 45.16% | +2.48 |
| MLP | 29 | 42.30% | 44.92% | +2.62 |
| MLP | 43 | 41.99% | 44.90% | +2.91 |
三个 CNN 配对的增益均为正,范围为 14.31—17.67 个百分点;MLP 的三个配对也均为正,但增益较小。这里的结论是“三次实测中方向一致”,没有把它包装成已完成统计显著性检验。

这支持 H2:在本次配置下,训练时接触遮挡,提高了测试时同类遮挡下的表现。较强的 39.06% 遮挡中,CNN 增强组仍达到 59.02%,普通组只有 39.25%;但它依然明显低于干净图片准确率,不能说增强之后就“不怕遮挡”。
4. 干净图片上没有出现统一的巨大代价
CNN 的干净准确率从 78.88% 变为 78.90%,平均差异约 +0.02 个百分点;逐种子差异有正有负,数值很小,不能据此宣称增强提高了干净图片表现。
MLP 则从 53.05% 变为 52.78%,平均下降 0.27 个百分点,三个配对差异都是轻微负值。
因此,H3 应当保留为条件性的判断:遮挡增强可能有代价,但本次 CNN 没有观察到明显的平均干净准确率损失;本次 MLP 出现了小幅损失。这不等于其他增强强度、训练预算和数据集也会如此。

右图是干净验证集曲线。两种 CNN 的末轮验证准确率接近,这与最终干净测试结果相吻合。左图中增强组训练准确率较低,是在不同难度的训练输入上得到的数值,不应只比较训练曲线高低。
5. 遮挡位置也会改变结论
相同 16×16 方块,放到固定中心后,四组平均准确率都比随机位置更低:
| 训练组 | 随机位置遮挡 25% | 中心遮挡 25% | 中心相对随机位置少(百分点) |
|---|---|---|---|
| MLP / 正常训练 | 42.32% | 38.89% | 3.43 |
| MLP / 遮挡增强 | 45.00% | 42.10% | 2.89 |
| CNN / 正常训练 | 52.93% | 45.97% | 6.97 |
| CNN / 遮挡增强 | 68.73% | 65.33% | 3.40 |
这支持 H4。一个可能的解释是,CIFAR-10 部分图片的关键对象位于中心,但本次没有使用目标框或分割标注测量“遮掉了多少物体”,因此这只是合理解释,尚不是被本实验单独证明的原因。
六、实验过程中实际遇到与处理的问题
1. 上一篇的训练增强和测试鲁棒性口径混在了一起
原代码在训练集中加入随机擦除,却仍然在干净测试集上计算准确率。它并非完全没有价值,但不能单独支持“测试图片被遮挡时更稳”的结论。
本次增加独立测试遮挡,并让四组模型经历完全相同的测试条件。每一条测试记录都保留 condition、side、mask_seed、correct、total。
2. 有 Conda 环境,不代表当前 shell 能直接找到 conda
本次当前 shell 的 command -v conda 没有返回路径,但 /Users/zhcho/miniconda3 下确实存在环境。最终使用绝对路径启动环境中的 Python;本文复现命令也使用 Conda 的绝对路径,避免把“命令不在 PATH”误判为“没装环境”。
3. 参数规模和随机遮挡都可能成为干扰变量
如果 MLP 比 CNN 大很多,或者每个模型各自生成一批遮挡,就很难说清比较的是模型、规模还是样本难度。本次把参数量控制到约 39 万,并固定数据划分、配对初始化、样本顺序和测试遮挡。
脚本额外检查了四个具体条件:遮挡面积精确、重复生成结果相同、不会原地改坏原始输入、改变 batch 切片不会改变遮挡结果。
这些控制降低了干扰,却没有消除所有因素:模型深度、BatchNorm、计算量和优化难度仍不相同。
4. 归一化后的填充值容易解释错
如果直接写 image[..., region] = 0,必须先说明图像所处的数值空间。在本次 [-1,1] 空间中,它代表中灰色。不能把灰块测试结果描述成“黑色遮挡”的结果,更不能默认填充颜色没有影响。
5. 多个遮挡副本不能冒充多次独立训练
每个训练模型做了三个随机遮挡副本,目的是减小某一套遮挡位置的偶然影响。它们共享同一个模型,所以统计时先在模型内部平均;最终误差棒仍然只有三个独立训练种子作为重复单位。
6. 图表初版的底部说明与横轴标签重叠
首次渲染后检查图片,发现底部的统计口径说明压到了横轴标签上。随后增加底部留白,并缩短种子标注,再重新生成和检查图表。这个修正只涉及展示,不改变训练或统计数据。
本次正式训练没有出现 OOM、非有限损失或下载失败。已有数据校验通过,因此也没有把另一次实验的数据下载经历写成本次遇到的问题。
七、数据怎样核对,结论到哪里为止
本次共保存并检查了:
- 12 次完整训练、120 条 epoch 记录。
- 每次模型 14 个测试条件,共 168 条测试记录。
- 1,680,000 条预测与标签比较记录(同一批 10,000 张测试图在多个模型/条件下重复评估,不是 168 万张独立图片)。
- 六对普通组/增强组初始权重哈希一致。
- 全部准确率重新从保存的预测计算,与结果文件一致。
每次模型的 14 个测试条件包括:1 个干净条件、4 档面积 × 3 个随机遮挡副本,以及 1 个中心遮挡条件。
可直接检查这些文件:
results/config.json # 正式配置和训练脚本 SHA256
results/environment.json # 环境与实际执行参数
results/data_integrity.json # 数据校验、分布、遮挡检查
results/split_indices.npz # 固定训练/验证/测试索引
results/all_history.csv # 120 条训练记录
results/all_results.csv # 168 条测试条件记录
results/per_training_seed.csv # 先在每次训练内平均遮挡副本
results/summary.csv # 均值与样本标准差
results/paired_augmentation_gains.csv # 增强的配对增益
results/verification.json # 数据核对结果
results/CNN_erase_17/predictions.npz # 示例:逐样本预测
results/CNN_erase_17/model.pt # 训练后生成:末轮模型参数(附件未含权重)
需要保留的边界也很明确:
- 只有 CIFAR-10、两个特定模型和 10 epoch 的训练预算,不代表所有 CNN 和 MLP。
- 遮挡是合成灰色方块,不代表真实物体遮挡、噪声、旋转或其他分布变化。
- 只有三个训练种子,能看方向和波动,尚不足以得出广泛统计结论。
- 参数量接近不等于结构与优化难度完全公平;统一超参也不等于两类模型各自的最优配置。
- 固定 CPU 随机数流与遮挡种子提高复现性,但不承诺不同机器、PyTorch 版本和 MPS 后端之间逐位一致。
八、最后的结论
这次数据验证了一个比“CNN 天然抗遮挡”更具体的结论:
在本次 CIFAR-10 对照实验中,普通 CNN 在遮挡后的绝对准确率高于普通 MLP,但下降幅度和相对退化更大。加入方形遮挡训练后,CNN 在 25% 随机遮挡下从 52.93% 提高到 68.73%,且干净图片平均准确率基本持平。
因此,我更愿意把本次得到的认识写成:
模型结构决定了它能怎样利用图像信息;训练分布影响它面对信息缺失时的表现。判断鲁棒性,必须让模型真正经历要验证的测试扰动,并把基线、退化和训练方法一起交代清楚。
下一步如果继续研究,可以预先固定黑色、随机纹理或真实遮挡物等新测试条件,再做一次独立评估;也可以增加训练种子和模型结构。这些是后续实验方向,不属于本次已经完成的结果。
参考资料
- CIFAR-10 官方数据与技术报告入口
- PyTorch MPS 后端说明
- torchvision RandomErasing 参数说明
- Random Erasing Data Augmentation
- Improved Regularization of Convolutional Neural Networks with Cutout
附录:完整训练与评估脚本
保存为 run_experiment.py,使用上文命令运行。脚本默认拒绝覆盖已有输出,完整复现说明见同目录 README.md;统计和绘图脚本为 analyze_results.py。
#!/usr/bin/env python3
"""CIFAR-10: separate training augmentation from test-time occlusion.
No torchvision dependency. Uses verified official CIFAR-10 Python batch files.
All model groups share splits, data order and test masks. Test data is evaluated
only after the fixed final training epoch. See EXPERIMENT_PLAN.md.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import math
import pickle
import platform
import random
import sys
import time
from datetime import datetime, timezone
from pathlib import Path
import numpy as np
import torch
from torch import nn
MD5 = {
'data_batch_1': 'c99cafc152244af753f735de768cd75f',
'data_batch_2': 'd4bba439e000b95fd0a9bffe97cbabec',
'data_batch_3': '54ebc095f3ab1f0389bbae665268c751',
'data_batch_4': '634d18415352ddfa80567beed471001a',
'data_batch_5': '482c414d41f54cd18b22e5b47cb7c3cb',
'test_batch': '40351d587109b95175f43aff81a1287e',
'batches.meta': '5ff9c542aee3614f3951f8cda6e48888',
}
TEST_SIDES = (0, 8, 12, 16, 20)
MASK_SEEDS = (101, 202, 303)
class MLP(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(nn.Flatten(), nn.Linear(3072, 128),
nn.ReLU(), nn.Linear(128, 10))
def forward(self, x):
return self.net(x)
class CNN(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(3, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2),
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2),
nn.Conv2d(64, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(),
nn.AdaptiveAvgPool2d((4, 4)), nn.Flatten(), nn.Linear(2048, 144),
nn.ReLU(), nn.Linear(144, 10))
def forward(self, x):
return self.net(x)
def json_write(path, obj):
path.write_text(json.dumps(obj, ensure_ascii=False, indent=2) + '\n')
def write_csv(path, rows):
if rows:
with path.open('w', newline='') as f:
w = csv.DictWriter(f, fieldnames=list(rows[0]))
w.writeheader()
w.writerows(rows)
def sync(device):
if device.type == 'mps':
torch.mps.synchronize()
elif device.type == 'cuda':
torch.cuda.synchronize()
def seed_all(seed):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def load_cifar(root):
"""Verify checksums BEFORE unpickling the official, trusted dataset."""
if (root / 'cifar-10-batches-py').is_dir():
root = root / 'cifar-10-batches-py'
for name, expected in MD5.items():
path = root / name
if not path.is_file():
raise FileNotFoundError(f'{path}: download/extract the official CIFAR-10 Python archive first.')
actual = hashlib.md5(path.read_bytes()).hexdigest()
if actual != expected:
raise ValueError(f'CIFAR-10 checksum mismatch for {name}: {actual}')
def read(name):
with (root / name).open('rb') as f:
batch = pickle.load(f, encoding='bytes')
return batch[b'data'].reshape(-1, 3, 32, 32), np.asarray(batch[b'labels'], dtype=np.int64)
parts = [read(f'data_batch_{i}') for i in range(1, 6)]
tx, ty = read('test_batch')
return np.concatenate([p[0] for p in parts]), np.concatenate([p[1] for p in parts]), tx, ty
def split_indices(labels, train_limit, val_limit):
"""One fixed stratified split, independent of model initialization seeds."""
rng = np.random.default_rng(20260925)
train, val = [], []
for cls in range(10):
ids = rng.permutation(np.flatnonzero(labels == cls))
val.extend(ids[:500][:val_limit // 10])
train.extend(ids[500:][:train_limit // 10])
return np.asarray(train), np.asarray(val)
def test_indices(labels, limit):
# Full run preserves official test order; smoke mode takes balanced subsets.
if limit == len(labels):
return np.arange(len(labels))
return np.concatenate([np.flatnonzero(labels == c)[:limit // 10] for c in range(10)])
def device_tensor(data, device):
return torch.from_numpy(data.astype(np.float32) / 127.5 - 1.0).to(device)
def square_mask(x, sides, top, left, active=None):
"""Vectorized independent square masks; normalized 0 == RGB 0.5 grey."""
n, _, h, w = x.shape
def tensor(v):
return torch.as_tensor(np.asarray(v), device=x.device).reshape(n, 1, 1, 1)
side, t, l = tensor(sides), tensor(top), tensor(left)
yy = torch.arange(h, device=x.device).reshape(1, 1, h, 1)
xx = torch.arange(w, device=x.device).reshape(1, 1, 1, w)
mask = (yy >= t) & (yy < t + side) & (xx >= l) & (xx < l + side)
if active is not None:
mask = mask & tensor(active)
return x.masked_fill(mask, 0.0)
def training_occlusion(x, rng):
n = len(x)
sides = rng.integers(6, 17, n)
top = (rng.random(n) * (33 - sides)).astype(np.int64)
left = (rng.random(n) * (33 - sides)).astype(np.int64)
active = rng.random(n) < 0.5
return square_mask(x, sides, top, left, active)
def test_positions(n, side, seed, center=False):
if center:
pos = np.full(n, (32 - side) // 2, dtype=np.int64)
return pos, pos.copy()
rng = np.random.default_rng(seed)
# Same random positions for every architecture/training condition/init seed.
top = (rng.random(n) * (33 - side)).astype(np.int64)
left = (rng.random(n) * (33 - side)).astype(np.int64)
return top, left
@torch.inference_mode()
def predict(model, x, batch_size, side=0, mask_seed=0, center=False):
model.eval()
top, left = test_positions(len(x), side, mask_seed, center)
preds = []
for start in range(0, len(x), batch_size):
end = min(start + batch_size, len(x))
xb = x[start:end]
if side:
xb = square_mask(xb, np.full(end-start, side), top[start:end], left[start:end])
preds.append(model(xb).argmax(1).cpu().numpy())
return np.concatenate(preds)
def check_mask(device):
# Meaningful invariants: exact area, no input mutation, same masks on repeat,
# and batch slicing must not alter test masks.
x = torch.ones(11, 3, 32, 32, device=device)
for side in TEST_SIDES:
top, left = test_positions(len(x), side, 101)
masked = square_mask(x, np.full(len(x), side), top, left)
count = (masked[:, 0] == 0).sum((1, 2)).cpu().numpy()
assert np.all(count == side * side)
again = square_mask(x, np.full(len(x), side), top, left)
assert torch.equal(masked, again)
halves = torch.cat([square_mask(x[:5], np.full(5, side), top[:5], left[:5]),
square_mask(x[5:], np.full(6, side), top[5:], left[5:])])
assert torch.equal(masked, halves)
assert bool((x == 1).all())
return 'exact area, reproducible masks, unchanged input and batch invariance: passed'
def train_one(name, aug, seed, x, y, vx, vy, args, device, run_dir):
run_dir.mkdir(exist_ok=True)
checkpoint = run_dir / 'model.pt'
if args.resume and checkpoint.exists() and (run_dir / 'training.json').exists():
print(f'RESUME trained {run_dir.name}', flush=True)
return
seed_all(seed)
model = (MLP() if name == 'MLP' else CNN()).to(device)
initial_hash = hashlib.sha256(b''.join(t.detach().cpu().numpy().tobytes() for t in model.state_dict().values())).hexdigest()
optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=1e-4)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs, eta_min=args.lr / 10)
order_rng = np.random.default_rng(seed + 10000)
augmentation_rng = np.random.default_rng(seed + 20000)
loss_fn = nn.CrossEntropyLoss()
history = []
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
model.train()
order = order_rng.permutation(len(x))
total_loss = torch.zeros((), device=device)
correct = torch.zeros((), dtype=torch.int64, device=device)
sync(device)
epoch_start = time.perf_counter()
for start in range(0, len(x), args.batch_size):
ids = torch.as_tensor(order[start:start+args.batch_size], device=device)
xb, yb = x[ids], y[ids]
if aug == 'erase':
xb = training_occlusion(xb, augmentation_rng)
optimizer.zero_grad(set_to_none=True)
logits = model(xb)
loss = loss_fn(logits, yb)
loss.backward()
optimizer.step()
total_loss += loss.detach() * len(ids)
correct += (logits.detach().argmax(1) == yb).sum()
sync(device)
train_seconds = time.perf_counter() - epoch_start
preds = predict(model, vx, args.batch_size)
val_acc = float((preds == vy).mean())
row = dict(model=name, augmentation=aug, seed=seed, epoch=epoch,
lr=optimizer.param_groups[0]['lr'], loss=float(total_loss.cpu()) / len(x),
train_accuracy=float(correct.cpu()) / len(x), val_clean_accuracy=val_acc,
train_seconds=train_seconds)
if not math.isfinite(row['loss']):
raise RuntimeError('Non-finite training loss')
history.append(row)
write_csv(run_dir / 'history.csv', history)
print(f'{run_dir.name} epoch {epoch:02}/{args.epochs} loss={row["loss"]:.4f} '
f'train={row["train_accuracy"]:.4f} val={val_acc:.4f} sec={train_seconds:.2f}', flush=True)
scheduler.step()
sync(device)
elapsed = time.perf_counter() - started
# Always the fixed final epoch; no test-set or validation-based checkpoint selection.
torch.save({k: v.detach().cpu() for k,v in model.state_dict().items()}, checkpoint)
json_write(run_dir / 'training.json', dict(model=name, augmentation=aug, seed=seed,
epochs=args.epochs, parameters=sum(p.numel() for p in model.parameters()),
initial_state_sha256=initial_hash, train_loop_seconds=sum(r['train_seconds'] for r in history),
training_and_validation_seconds=elapsed, final_val_clean_accuracy=val_acc))
del model, optimizer
if device.type == 'mps':
torch.mps.empty_cache()
def evaluate_one(name, aug, seed, x, labels, args, device, run_dir):
if args.resume and (run_dir / 'evaluation.json').exists():
print(f'RESUME evaluated {run_dir.name}', flush=True)
return
model = (MLP() if name == 'MLP' else CNN()).to(device)
model.load_state_dict(torch.load(run_dir / 'model.pt', map_location='cpu', weights_only=True))
rows, saved = [], {'labels': labels}
started = time.perf_counter()
conditions = [(0, 0, False)] + [(s, ms, False) for s in TEST_SIDES[1:] for ms in MASK_SEEDS] + [(16, 0, True)]
for side, ms, center in conditions:
pred = predict(model, x, args.batch_size, side, ms, center)
condition = 'clean' if side == 0 else ('center' if center else 'random')
key = f'{condition}_s{side}_m{ms}'
saved[key] = pred.astype(np.int16)
correct = int((pred == labels).sum())
rows.append(dict(model=name, augmentation=aug, seed=seed, condition=condition,
side=side, area_fraction=side * side / 1024, mask_seed=ms,
correct=correct, total=len(labels), accuracy=correct/len(labels)))
np.savez_compressed(run_dir / 'predictions.npz', **saved)
write_csv(run_dir / 'evaluation.csv', rows)
json_write(run_dir / 'evaluation.json', dict(seconds=time.perf_counter()-started, conditions=len(rows)))
print(f'EVALUATED {run_dir.name}: clean={rows[0]["accuracy"]:.4f}; '
f'random25%={np.mean([r["accuracy"] for r in rows if r["condition"]=="random" and r["side"]==16]):.4f}', flush=True)
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument('--data-root', type=Path, required=True)
ap.add_argument('--out', type=Path, default=Path('results'))
ap.add_argument('--device', choices=['auto','mps','cpu','cuda'], default='auto')
ap.add_argument('--epochs', type=int, default=10)
ap.add_argument('--batch-size', type=int, default=256)
ap.add_argument('--lr', type=float, default=1e-3)
ap.add_argument('--seeds', type=int, nargs='+', default=[17,29,43])
ap.add_argument('--models', nargs='+', choices=['MLP','CNN'], default=['MLP','CNN'])
ap.add_argument('--augmentations', nargs='+', choices=['plain','erase'], default=['plain','erase'])
ap.add_argument('--train-limit', type=int, default=45000)
ap.add_argument('--val-limit', type=int, default=5000)
ap.add_argument('--test-limit', type=int, default=10000)
ap.add_argument('--phase', choices=['all','train','eval','check'], default='all')
ap.add_argument('--resume', action='store_true')
args = ap.parse_args()
for value, maximum in [(args.train_limit,45000),(args.val_limit,5000),(args.test_limit,10000)]:
if not 0 < value <= maximum or value % 10:
ap.error('dataset limits must be positive multiples of 10 within official split sizes')
if args.epochs < 1 or args.batch_size < 1:
ap.error('epochs and batch-size must be positive')
torch.set_num_threads(4)
selected = ('mps' if torch.backends.mps.is_available() else 'cuda' if torch.cuda.is_available() else 'cpu') if args.device=='auto' else args.device
device = torch.device(selected)
args.out.mkdir(parents=True, exist_ok=True)
config = {k:str(v) if isinstance(v,Path) else v for k,v in vars(args).items() if k not in ('resume','phase','out')}
config.update(actual_device=selected, source_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest())
config_path = args.out/'config.json'
if config_path.exists():
if not args.resume:
raise FileExistsError('Output already contains a run; use a new --out or explicit --resume.')
if json.loads(config_path.read_text()) != config:
raise ValueError('Configuration/source changed; use a new output directory.')
else:
json_write(config_path, config)
environment = dict(started_utc=datetime.now(timezone.utc).isoformat(), python=sys.version,
executable=sys.executable, torch=torch.__version__, numpy=np.__version__, platform=platform.platform(),
mps_available=torch.backends.mps.is_available(), mps_built=torch.backends.mps.is_built(),
device=selected, threads=torch.get_num_threads(), argv=sys.argv,
determinism='CPU RNG streams are seeded; bitwise MPS determinism across machines/versions is not promised.')
if not (args.out/'environment.json').exists():
json_write(args.out/'environment.json', environment)
print(json.dumps(environment,ensure_ascii=False),flush=True)
check = check_mask(device)
print('MASK CHECK:',check,flush=True)
if args.phase=='check':
return
raw, labels, raw_test, labels_test = load_cifar(args.data_root)
train_ids, val_ids = split_indices(labels,args.train_limit,args.val_limit)
test_ids = test_indices(labels_test,args.test_limit)
assert not np.intersect1d(train_ids,val_ids).size
np.savez_compressed(args.out/'split_indices.npz',train=train_ids,validation=val_ids,test=test_ids)
json_write(args.out/'data_integrity.json',dict(md5=MD5, train=len(train_ids), validation=len(val_ids),
test=len(test_ids),train_per_class=np.bincount(labels[train_ids]).tolist(),
validation_per_class=np.bincount(labels[val_ids]).tolist(),
test_per_class=np.bincount(labels_test[test_ids]).tolist(),mask_invariants=check))
vx=device_tensor(raw[val_ids],device)
if args.phase in ('all','train'):
x=device_tensor(raw[train_ids],device)
y=torch.as_tensor(labels[train_ids],device=device)
for seed in args.seeds:
for name in args.models:
for aug in args.augmentations:
train_one(name,aug,seed,x,y,vx,labels[val_ids],args,device,args.out/f'{name}_{aug}_{seed}')
del x,y
if args.phase in ('all','eval'):
tx=device_tensor(raw_test[test_ids],device)
for seed in args.seeds:
for name in args.models:
for aug in args.augmentations:
evaluate_one(name,aug,seed,tx,labels_test[test_ids],args,device,args.out/f'{name}_{aug}_{seed}')
all_rows=[]
for seed in args.seeds:
for name in args.models:
for aug in args.augmentations:
with (args.out/f'{name}_{aug}_{seed}'/'evaluation.csv').open() as f:
all_rows.extend(csv.DictReader(f))
write_csv(args.out/'all_results.csv',all_rows)
json_write(args.out/f'completed_{args.phase}.json',dict(completed_utc=datetime.now(timezone.utc).isoformat()))
if __name__ == '__main__':
main()

浙公网安备 33010602011771号