ARTICLE DETAIL

资讯详情

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

ResNet真假图识别实战:从模型选型到阈值校准

ResNet真假图识别实战:从模型选型到阈值校准 简介一份基于PyTorch的ResNet图像真假识别训练代码面向具备Python基础、希望快速上手图像分类项目的开发者。资源不含数据集图片下载后按提示自行收集真图/假图放入对应文件夹即可投入训练整体流程清晰。压缩包共7个文件含3个Python脚本——分别负责生成训练集/验证集txt列表、启动CNN训练、提供PyQt图形界面另配2张放置提示图、1份依赖清单和1份说明文档资源包仅190KB。已有46人学习浏览。代码内置逐行注释并附说明文档可降低入门门槛训练脚本会自动读取分类文件夹数量增加类别时无需修改代码并能实时显示进度条、每轮准确率与损失值训练结束后保存日志及模型权重方便复盘和二次调参。适合需要快速验证图像分类流程、或希望拥有可扩展训练框架的开发者参考。1. 一份不带数据集图片的 resnet 真假图识别项目到底能教会你什么很多人拿到这份 resnet 模型代码包第一反应是先翻 dataset 文件夹结果发现 zip 里除去逐行注释的 Python 脚本、说明文档和必要的配置文件外一张训练图片都没有心里就开始打鼓没有数据这项目还能落地吗我的看法正相反不带数据反而更接近真实工程状态。自建真假图片数据集本来就是这个方向最值钱的部分依赖别人打包好的图片你反而很难判断模型到底是真的学到了伪造痕迹还是单纯背下了文件夹顺序。这个项目标题里锁定的技术路径很清楚用 CNN 训练识别真假图片主模型选用 resnet交付时以代码和文档为主。它适合两类人一类是刚做完图像分类入门想搞明白 resnet 预训练模型怎么微调以及训练脚本里每个参数是干什么用的新手另一类是需要给内容审核、商品图校验、存证系统做辅助鉴别模型的工程师。我会按我平时做这类活儿的顺序从模型选型、自建数据、训练脚本、踩坑排查一直讲到可以落地的阈值校准把每一段能复用、能抄作业的部分都拆开讲。2. 为什么是 ResNet真假图识别里的 CNN 卷积神经网络选型逻辑2.1 伪造痕迹藏在高频细节里CNN 卷积神经网络到底在看什么真实照片和伪造图片之间的差异跟“猫和狗”的差异完全不是一个层级。真实照片有传感器噪声、光照渐变、镜头畸变、JPEG 压缩痕迹这些信号大多集中在图像的高频局部区域。而 AI 合成图、修图软件涂抹过的区域往往表现出过度平滑、边缘过渡异常、颜色统计分布崩塌、压缩伪影不够自然。CNN 卷积神经网络依靠卷积核在局部感受野里逐层扫描浅层卷积核捕捉边缘、纹理和噪声模式深层卷积核再组合出语义级别的特征。所以它天然比传统手工特征更适合干这件事。但这里有个容易忽视的边界resnet 的默认输入是 224×224这个分辨率对于猫狗分类够用因为语义信息不会因为缩小而消失。对于真假图检测如果伪造区域只有 32×32甚至更小整张图缩到 224 会把最关键的高频痕迹抹掉。我常见的做法是先用 256 或 384 的输入尺寸训练配合随机裁剪让网络有机会看到局部细节只有推理性能吃紧时才退回 224 并叠加多尺度预测。换句话说模型选型不只是选 resnet 还是 vgg还要连同输入分辨率、数据增强策略一起定。从实际项目经验看还有一点值得提前说不要把真假图的判断理解成一个纯粹的二分类。很多数据集中“假”这一类别内部差异极大可能是 GAN 生成的假人像也可能是 Photoshop 把商品标签抠掉再贴上去的假图。这两类伪造的痕迹完全不同一个模型很难同时吃下。如果压缩包说明文档里没有强调这一点建议自己先把 fake 类做成分桶标签先训练一个“真实 vs 生成 vs 编辑”的三分类模型再在业务端把后两类合并成“假”。这件事想不清楚后面所有训练都容易翻车。2.2 残差连接为什么能让真假图分类训练更稳早期做图像分类大家喜欢堆 VGG 那种纯卷积层。层数一深梯度要从最后一层一路传回第一层链式法则里连续相乘会让梯度迅速变小网络浅层几乎学不到东西这就是常说的梯度消失。ResNet 的做法是在每个残差块里加了一条“抄近道”的恒等连接让输入 x 可以直接加到卷积输出 F(x) 上整个块变成 F(x)x。这样一来梯度可以从深层的损失函数直接经这条短路传回浅层即使中间卷积分支的梯度很小骨干网络也能持续更新。这个结构对真假图识别特别有价值。因为真实图和伪造图在整体视觉上都是自然图像像素差异很小网络想要学到的其实是一个“偏离量”而不是全新的特征。残差块的恒等路径相当于告诉网络你已经有预训练模型带来的强先验了卷积分支只需要在小范围内修正这个先验即可。学习修正量比从零学习一套新特征容易得多收敛也更稳定。我在一个自建的商品图真伪集上做过快速对照同样的学习率、同样的 30 个 epochResNet18 的验证集波动明显小于 VGG16尤其在前 10 个 epochVGG16 的 loss 下降很慢。这类项目里稳定性比极限精度更重要因为真假图分布随时会变模型不能一换训练集就失控。另外要提一下残差的另一个好处它让网络宽度不被层数卡死。ResNet18 只有 18 层参数量约 1100 万但在中小数据集上经常比 50 层版本更合适。深层 resnet 的容量大却对数据量足够多、伪造模式足够丰富有更高要求。数据只有两三千张时强行上 ResNet50训练集能拟合得很好验证集却会因为特征过拟合而表现飘忽。这不是说 ResNet50 不行而是说你得清楚它在什么时候才值得用。2.3 resnet 预训练模型怎么选18 层还是 50 层冻结还是全量微调既然项目标题点名了 resnet就必须把 resnet 预训练模型的使用方式说透。绝大多数训练脚本都会写成下面这种可切换结构。# model_builder.py import torch.nn as nn import torchvision.models as models def build_model(archresnet18, num_classes2, pretrainedTrue): if arch resnet18: model models.resnet18(pretrainedpretrained) # 加载 ImageNet 预训练权重 elif arch resnet50: model models.resnet50(pretrainedpretrained) else: raise ValueError(到这里只支持 resnet18 / resnet50) # resnet 最后一层全连接输入维度由卷积部分决定 in_features model.fc.in_features # 把原来 1000 类分类头替换成我们自己的 2 类 model.fc nn.Linear(in_features, num_classes) return model这段代码最关键的地方是in_features model.fc.in_features。ResNet18 和 ResNet50 最后一个池化层输出的特征维度不同前者是 512后者是 2048写死数字以后换网络就得改代码动态读取则永远正确。pretrainedTrue会从 torchvision 下载 ImageNet 权重如果你的环境不能联网也可以先把权重文件放到 torchvision 缓存目录再把pretrained参数换成weights_path效果是一样的。ImageNet 权重在这里是正收益因为真实照片的纹理基元、边缘响应、色彩统计都跟 ImageNet 里的自然图片高度重合预训练模型对浅层特征的初始化比随机权重靠谱得多。pretrained 权重拿到之后还有一个选择冻结浅层还是全量微调。我的经验可以用一张表说清楚。每类图片数量建议方案学习率参考少于 500 张只用 ResNet18冻结卷积前三个阶段只训练最后阶段和全连接层1e-4500 到 5000 张ResNet18 或 ResNet50 全量微调配合数据增强1e-4 到 3e-4超过 5000 张优先 ResNet50 全量微调输入尺寸可以提到 3841e-4 并配合余弦衰减冻结操作很容易遍历模型的参数把前面阶段的requires_grad设为 False优化器只接收需要更新的参数。但有一点要特别注意BatchNorm层的均值和方差在冻结状态下要用预训练统计量不能继续更新否则会出现训练 loss 降不下去的问题。很多翻车案例都是只冻了卷积、没冻 BN导致推理时 BN 统计量和训练时对不上。这个点我后面在避坑章节还会展开。3. 数据集图片不在压缩包里把自备真伪图片整理成 ResNet 能吃的格式3.1 先做目录结构再做训练集、验证集、测试集划分zip 里没有数据集图片意味着首次运行前必须自己搭建数据目录。最常见的组织方式有两种第一种是直接按类别建文件夹ImageFolder可以直接读第二种是维护一份 CSV 清单里面记录图片路径和标签适合图片分散在多台机器上的场景。我推荐把两种结合起来原始图片放在一个只读目录里再用脚本生成一份统一的划分清单最后按清单把文件复制到标准的 train/val/test 目录。这样做的原因很简单训练代码只认目录结构不关心你的原始素材是从相机里导出的、还是从生成模型里批量跑出来的。# make_split.py import os import random import shutil # 固定随机种子让同一份原始素材每次都得到相同划分 random.seed(2024) # 原始素材目录real 放真实照片fake 放编辑/生成图 raw_folders { real: raw/real, fake: raw/fake, } train_ratio, val_ratio 0.7, 0.15 split_names [train, val, test] for cls_name, src_dir in raw_folders.items(): files [f for f in os.listdir(src_dir) if f.lower().endswith((.jpg, .jpeg, .png))] random.shuffle(files) n_train int(len(files) * train_ratio) n_val int(len(files) * val_ratio) # 剩下自动进入 test避免生成集和测试集混淆 flists { train: files[:n_train], val: files[n_train:n_train n_val], test: files[n_train n_val:], } for split_name in split_names: dst_dir os.path.join(data, split_name, cls_name) os.makedirs(dst_dir, exist_okTrue) for fname in flists[split_name]: src os.path.join(src_dir, fname) dst os.path.join(dst_dir, fname) shutil.copy2(src, dst) # copy2 会保留拍摄时间等元信息这段脚本把 70% 的图分给训练、15% 分给验证、15% 分给测试。有人会问验证集和测试集都是不参与训练的为什么要分两份因为验证集承担的是“训练过程中的模型挑选”功能测试集承担的是“最终效果评估”功能。如果只用一个验证集你很容易在反复调优的过程中把验证集当成隐式训练集来用最后得到虚高指标。shutil.copy2保留 EXIF 信息这一点在真假图检测里非常重要。修图工具通常会改写图片元数据保留原始 EXIF 能在后期做跨模型交叉验证时排掉很多干扰。如果你更在意磁盘空间也可以不复制文件改成把划分结果写进 CSV自定义 Dataset 按路径读取。两种方式我都用过小项目复制文件最省事路径一乱起来毛病少。3.2 编写自定义 Dataset别让标签顺序变成黑匣子使用torchvision.datasets.ImageFolder是最快的启动方式但它有个隐性坑类别顺序按名称排序。如果你的目录是data/train/real和data/train/fakeImageFolder 会把fake映射到 0real映射到 1而不是你以为的 real 是 0。CrossEntropyLoss 不会因为这个顺序出错但你在计算准确率、混淆矩阵、以及最后给业务方解释时可能要处理一堆反逻辑的输出。所以我在这类项目里更愿意写一份非常轻的自定义 Dataset至少把标签映射钉死在代码里。# fake_real_dataset.py import os from PIL import Image from torch.utils.data import Dataset class FakeRealDataset(Dataset): # label: 0 表示 fake1 表示 real集中定义便于后续改阈值 CLASS_TO_LABEL {fake: 0, real: 1} def __init__(self, root, transformNone): self.samples [] # 保存 (路径, label) self.transform transform for cls_name in self.CLASS_TO_LABEL.keys(): cls_dir os.path.join(root, cls_name) if not os.path.isdir(cls_dir): continue for fname in sorted(os.listdir(cls_dir)): if fname.lower().endswith((.jpg, .jpeg, .png)): path os.path.join(cls_dir, fname) label self.CLASS_TO_LABEL[cls_name] self.samples.append((path, label)) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label这个类把真假标签显式定义在CLASS_TO_LABEL字典里后面推理时你可以直接按这个字典解释输出。convert(RGB)很重要因为真实素材里可能会混入 RGBA 的 PNG、灰度 JPEG不统一到三通道模型训练时会因为通道数不一致直接报错。在数据量比较大的场景还可以在__getitem__里加入读取失败保护如果Image.open抛异常就返回同类的其他样本。这样脚本不会因为单张损坏图片中途崩掉。我处理过一份从电商平台抓下来的图片集里面有约 0.3% 的图片是空文件如果没有这个保护训练跑两小时直接中断只能白等。3.3 数据增强的边界能翻转裁剪但别用高斯模糊真假图识别跟普通图像分类的增强策略有一个显著差异普通分类任务为了提升鲁棒性喜欢加高斯模糊、随机擦除、强色彩抖动因为猫即使被模糊了也还是猫。但真假图检测想要抓的是局部噪声和伪造伪影高斯模糊会把真伪差距磨平随机擦除可能会把唯一的伪造痕迹擦掉这些增强都要慎用。我平时会用一个偏保守的增强组合。# transforms.py from torchvision import transforms train_transform transforms.Compose([ # 随机裁剪出 0.8 到 1.0 的区域再缩放保留局部细节 transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(p0.5), # 真伪识别对色偏敏感所以饱和度抖动系数尽量小 transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2, hue0.05), transforms.ToTensor(), # ImageNet 归一化配合 resnet 预训练权重 transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])这里的RandomResizedCrop相当于把网络注意力强制分散到不同局部区域对“伪造区域只占画面一小块”的检测特别有效。旋转和透视变换要看业务场景商品图、证件照可以做小角度旋转但如果是带方向性的图像比如人脸、车牌旋转超过 10 度就可能导致语义颠倒反而学了错误先验。ColorJitter 的 hue 参数我一般控制在 0.05 以内。伪造算法经常在颜色统计上露马脚如果训练时把色相抖动设得太大模型会认为颜色失真也可以当真实图等于自己把有效特征削弱了。如果你担心部署时会遇到重新压缩可以在增强链路里加一层 JPEG 重压缩用PIL.Image.save(qualityrandom.randint(70, 95))保存到临时对象再读回来。这一步能让模型学会抵抗压缩伪影的变化代价是训练速度会慢一些但部署后效果通常更稳。4. 跑通一次完整的 CNN 训练ResNet 微调脚本逐行注释4.1 拿到压缩包以后先检查这三处再动手打开 zip 之后我建议不要直接双击运行先看三样东西依赖文件、数据读取方式、训练入口。依赖文件通常叫 requirements.txt里面会列出 PyTorch、torchvision、numpy 的版本范围。真假图识别一般不挑显卡CPU 也能跑但 PyTorch 版本和 torchvision 版本必须匹配否则 import torchvision 时会因为底层算子不兼容报错。数据读取方式要看代码里用的是 ImageFolder 还是自定义 Dataset前者要求目录结构严格按类名分好后者需要字段路径。训练入口一般叫 train.py 或者 main.py它决定了你后面要传哪些命令行参数。这个项目标题里明确说了“含逐行注释和说明文档”正常情况下说明文档会先讲依赖再讲如何准备目录最后讲训练命令。我下面给出的脚本骨架也是按 PyTorch 风格训练脚本最常见的写法来展开的你可以拿着它跟手头的源码逐行对比。4.2 训练主脚本拆解从数据加载到保存最优权重# train_fake_real.py import argparse import os import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from model_builder import build_model def parse_args(): parser argparse.ArgumentParser() parser.add_argument(--data-root, defaultdata) parser.add_argument(--arch, defaultresnet18, choices[resnet18, resnet50]) parser.add_argument(--epochs, typeint, default30) parser.add_argument(--batch-size, typeint, default32) parser.add_argument(--lr, typefloat, default1e-4) parser.add_argument(--weight-decay, typefloat, default1e-4) return parser.parse_args() def make_loader(root, split, batch_size, shuffle): # 训练和验证共用大部分 transform区别在于是否做随机增强 tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), if split train else transforms.Resize(256), ... ]) dataset datasets.ImageFolder(os.path.join(root, split), transformtf) loader DataLoader(dataset, batch_sizebatch_size, shuffleshuffle, num_workers4, pin_memoryTrue) return loader, dataset.classes def main(): args parse_args() torch.manual_seed(0) # 固定随机种子方便复现每次结果 train_loader, classes make_loader(args.data_root, train, args.batch_size, True) val_loader, _ make_loader(args.data_root, val, args.batch_size, False) model build_model(args.arch, num_classeslen(classes), pretrainedTrue) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() # AdamW 比 Adam 更适合权重衰减和预训练模型微调搭配更稳 optimizer torch.optim.AdamW(model.parameters(), lrargs.lr, weight_decayargs.weight_decay) # 余弦退火学习率到训练末尾学习率接近 0避免震荡 scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxargs.epochs) best_acc 0.0 for epoch in range(args.epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) scheduler.step() model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) val_acc correct / total print(fepoch {epoch1:03d} loss{running_loss/len(train_loader):.4f} fval_acc{val_acc:.4f}) # 只在验证集表现更好时保存权重防止最后一个 epoch 过拟合 if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), fbest_{args.arch}.pth) if __name__ __main__: main()这段脚本里值得逐行解释的地方不少。RandomResizedCrop(224)里的 224 与Resize(256)的 256 形成一组对比训练时随机裁剪验证时先放大到 256 再中心裁剪到 224这是 torchvision 官方微调里的常见配置能让验证过程更稳定。num_workers4表示用 4 个子进程加载图片避免 GPU 等待数据pin_memoryTrue能减少数据从内存搬运到显存的耗时CPU 机器上可以忽略。损失函数用的是CrossEntropyLoss它内部已经把 softmax 和 NLLLoss 合在一起所以模型输出层不需要额外加 softmax。如果你直接在训练代码里看到F.softmax那通常只是为了打印概率真正算 loss 时再用一次 softmax 会出错。优化器选AdamW而不是 Adam因为权重衰减在 Adam 里实现方式有偏差AdamW 把权重衰减和自适应学习率解耦对大模型微调更友好。学习率 1e-4 是 resnet 预训练权重微调的安全起点直接上 1e-3 很容易在前几个 epoch 出现 loss 上升。CosineAnnealingLR的T_max必须和epochs一致否则最后一个 epoch 不会真正退火到最小值附近。4.3 用验证集挑模型而不是看最后一个 epoch很多刚入门的朋友习惯把model.state_dict()在训练结束后保存这其实是个坑。最后一个 epoch 往往已经进入过拟合阶段验证集准确率可能正在下降。正确做法是在每个 epoch 结束后跑一次验证只有验证集准确率比历史最好值更高时才保存。我在上面的代码里用了best_acc记录历史最优然后用torch.save覆盖旧文件。这一步不复杂但能减少很多无效训练。验证时有一句代码容易踩坑outputs.argmax(dim1)。模型输出是一个二维张量形状是[batch_size, num_classes]dim1表示在类别维度上取最大值下标。如果误写成dim0你会得到每个类别上最大响应的样本序号结果完全错乱。类似的计算correct时preds labels返回布尔张量sum().item()转成整数注意不要直接sum()而不取.item()否则得到的是张量后续打印格式会很奇怪。4.4 从训练脚本到推理脚本输出概率而不是只输出标签业务系统最终需要的通常不是“真/假”标签而是“这张图有多大概率是假的”。所以推理脚本里不要直接argmax而是输出 softmax 概率。# inference.py import torch import torch.nn.functional as F from torchvision import transforms from PIL import Image from model_builder import build_model def predict(model, image_path, device): tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) img Image.open(image_path).convert(RGB) x tf(img).unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits model(x) prob F.softmax(logits, dim1) # prob[0][0] 是 fake 概率prob[0][1] 是 real 概率 return prob[0].cpu().tolist()unsqueeze(0)很重要模型期望的输入是四维张量[batch, channel, height, width]单张图片只有三维必须加一维 batch。F.softmax(logits, dim1)会把 logits 转换到 0 到 1 之间两个类别的概率加起来等于 1。最终返回的列表里下标 0 对应之前定义的 fake 类下标 1 对应 real 类。我建议在推理脚本里写一个明确注释防止一个月后自己都忘了顺序。5. 避坑排查真假图识别训练最容易翻车的 5 个常见问题5.1 现象训练 loss 一直不降准确率卡在 50% 附近出现这个现象时先把学习率打印出来看看。原因最常见的是学习率设置过大预训练模型的初始输出已经比较极端过大的更新一步就把之前学到的特征冲乱另一个原因是输入没有做 ImageNet 归一化直接把 0 到 255 的像素值喂给网络会让最开始的梯度变化幅度跟预训练分布差距太大。解决把--lr调到 1e-4 以下确认 transform 里包含Normalize(mean[0.485,0.456,0.406], std[0.229,0.224,0.225])。如果还不行把最后一个全连接层单独用一个更大的学习率卷积层用更小的学习率这种分组配置在迁移学习里能明显改善早期收敛速度。5.2 现象训练集准确率 99%验证集只有 70%这是典型的过拟合信号但并不一定是模型问题。原因训练集和验证集来自同一个数据源而且两张图片可能内容重叠。比如从同一段视频里连续抽帧或者同一张原图被轻微压缩后复制进两个目录模型看到的其实是非常相似的图片。另一个常见原因是训练集图片数量太少模型没有足够样本去泛化。解决做文件名 MD5 去重把完全相同的图片从验证集里剔除视频抽帧至少间隔 5 帧以上再就是增加 ResizedCrop 和颜色抖动人为打断“背题”路径。如果数据量实在不够干脆把模型从 ResNet50 降级到 ResNet18减少容量反而能让验证集更稳定。5.3 现象验证集分数不错但换一批新图后误判率暴增这类现象最让人头疼因为它不是代码错误而是训练分布和真实分布不一致。原因训练时过度使用了中心裁剪。验证集图片恰好是大头照伪造区域集中在脸中央模型只学会了看中间区域换到新场景伪造痕迹在边缘或者角落模型根本看不见。另外训练时全部用高质量原图部署时图片经过微信传输、平台压缩高频细节已经被抹掉模型面对的是分布外数据。解决训练时用RandomResizedCrop(scale(0.5,1.0))强制模型看不同空间位置增强链路里加入 JPEG 重压缩步骤模拟真实传输损耗。部署前的测试集最好独立收集一批没进过训练流程的真实业务图不要拿训练脚本里 split 出来的 test 目录自欺欺人。5.4 现象所有图片都被判成“真”或者都被判成“假”模型输出出现严重倾向性通常不是模型本身的问题而是训练数据分布出了问题。原因真实类和伪造类的数量差距过大比如 real 有 5000 张fake 只有 200 张。CrossEntropyLoss 在类别不平衡下会让网络倾向于预测样本量多的那一类。另外如果训练时数据加载器没开 shuffle每一轮迭代看到的类别顺序固定模型也会学到周期性的输出偏差。解决在DataLoader里设置shuffleTrue把类别权重传入CrossEntropyLoss(weighttorch.tensor([w_fake, w_real]))。更稳妥的做法是做类别均衡采样让每个 batch 里真实图和伪造图数量接近。验证指标不要只看 Accuracy要看每一类的召回率用混淆矩阵判断到底哪一类被牺牲了。5.5 现象GPU 占用率上不去一个 epoch 要跑半小时训练速度慢不一定是显卡差更多时候是数据加载和模型配置没匹配上。原因num_workers0导致数据加载单线程GPU 在等待图片读取也有可能是 batch size 太大导致每次反向传播前显存不足程序自动回退到极小的 batch 范围还有可能是开了pin_memory但数据转换全在 GPU 端做造成传输瓶颈。解决优先把num_workers调成 4 或 8并把pin_memoryTrue保留batch size 按显存余量调整ResNet18 在 8GB 显存下 batch size 32 是安全的ResNet50 则降到 16 更稳。如果显存还是不够可以在训练代码里加混合精度torch.cuda.amp.autocast()能省接近一半显存同时训练速度提升 20% 左右。6. 进阶技巧用 Grad-CAM 看模型注意力再用阈值校准收尾6.1 Grad-CAM让 CNN 指给你看它到底在关注哪里模型训完以后我一直习惯做一次可视化再交付因为真假图识别最怕“模型靠背景也能分类”。Grad-CAM 能利用特征图的梯度生成一份高亮热力图告诉我们模型决策时主要看画面哪个区域。# grad_cam.py import torch from torchvision import transforms def gradcam(model, image_path, device, target_layer): model.eval() # 注册 forward 和 backward hook临时抓取特征图和梯度 feat_map {} def forward_hook(module, input, output): feat_map[activation] output def backward_hook(module, grad_input, grad_output): feat_map[grad] grad_output[0] handle_forward target_layer.register_forward_hook(forward_hook) handle_backward target_layer.register_full_backward_hook(backward_hook) tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) x tf(Image.open(image_path).convert(RGB)).unsqueeze(0).to(device) output model(x) # 让梯度回传到目标层的中间特征 model.zero_grad() output[0, 0].backward() handle_forward.remove() handle_backward.remove() # 把特征图张量从计算图里拆出来避免无法切片 activation feat_map[activation].detach() grad feat_map[grad].detach() # 对梯度在空间维度求均值得到每个通道的权重 weights grad.mean(dim(2, 3), keepdimTrue) # 热力图 通道权重和特征图的加权和 cam (weights * activation).sum(dim1, keepdimTrue).relu() cam torch.nn.functional.interpolate(cam, size(224, 224), modebilinear) return cam.squeeze().cpu().numpy()以上代码直接复制就能用但要注意register_full_backward_hook的 PyTorch 版本要求如果你用的是 1.8 之前的版本需要改成register_backward_hook。生成的 heatmap 最好叠加在原图上保存如果发现模型只盯着角落的黑色背景而不是图片主体那就要检查数据增强是否把主体裁得太偏。6.2 阈值校准不要死等默认的 0.5二分类模型默认把输出概率 0.5 当作判定边界但这个边界往往不是最优选择。真伪检测业务通常更重视“别把伪造图放过去”所以宁可牺牲一点真实图的召回率也要保证伪造图被拦住。最佳阈值应该从验证集上统计出来的 F1 曲线里找。# calibrate_threshold.py import numpy as np from sklearn.metrics import precision_recall_curve # y_true 为真实标签y_prob 为模型预测的 fake 概率 prec, recall, thresholds precision_recall_curve(y_true, y_prob) f1_scores 2 * prec * recall / (prec recall 1e-9) best_idx int(np.argmax(f1_scores)) best_threshold thresholds[best_idx] # 这里把阈值设置成能使 F1 最大的值 final_threshold round(best_threshold, 4)precision_recall_curve返回的thresholds比prec和recall少一个元素所以直接用best_idx索引不会越界。我一般会把人工审核成本考虑进去如果审核端能承载比较高的人工量就把阈值往低调如果审核量有限就要求 F1 和准确率之间的一个平衡。这个参数最终要写进配置文件不要每次推理都从模型代码里去抠。6.3 我自己的一个训练习惯把随机种子、增强参数和阈值一起记下来这个项目里最磨人的不是模型而是复现和排错。我吃过大亏同一个脚本跑两次因为没固定随机种子第一次准确率 92%第二次只有 85%整个调参过程像在做玄学实验。后来养成一个习惯每次训练前生成一份训练记录 JSON把torch.manual_seed、dataloader 的 worker 种子、增强开关、学习率、阈值全部写进去之后无论是复现最好成绩还是排查新数据的分布偏移都能快速定位变量。这一招不花什么成本但对真假图识别这种特别依赖数据分布的项目来说比多调 10 个 epoch 都有价值。希望这个习惯也能帮到正在复现 resnet 真假图训练的同行少走一点我走过的弯路。本文还有配套的精品资源点击获取
返回列表