ARTICLE DETAIL

资讯详情

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

AI合成图片检测最小工程闭环:从ResNet18到实战

AI合成图片检测最小工程闭环:从ResNet18到实战 AI 生成媒体检测正在从研究论文里的实验指标变成实际业务里的工程需求。Grove Research 这类新研究机构在近期亮相长期关注合成媒体问题的技术账号 deepfates 也重新出现在讨论中无论这些消息背后的具体组织关系如何技术社区真正关心的仍然是同一个问题当一张图片、一段语音、一个视频由模型生成时工程上应该用什么手段把它识别出来。本文不追踪具体公司的组织架构也不讨论人员变动而是把这则消息当作一个切入口完整走一遍 AI 合成图片检测的最小工程闭环。你会看到检测任务为什么不是简单的图像分类、环境需要准备到什么程度、一个可训练的检测器怎么写、训练后如何验证效果以及上线前哪些坑最容易踩。完成之后你得到的不是一篇概念科普而是一套可以替换数据后继续迭代的检测代码骨架。1. 先理解合成媒体检测到底要解决什么问题1.1 从 deepfake 到通用合成媒体deepfake 最初指用深度学习生成的换脸视频和图片核心是把一个人的脸换到另一个人的身体上。随着扩散模型和自回归图像模型的普及问题已经扩展到整张人脸生成、场景合成、局部编辑、语音克隆和多模态换声。现在更准确的说法是 AI 生成媒体或者合成媒体。检测任务的目标也发生了变化。早期检测主要针对特定生成器的固定伪造痕迹比如脸部边缘过渡不自然、眨眼频率异常、肤色不均。现在模型可以一次性生成高分辨率图像肉眼很难找到明显破绽检测器必须从更底层的特征入手例如局部噪声分布、频率域残留、物体结构不一致等。这里最关键的技术判断是检测器和生成器之间存在持续的对抗关系。生成模型迭代一次检测模型往往也要重新训练或调整特征提取方式。因此任何检测方案都必须预留“数据更新、模型重训、效果复评”的环节不能当成一次性交付。1.2 检测任务的三条技术路线实际工程中合成图片检测通常从三个层面同时考虑而不是只选一条路走到黑。第一条是基于空间特征的分类路线。把图片缩放成固定尺寸用卷积神经网络提取特征最后输出 real 和 fake 两个类别的概率。这条路线实现成本低、训练速度快适合作为第一个可运行的基线模型。缺点是容易被后处理和压缩干扰也容易过拟合到数据来源。第二条是频域和噪声分析。真实图片经过成像传感器处理会留下固定的噪声模式生成模型合成的图片在频域上往往过于平滑或者在高频分量上呈现出与真实图片不同的统计规律。常见做法是把图片做离散余弦变换或小波变换后对系数分布进行分析。这条路线对压缩更鲁棒但特征工程难度更高需要实验经验。第三条是取证溯源。C2PA、Content Credentials 这类标准让相机和编辑软件在图片元数据里写入签名、来源和编辑记录。检测器结合元数据可以判断图片是否经过 AI 编辑。它的优点是可信度高缺点是只有主动记录元数据的设备和软件才会留下痕迹作用范围有限。实际项目常常是三条路线组合使用先跑分类基线再用频域分析覆盖压缩场景最后接入溯源标准处理有元数据的素材。1.3 为什么一条简单 CNN 值得先跑通很多初学者一上来就想使用大规模视觉模型或专门针对 deepfake 的复杂网络结果卡在环境配置和算力不足上连一个像样的准确率都拿不出来。正确的做法是先跑通最基础的 CNN 检测器。它虽然不会是最强方案但能帮你完成三件事验证数据标注是否可靠、确认训练推理链路是否通畅、拿到一个可以对照的准确率基线。后续无论换成更大模型还是接入频率分析模块都能用这个基线评估新方案是否真的有效。本文后面给出的就是一个基于 ResNet18 的最小检测器它的价值在于结构完整、可替换、可复现。2. 准备环境最小依赖搭出检测项目2.1 版本与依赖选择项目以 Python 3.10 为基础深度学习框架选 PyTorch。选择 PyTorch 的原因是生态成熟torchvision自带常用网络结构和 ImageNet 预训练权重适合快速做迁移学习。在正式安装前先确认本机是否具备 NVIDIA GPU。没有 GPU 也可以运行只是训练会慢很多。建议先创建独立的虚拟环境避免和系统环境或其它项目冲突。python -m venv .venv source .venv/bin/activate然后在虚拟环境中安装依赖。以下是本文代码对应的最小依赖清单pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install pillow numpy tqdm matplotlib如果使用 CPU 环境去掉--index-url参数直接安装默认发行版即可pip install torch torchvision pillow numpy tqdm matplotlib依赖用途版本建议torch模型定义、自动求导、训练2.0 及以上torchvision数据集加载、预训练模型0.15 及以上pillow图片读取与转换9.0 及以上numpy数组操作与概率计算1.24 及以上注意PyTorch 版本与 CUDA 版本必须匹配。用python -c import torch; print(torch.cuda.is_available())检查输出True才说明 GPU 可用。2.2 数据集组织方式检测器训练需要成对的真实图片和合成图片。真实图片可以来自公开人像数据集也可以从自己业务场景中收集合成图片则需要覆盖多种生成方式比如不同的人脸生成模型、图像编辑模型、变脸工具等。本文使用torchvision.datasets.ImageFolder加载数据它要求数据按照下面的目录结构组织data/ train/ real/ fake/ val/ real/ fake/ test/ real/ fake/ImageFolder会自动按子目录名把图片映射为类别标签。这里的关键点是子目录名会按字母顺序排序fake在字母序上排在real前面因此类别索引是0fake, 1real。后续推理脚本中的类别名列表必须与这个顺序保持一致否则会出现预测结果倒置的问题。图片数量方面学习环境可以先用每个类别 200 到 500 张图片跑通流程生产级效果则需要每个类别数万张并且按来源、分辨率、压缩率分层采样。2.3 项目结构工程上推荐把训练脚本、推理脚本、依赖清单和输出目录分开组织方便后续替换数据和模型版本。fake_detector/ train.py infer.py requirements.txt data/ train/ val/ test/ output/output目录存放训练得到的模型权重、评估结果和日志。不要直接把权重文件提交到代码仓库大文件应使用模型仓库或对象存储管理。3. 核心实现写一个可训练的合成图片检测器3.1 数据读取与预处理图片在进入网络之前必须经过统一的预处理流程包括缩放、归一化和数据增强。归一化的均值和标准差取自 ImageNet 统计值因为接下来使用的预训练模型就是在这些参数下训练的推理阶段必须沿用同一套参数。from torchvision import transforms, datasets def build_transforms(size224, trainTrue): if train: return transforms.Compose([ transforms.Resize((size, size)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(10), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) return transforms.Compose([ transforms.Resize((size, size)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])训练阶段加入随机水平翻转和随机旋转是为了让模型不依赖图片的方向信息增强泛化能力。验证阶段不加入随机增强保证每次评估结果稳定可复现。3.2 模型结构选择本文使用 ResNet18并在 ImageNet 预训练权重基础上做微调。之所以用预训练模型是因为真实世界图片的低层特征例如边缘、纹理、颜色分布具有共性。预训练权重让模型从这些通用特征出发只用少量合成媒体数据就能学习到 fake 和 real 的差异。import torch.nn as nn from torchvision import models def build_model(num_classes2): model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) return model代码中只替换了最后一层全连接层输出从 1000 类改成 2 类。训练时整网参与微调学习率需要设置得比从头训练低一般取1e-4到3e-4。3.3 训练脚本训练脚本的核心逻辑是加载 ImageFolder 数据、构建模型、定义交叉熵损失、使用 Adam 优化器、逐轮更新权重并输出验证集准确率。import os import argparse import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms def build_transforms(size224, trainTrue): if train: return transforms.Compose([ transforms.Resize((size, size)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(10), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) return transforms.Compose([ transforms.Resize((size, size)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) def build_model(num_classes2): model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) return model def main(): parser argparse.ArgumentParser() parser.add_argument(--data_dir, defaultdata) parser.add_argument(--epochs, typeint, default10) parser.add_argument(--batch_size, typeint, default32) parser.add_argument(--lr, typefloat, default1e-4) parser.add_argument(--num_workers, typeint, default2) parser.add_argument(--device, defaultcuda) args parser.parse_args() device torch.device(args.device if torch.cuda.is_available() else cpu) train_set datasets.ImageFolder( os.path.join(args.data_dir, train), transformbuild_transforms(trainTrue) ) val_set datasets.ImageFolder( os.path.join(args.data_dir, val), transformbuild_transforms(trainFalse) ) train_loader DataLoader(train_set, batch_sizeargs.batch_size, shuffleTrue, num_workersargs.num_workers) val_loader DataLoader(val_set, batch_sizeargs.batch_size, shuffleFalse, num_workersargs.num_workers) model build_model().to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lrargs.lr) for epoch in range(args.epochs): model.train() train_loss 0.0 train_correct 0 train_total 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() train_loss loss.item() * images.size(0) _, preds torch.max(outputs, 1) train_correct (preds labels).sum().item() train_total labels.size(0) model.eval() val_loss 0.0 val_correct 0 val_total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) loss criterion(outputs, labels) val_loss loss.item() * images.size(0) _, preds torch.max(outputs, 1) val_correct (preds labels).sum().item() val_total labels.size(0) print(fEpoch {epoch1}/{args.epochs} ftrain_loss{train_loss/train_total:.4f} ftrain_acc{train_correct/train_total:.4f} fval_loss{val_loss/val_total:.4f} fval_acc{val_correct/val_total:.4f}) os.makedirs(output, exist_okTrue) torch.save(model.state_dict(), output/fake_detector.pth) print(saved to output/fake_detector.pth) if __name__ __main__: main()训练过程中需要同时观察两个指标。训练准确率反映模型对训练集的拟合程度验证准确率反映模型对未见数据的泛化能力。如果训练准确率持续上升而验证准确率停滞不前说明模型过拟合应该增加数据量、增强数据增强强度或者提前停止训练。运行训练的命令如下python train.py --data_dir data --epochs 10 --batch_size 32 --lr 1e-4 --device cuda3.4 推理脚本推理脚本加载训练好的权重对单张图片输出类别和置信度。推理时的预处理必须与训练验证阶段完全一致特别是归一化参数和图片尺寸否则模型预测会明显退化。import argparse import torch import torch.nn as nn from PIL import Image from torchvision import transforms, models def build_model(num_classes2): model models.resnet18(weightsNone) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) return model def main(): parser argparse.ArgumentParser() parser.add_argument(--checkpoint, defaultoutput/fake_detector.pth) parser.add_argument(--image, requiredTrue) parser.add_argument(--device, defaultcuda) args parser.parse_args() device torch.device(args.device if torch.cuda.is_available() else cpu) model build_model().to(device) model.load_state_dict(torch.load(args.checkpoint, map_locationdevice)) model.eval() 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]), ]) image Image.open(args.image).convert(RGB) x transform(image).unsqueeze(0).to(device) with torch.no_grad(): logits model(x) prob torch.softmax(logits, dim1) idx int(torch.argmax(prob, dim1)) labels [fake, real] print(flabel{labels[idx]} prob{prob[0][idx].item():.4f}) if __name__ __main__: main()运行命令python infer.py --checkpoint output/fake_detector.pth --image test/real/001.jpg这里有一个非常容易踩的坑ImageFolder的类别顺序按目录名排序所以fake是索引 0real是索引 1。推理脚本中的labels [fake, real]必须与之对应否则会把真实图片预测成合成图片。4. 关键参数与评估指标怎样判断检测器真的有效4.1 训练参数速查参数含义示例值调大的影响调小的影响--epochs完整遍历训练集的轮数10拟合更充分但容易过拟合欠拟合验证准确率低--batch_size每次参数更新使用的图片数32梯度更稳定显存占用高梯度噪声大训练震荡--lr学习率1e-4收敛快但可能越过最优点收敛慢容易停在局部最优--num_workers数据加载进程数2数据供应更快内存压力大CPU 成为瓶颈GPU 等待调参顺序建议是先用小的epochs跑通流程再根据曲线决定是否增加轮数。learning rate优先从1e-4开始如果训练 loss 出现明显震荡降到3e-5。batch_size受显存限制8GB 显存建议不超过 64。4.2 评估指标不是只有准确率在 real 与 fake 类别失衡时准确率会掩盖模型的实际能力。比如测试集里 95% 是真实图片模型把所有图片都预测为 real准确率也有 95%但实际检测能力为零。因此至少要同时观察四个指标指标含义关注点Precision预测为 fake 的图片中有多少真为 fake高代表误报少业务误伤低Recall真正的 fake 图片被识别出多少高代表漏报少安全覆盖好F1Precision 和 Recall 的调和平均整体检测能力AUC不同阈值下的综合分类能力不依赖具体阈值适合模型选型业务场景需要根据成本决定偏向。如果检测结果用于人工复核宁可提高召回率把可疑图片都送进去如果检测结果直接触发封禁或拦截就要提高精确率控制误报。4.3 最容易低估的一个问题很多入门项目把 train、val、test 数据按文件列表随机打散结果验证集和训练集里包含了来自同一批生成模型的图片。模型真实记忆的是“这批生成器的特征”而不是“合成媒体的通用特征”。一旦换成新生成器准确率立刻大幅下降。正确的划分方式是按照图片来源划分同一生成模型产出的图片全部进入同一个集合不能一部分在训练集、一部分在测试集。这样才能判断模型是否真的具备跨生成器泛化能力。5. 运行验证把训练到推理的闭环跑通5.1 训练过程输出在图片数量充足、数据标注正确的前提下训练输出大致会呈现如下趋势。下面输出用于说明格式实际数值取决于数据集规模和质量。Epoch 1/10 train_loss0.6241 train_acc0.6542 val_loss0.4378 val_acc0.8120 Epoch 2/10 train_loss0.3018 train_acc0.9120 val_loss0.2215 val_acc0.9304 ... Epoch 10/10 train_loss0.0823 train_acc0.9741 val_loss0.0936 val_acc0.9615 saved to output/fake_detector.pth如果第一轮 val_acc 就接近 100%反而要警惕。这通常说明数据划分有问题或者真实图片和合成图片在目录层面就存在明显的外观差异模型学到的不是伪造痕迹而是背景、色调、拍摄设备等无关特征。5.2 推理验证预期对测试集中的真实图片运行推理预期输出labelreal prob0.9821对测试集中的合成图片运行推理预期输出labelfake prob0.9934工程上不要只验证一张图。建议把测试集全部图片跑一遍输出混淆矩阵分别统计 real 被误判为 fake 的数量和 fake 被漏判为 real 的数量。5.3 一个快速验证脚本可以用下面的简单脚本遍历测试集统计准确率import os import torch from PIL import Image from torchvision import transforms from tqdm import tqdm def evaluate_folder(model, image_dir, label, transform, device): correct 0 total 0 for name in tqdm(os.listdir(image_dir)): path os.path.join(image_dir, name) img Image.open(path).convert(RGB) x transform(img).unsqueeze(0).to(device) with torch.no_grad(): prob torch.softmax(model(x), dim1) pred int(torch.argmax(prob, dim1)) total 1 if (label fake and pred 0) or (label real and pred 1): correct 1 return correct, total这个脚本的价值不在于效率而在于快速确认模型在每个类别上的表现从而判断是否存在偏向某一类的现象。6. 常见问题排查按链路逐层定位6.1 现象、原因与处理对照问题现象常见原因检查方式处理建议训练 loss 不下降学习率过大或过小、数据未归一化、标签错误打印首个 batch 的 loss确认其能在若干步内下降学习率调整到 1e-4 量级检查 transform 是否包含 Normalize验证准确率极高而实战失效训练集和验证集同源模型记住来源查看 val 和 test 图片的生成来源按来源划分数据不以文件列表随机划分推理结果全部是同一个类别类别索引与代码顺序不一致打印train_set.classes查看类别映射让推理脚本类别列表与class_to_idx一致显存不足batch_size 过大、图片尺寸过大观察nvidia-smi显存占用降低 batch_size或把图片缩到 224CPU 推理非常慢模型较大、未做推理优化统计单张推理耗时转 ONNX 格式或换更小模型压缩后的图片误判严重训练数据缺乏压缩样本对比原图和压缩图的预测结果训练数据加入多档 JPEG 压缩增广6.2 一个典型排查链路假设出现“验证集准确率 95%但新来的图片预测错误”的现象。第一步检查数据来源。确认验证集和训练集是否来自同一批生成器如果是先重新划分数据。第二步检查预处理。确认推理脚本的尺寸、均值、标准差与训练时一致差异会直接改变输入分布。第三步检查测试图片本身。将图片保存后重新打开确认是否经过了重压缩、裁剪、加水印等后处理。这一类操作会抹掉部分伪造痕迹也会引入新的图像处理特征。第四步检查概率输出。如果模型对错误图片输出的 confidence 都很接近 0.5说明模型自身不确定而不是简单阈值问题需要考虑增加该类图片的训练样本。如果 confidence 接近 1.0说明模型对某个特征非常自信而这个特征在真实环境里不具备代表性。排查顺序的优先级是输入是否正确路径和命名是否正确依赖版本是否匹配配置是否生效数据分布是否合理最后才是模型结构问题。注意不要只验证程序能启动。检测类的项目必须验证输入图片、输出概率、错误类别三个层面都符合预期否则一个标签顺序错误就能让整个系统失效。7. 从 Demo 到生产工程化边界与研究趋势7.1 学习环境与生产环境的差异上面的代码在单机、离线、数据量适中的场景下可以正常工作但直接搬到生产会面临一系列问题。生产环境至少需要补齐以下能力训练数据的版本管理每次重训都要记录数据来源、采样策略和标注规则否则模型效果变化时无法定位原因。模型版本管理检测器要支持灰度发布新模型先在少量流量上观察误报率再逐步放开。日志与监控记录每次检测的置信度、图片来源、耗时方便事后分析和误报追查。推理性能约束在线检测接口通常要求在几百毫秒内返回结果ResNet18 在 GPU 上可以满足在 CPU 上则要考虑量化或用更小模型。回滚机制生成器升级导致检测准确率骤降时需要快速回退到上一版本模型并保留现场样本用于重训。如果业务场景对隐私要求高还需要考虑图片是否允许离线存储、日志中是否要脱敏、人工复核时是否展示原图。7.2 溯源标准与多模态检测单靠图像分类模型做检测本质上是在和生成器拼迭代速度永远存在滞后。更稳的方案是结合溯源标准让真实内容的产生过程留下可验证的元数据。C2PA 规范定义了内容凭证在图片生成或编辑时写入签名验证方通过公钥确认内容是否经过修改。工程上推荐采用“分类模型 溯源校验”的分层策略第一层用模型输出可疑度第二层用元数据校验真伪两层都通过才标记为可信内容。语音和视频场景可以复用同样的分层思路额外加入音频特征分析和帧间一致性检查。这些方向正是新研究机构和技术社区持续投入的原因。检测任务不会因为某个模型出现就结束它会随着生成技术的演进不断产生新的对抗样本、新的评估基准和新的工程需求。7.3 可复用的工程检查清单无论是学习项目还是生产项目上线前都可以对照以下清单逐项检查数据来源是否明确real 和 fake 是否有误标。数据集是否按来源划分而不是按文件列表随机划分。测试集是否覆盖多种生成工具、多种分辨率、多种压缩率。训练与推理的预处理是否完全一致。类别映射是否与ImageFolder的排序顺序一致。评估是否包含 Precision、Recall、F1、AUC而不仅是准确率。训练参数、数据版本、模型版本是否记录在日志中。推理接口是否有超时、限流、失败降级和异常处理。新模型上线前是否有灰度计划、回滚方案和误报监控。是否定期用新生成器产出的样本重新评估模型效果。这套清单也可以直接当作代码 Review 的检查项。检测类项目最容易出问题的环节不是模型网络本身而是数据划分、类别映射和上线后的持续维护这三项做好了替换更强模型只是时间问题。
返回列表