将MNIST的FGSM对抗训练代码迁移到CelebA数据集的报错求助
问题描述
正在进行对抗训练,使用基于PyTorch在MNIST数据集上训练模型的代码,希望将数据集替换为CelebA,但运行原代码时出现如下错误:RuntimeError: Given groups=1, weight of size [6, 1, 5, 5], expected input[128, 3, 218, 178] to have 1 channels, but got 3 channels instead
已尝试操作
修改第一层卷积的输入通道数,将:self.conv1 = torch.nn.Conv2d(1, 6, 5, padding=2)
改为:self.conv1 = torch.nn.Conv2d(3, 6, 5, padding=2)
但此时又在x=x.view(-1, 16*5*5)行出现新错误:
RuntimeError: shape '[-1, 400]' is invalid for input of size 4472832
尝试将-1改为3后出现其他错误,因不清楚原理停止修改。相关核心代码片段:
class LeNet5(torch.nn.Module): def __init__(self): super(LeNet5, self).__init__() self.conv1 = torch.nn.Conv2d(1, 6, 5, padding=2) self.conv2 = torch.nn.Conv2d(6, 16, 5) self.fc1 = nn.Linear(16*5*5, 120) self.fc2 = nn.Linear(120, 84) self.fc3 = nn.Linear(84, 10) def forward(self, x): x = F.relu(self.conv1(x)) x = F.max_pool2d(x, 2) x = F.relu(self.conv2(x)) x = F.max_pool2d(x, 2) x = x.view(-1, 16*5*5) x = F.relu(self.fc1(x)) x = F.relu(self.fc2(x)) x = self.fc3(x) return F.log_softmax(x,dim=-1) # 数据加载部分 train_loader = torch.utils.data.DataLoader( datasets.CelebA('data', split='train', transform=transforms.ToTensor(), download="True"), batch_size=128, shuffle=True)
需求
- 理解并解决上述错误,使CelebA数据集适配该代码;
- 解决上述问题后,针对CelebA的对抗训练是否还需其他修改?
- 最终目标是在人脸识别数据集上实现基于PyTorch的FGSM对抗训练,若当前方案不合适,希望获取相关资源或实现建议。
1. 错误原因及解决方法
错误根源
原LeNet5是为单通道28x28的MNIST图片设计的,而CelebA是3通道218x178的人脸图片,两者的输入尺寸、通道数、任务类型(MNIST是10分类,CelebA是40个属性的多标签分类)完全不同,导致两个核心错误:
- 通道数不匹配:原conv1输入通道为1,CelebA是3通道;
- 特征图尺寸不匹配:卷积池化后的特征图尺寸远大于MNIST的5x5,导致全连接层输入维度计算错误。
具体修改步骤
步骤1:修正卷积层后的特征图维度计算
计算CelebA图片经过卷积池化后的特征图尺寸:
- 输入尺寸:
3x218x178 - conv1(3→6通道,5x5卷积,padding=2):输出尺寸保持
218x178,maxpool2x2后变为109x89 - conv2(6→16通道,5x5卷积,无padding):输出尺寸为
(109-5+1)x(89-5+1)=105x85,maxpool2x2后变为52x42(向下取整) - 特征图总元素数:
16 * 52 * 42 = 34944
步骤2:修改模型结构
class LeNet5_CelebA(torch.nn.Module): def __init__(self): super(LeNet5_CelebA, self).__init__() # 修正输入通道为3 self.conv1 = torch.nn.Conv2d(3, 6, 5, padding=2) self.conv2 = torch.nn.Conv2d(6, 16, 5) # 修正全连接层输入维度为34944 self.fc1 = nn.Linear(34944, 120) self.fc2 = nn.Linear(120, 84) # CelebA是40个属性的多标签分类,输出维度改为40 self.fc3 = nn.Linear(84, 40) def forward(self, x): x = F.relu(self.conv1(x)) x = F.max_pool2d(x, 2) x = F.relu(self.conv2(x)) x = F.max_pool2d(x, 2) # 修正view的维度为34944 x = x.view(-1, 34944) x = F.relu(self.fc1(x)) x = F.relu(self.fc2(x)) x = self.fc3(x) # 多标签分类用sigmoid代替softmax return torch.sigmoid(x)
步骤3:修正损失函数
原代码用F.nll_loss(适用于单标签分类),CelebA是多标签任务,需替换为二元交叉熵损失:
在trainTorch函数中,将:
loss = F.nll_loss(preds, ys)
改为:
# BCEWithLogitsLoss可以直接处理未经过sigmoid的输出,若模型用sigmoid则用BCELoss loss = F.binary_cross_entropy(preds, ys.float())
注意:CelebA的target是40维的二进制张量,需要转为float类型。
步骤4:可选:缩小图片尺寸(避免内存溢出)
218x178的图片会导致特征图过大,全连接层参数过多,建议在数据加载时添加Resize变换:
transform = transforms.Compose([ transforms.Resize((64, 64)), transforms.ToTensor() ]) train_loader = torch.utils.data.DataLoader( datasets.CelebA('data', split='train', transform=transform, download=True), batch_size=128, shuffle=True)
若用64x64尺寸,重新计算特征图维度:
- conv1后64x64 → maxpool→32x32
- conv2后32-5+1=28 →28x28 →maxpool→14x14
- 总元素数:161414=3136,此时fc1输入改为3136即可。
2. CelebA对抗训练的其他必要修改
- FGSM适配多标签任务:原FGSM代码是针对单标签损失计算梯度,需修改为基于多标签损失计算输入的梯度。核心逻辑不变,但损失函数要对应二元交叉熵:
在攻击时,计算损失用二元交叉熵,然后对输入求导。def fgsm_attack(image, epsilon, data_grad): sign_data_grad = data_grad.sign() perturbed_image = image + epsilon*sign_data_grad perturbed_image = torch.clamp(perturbed_image, 0, 1) return perturbed_image - 调整训练参数:CelebA数据集更大,图片尺寸更大,需调小batch_size(如32)、增加训练轮次(如10-20轮),或使用GPU加速。
- 数据预处理:添加归一化(如用ImageNet均值方差),提升模型稳定性:
transform = transforms.Compose([ transforms.Resize((64,64)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])
3. 人脸识别数据集的FGSM对抗训练建议
当前方案(LeNet+CelebA)不适合人脸识别,因为CelebA是属性分类任务,而非身份识别。以下是针对性建议:
数据集选择
使用专门的人脸识别数据集:
- LFW(Labeled Faces in the Wild):小型数据集,适合入门;
- CASIA-WebFace:包含大量人脸身份,适合训练;
- MS-Celeb-1M:大规模数据集,需处理数据清洗问题。
模型选择
LeNet的容量不足以处理人脸识别任务,建议使用:
- 预训练CNN模型:如ResNet50、MobileNetV2,在ImageNet上预训练后,替换最后一层全连接层为身份分类的输出维度;
- 人脸识别专用模型:如ArcFace、SphereFace,这些模型针对人脸识别的特征匹配优化,对抗训练效果更好。
FGSM对抗训练实现要点
- 损失函数:人脸识别常用交叉熵损失(身份分类)或ArcFace损失;
- 对抗样本生成:基于分类损失计算输入梯度,生成FGSM扰动,核心逻辑与单标签分类一致;
- 库工具:直接使用
torchattacks库,内置FGSM、PGD等多种对抗攻击实现,无需手动编写:from torchattacks import FGSM attack = FGSM(model, eps=0.03) perturbed_images = attack(images, labels) - 训练流程:交替训练正常样本和对抗样本,或在每轮训练中混合两种样本,提升模型的对抗鲁棒性。
内容的提问来源于stack exchange,提问作者ItWorksOnLocal

