联邦学习对抗攻击实战:符号翻转与后门注入课程大作业
发布时间:2026/9/28 12:25:45来源:尧图网络
简介这份资源是面向计算机、人工智能、通信工程等专业学生与教师的联邦学习对抗攻击课程大作业完整方案适合作为毕设、课设或项目立项演示的参考实现。项目围绕联邦学习模型展开对抗攻击实验包含攻击脚本、孪生网络结构、联邦基础函数与模型数据初始化等模块并配有详细注释便于理解攻击流程与模型交互逻辑。压缩包共18个文件以10个pth模型权重、6个py源码、1个md说明文档和1个license许可文件为主整体约623KB体积轻量但结构完整下载后可直接运行验证。目前已有58人学习关注作者提供远程教学支持代码均经测试运行成功答辩评审平均分达96分。读者可从中获得可复现的联邦学习对抗攻击实验代码、预训练模型与排错思路也能在此基础上修改扩展完成其他功能或二次开发。1. 联邦学习遇上对抗攻击课程大作业里最容易被忽略的那条暗线做联邦学习课程大作业的人十有八九会把精力砸在「怎么把 FedAvg 跑通」上等准确率曲线好看了就收工。但如果你翻一翻近两年安全方向的会议论文会发现另一条更值得写的暗线联邦学习本身并不安全参与方上传的模型更新可以被恶意构造从而在全局模型里埋下后门或直接拉低精度。这就是「基于联邦学习模型实现的对抗攻击」这个题目真正的价值所在——它不是让你复现一个训练流程而是让你亲手造一次攻击再观察全局模型怎么被污染。这个方向适合三类人一是要交课程大作业、想拿高分而不是交个 MNIST 分类了事的学生二是刚接触联邦学习、想搞懂「聚合」这一步到底脆弱在哪的工程师三是做安全评估、需要一套可复现攻击基线的人。整篇内容围绕一份 Python 源码展开包含详细注释和模型定义我会把数据怎么切、攻击怎么注入、参数怎么调、哪里最容易翻车讲清楚让你能照着跑出结果而不是对着论文干瞪眼。2. 先搞清楚攻击面联邦学习里到底哪一步能被动手脚2.1 联邦学习的标准流程与三个可攻击位置联邦学习的经典流程是服务器下发全局模型各客户端用本地数据训练若干轮把模型更新梯度或权重回传服务器聚合成新的全局模型。这个流程里攻击者能下手的地方有三处。第一处是数据层客户端本地数据被投毒比如把某些样本的标签翻转或者插入带触发器的后门样本。第二处是更新层客户端在回传更新前直接篡改数值这是最常被研究的因为它不需要控制数据分布只要控制代码就行。第三处是聚合层如果攻击者能影响服务器端的聚合逻辑比如伪造多个恶意客户端身份就能放大攻击效果。课程大作业通常选第二处因为实现成本最低你只要在客户端本地训练完之后对梯度做一次符号翻转或者加噪就能观察全局模型的变化。常见做法是让一部分客户端成为「恶意客户端」比例一般设在 10% 到 30% 之间太低看不出效果太高就变成明着破坏、失去隐蔽性。2.2 为什么选符号翻转和后门注入作为攻击实现对抗攻击在联邦学习里分两大类非目标攻击和目标攻击。非目标攻击就是让全局模型精度掉下去最粗暴的是符号翻转Sign Flipping把梯度乘以 -1 再上传聚合时正常客户端的更新会被抵消。目标攻击是让模型在特定输入上输出攻击者想要的标签典型是后门攻击在本地数据里混入带触发器的样本训练后模型学到「看到触发器就输出指定类别」。选这两个作为大作业实现理由是它们代码量小、效果直观、容易画图对比。符号翻转只需要一行grad -grad后门注入只需要在 Dataset 里加一个 patch 函数。而且两者可以放在同一套框架里对比一个看精度曲线怎么崩一个看后门成功率怎么涨。下面这张表是我一般用来给同学讲清楚区别的攻击类型攻击目标实现位置典型指标隐蔽性符号翻转降低全局精度客户端更新层测试准确率下降幅度低聚合时容易被检测后门注入特定输入误分类客户端数据层后门成功率 ASR高正常精度几乎不变梯度加噪降低收敛速度客户端更新层收敛轮数增加中取决于噪声方差模型替换完全控制全局模型客户端更新层全局模型行为偏移低需要多轮持续攻击提示课程大作业里不要同时上三种攻击选一到两种做深把对比实验做扎实比堆一堆半成品强得多。2.3 环境准备与依赖安装的最小命令集源码是 Python 写的依赖主要是 PyTorch 和 NumPy。我一般建议用 conda 建一个干净环境避免和系统里的包打架。下面这套命令在 Linux 和 Windows 的 WSL 下都验证过conda create -n fl_attack python3.9 -y conda activate fl_attack pip install torch2.0.1 torchvision0.15.2 --index-url https://download.pytorch.org/whl/cpu pip install numpy1.24.3 matplotlib3.7.1如果你有 GPU把cpu换成对应的cu118版本即可。这里固定版本号是因为联邦学习代码里经常用到torch.autograd的手动梯度操作不同版本对retain_graph的处理有差异版本飘了容易出玄学 bug。装完之后用python -c import torch; print(torch.__version__)确认一下输出 2.0.1 就对了。数据方面源码默认用 MNIST 和 CIFAR-10这两个数据集 torchvision 会自动下载。如果你的网络环境下载慢可以提前把./data目录准备好或者改downloadFalse指向本地路径。注意别用太新的 torchvision它默认的下载源有时候会变固定版本能省掉不少麻烦。3. 把源码跑起来数据切分、模型定义与联邦训练主循环3.1 非独立同分布切分用 Dirichlet 分布模拟真实客户端联邦学习最怕的就是数据独立同分布假设真实场景里每个客户端的数据分布都不一样。源码里用的是 Dirichlet 分布切分这是目前最主流的模拟方式。核心参数是alpha越小表示分布越不均衡一般取 0.1、0.5、1.0 做对比实验。import numpy as np def dirichlet_split(labels, num_clients, alpha0.5, seed42): 按 Dirichlet 分布把数据集切分给多个客户端 labels: 所有样本的标签数组 num_clients: 客户端数量 alpha: 浓度参数越小越不均衡 np.random.seed(seed) num_classes len(np.unique(labels)) # 为每个类别生成一个 Dirichlet 分布决定该类样本分给各客户端的比例 proportions np.random.dirichlet([alpha] * num_clients, num_classes) client_indices [[] for _ in range(num_clients)] for cls in range(num_classes): cls_indices np.where(labels cls)[0] np.random.shuffle(cls_indices) # 按比例切分用 cumsum 确定切分点 split_points (np.cumsum(proportions[cls]) * len(cls_indices)).astype(int) split_points np.clip(split_points, 0, len(cls_indices)) prev 0 for cid in range(num_clients): client_indices[cid].extend(cls_indices[prev:split_points[cid]]) prev split_points[cid] return client_indices这段代码的逻辑是对每个类别单独做一次 Dirichlet 采样得到该类样本在各客户端上的分配比例再按比例切分索引。alpha0.5时有的客户端可能拿到某类的大部分样本有的几乎拿不到这就模拟了真实的不均衡。参数上num_clients一般设 10 或 20alpha设 0.1 时极端不均衡适合观察攻击在恶劣条件下的效果。如果你发现某个客户端样本数为 0把alpha调大一点或者换种子。3.2 模型定义一个带详细注释的 CNN 与攻击注入点源码里的模型是一个轻量 CNN两层卷积加两层全连接参数量在 100 万左右CPU 上跑 MNIST 一轮大概几秒。关键是在客户端本地训练函数里预留了攻击注入点。import torch import torch.nn as nn import torch.nn.functional as F class SimpleCNN(nn.Module): def __init__(self, num_classes10): super().__init__() self.conv1 nn.Conv2d(1, 32, 3, padding1) self.conv2 nn.Conv2d(32, 64, 3, padding1) self.pool nn.MaxPool2d(2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, num_classes) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) return self.fc2(x) def local_train(model, dataloader, epochs1, lr0.01, attackNone): 客户端本地训练attack 参数控制是否注入攻击 attacksign_flip 时对梯度取反 attackbackdoor 时使用带触发器的数据 optimizer torch.optim.SGD(model.parameters(), lrlr) model.train() for _ in range(epochs): for data, target in dataloader: optimizer.zero_grad() output model(data) loss F.cross_entropy(output, target) loss.backward() if attack sign_flip: # 符号翻转把每个参数的梯度取反放大破坏效果 for param in model.parameters(): if param.grad is not None: param.grad -param.grad optimizer.step() return model.state_dict()local_train是攻击注入的核心位置。attacksign_flip时梯度取反会让本地模型朝着损失增大的方向更新聚合后全局模型精度会明显下降。注意这里是在backward()之后、step()之前改梯度顺序不能反。attackbackdoor时需要在dataloader里就把触发器加上这个在下一节讲。参数上lr设 0.01 比较稳epochs设 1 到 3太多轮恶意客户端会过度偏离反而容易被聚合算法识别。3.3 联邦训练主循环FedAvg 聚合与攻击客户端比例控制主循环负责下发模型、收集更新、聚合。源码里用了一个malicious_ratio参数控制恶意客户端比例实现方式是随机选一部分客户端在本地训练时传入attack参数。import copy def fed_avg(global_model, client_updates, weights): FedAvg 聚合按样本数加权平均各客户端参数 new_state copy.deepcopy(global_model.state_dict()) for key in new_state.keys(): new_state[key] torch.zeros_like(new_state[key], dtypetorch.float32) for update, w in zip(client_updates, weights): new_state[key] update[key].float() * w return new_state def federated_train(global_model, clients, rounds20, malicious_ratio0.2, attacksign_flip): num_malicious int(len(clients) * malicious_ratio) malicious_ids set(np.random.choice(len(clients), num_malicious, replaceFalse)) for r in range(rounds): updates, weights [], [] for cid, client in enumerate(clients): local_model copy.deepcopy(global_model) atk attack if cid in malicious_ids else None state local_train(local_model, client, attackatk) updates.append(state) weights.append(len(client.dataset)) # 权重归一化保证加权平均后总和为 1 total sum(weights) weights [w / total for w in weights] new_state fed_avg(global_model, updates, weights) global_model.load_state_dict(new_state) # 每轮评估全局模型精度用于画曲线 acc evaluate(global_model, test_loader) print(fRound {r1}, Test Acc: {acc:.4f}) return global_modelfed_avg里有个容易翻车的点torch.zeros_like默认是整型还是浮点取决于原张量如果原张量是整型比如某些统计量加权平均会丢精度。所以我在代码里显式加了.float()。malicious_ratio设 0.2 是常用起点你可以试 0.1、0.3 看攻击效果曲线。rounds设 20 轮足够看出趋势MNIST 上正常 FedAvg 大概 10 轮就到 98% 以上符号翻转攻击下会卡在 60% 到 80% 之间波动。注意恶意客户端的选择在每轮开始时固定不要每轮重新随机否则攻击效果会被稀释曲线会很难看。4. 攻击实现细节后门触发器构造与攻击效果评估4.1 后门触发器一个 3x3 白色方块就够了后门攻击的关键是触发器要足够小、足够独特模型能学会它但人眼不容易察觉。MNIST 上最常用的是右下角 3x3 白色方块CIFAR-10 上可以用 4x4 的彩色 patch。import torch def add_trigger(images, trigger_size3, positionbottom_right): 在图像上添加后门触发器返回修改后的图像 images: [N, C, H, W] 张量值域 [0, 1] triggered images.clone() _, _, H, W images.shape if position bottom_right: triggered[:, :, H-trigger_size:H, W-trigger_size:W] 1.0 return triggered def poison_dataset(dataset, poison_ratio0.1, target_label0): 把一部分样本加上触发器并改标签为目标类别 num_poison int(len(dataset) * poison_ratio) indices np.random.choice(len(dataset), num_poison, replaceFalse) for idx in indices: img, _ dataset[idx] img add_trigger(img.unsqueeze(0)).squeeze(0) dataset[idx] (img, target_label) return datasetadd_trigger把右下角像素置为 1.0poison_dataset随机选 10% 的样本加触发器并把标签改成目标类别比如 0。这样训练出来的本地模型会学到「右下角有白块 → 输出 0」的映射。参数上poison_ratio设 0.1 到 0.3太低学不会太高正常精度会掉。target_label一般选一个和触发器视觉上无关的类别选 0 或 1 都行。4.2 攻击效果评估精度曲线与后门成功率怎么算评估分两部分正常测试集上的准确率以及带触发器测试集上的后门成功率ASR。ASR 的定义是在带触发器的样本中被预测为目标类别的比例。def evaluate_backdoor(model, test_loader, target_label0): 计算后门成功率带触发器样本被预测为目标类别的比例 model.eval() correct, total 0, 0 with torch.no_grad(): for data, _ in test_loader: triggered add_trigger(data) output model(triggered) pred output.argmax(dim1) correct (pred target_label).sum().item() total data.size(0) return correct / total正常精度用标准的evaluate函数算后门成功率用上面这个。跑 20 轮每轮记录两个指标最后画两条曲线。正常情况下后门攻击下正常精度只掉 1 到 2 个百分点但 ASR 能从 10% 涨到 90% 以上这就是后门攻击的隐蔽性。符号翻转攻击则相反正常精度会掉到 60% 左右但 ASR 基本不变。攻击类型正常精度20轮后门成功率收敛轮数无攻击98.5%10%8符号翻转65.2%10%不收敛后门注入97.8%92.3%10梯度加噪88.1%10%15提示画图时把正常精度和后门成功率放在两个 y 轴或者分开两张图否则量纲差太多看不清趋势。4.3 参数敏感性恶意比例、攻击轮次与学习率的影响做实验最忌讳只跑一组参数就下结论。我一般会固定两个、变一个做三组对比。恶意比例malicious_ratio从 0.1 到 0.4符号翻转攻击下精度下降幅度从 10% 涨到 40%但超过 0.3 后聚合结果会剧烈震荡因为正常客户端已经压不住恶意更新了。后门攻击对比例没那么敏感0.1 就能达到 80% 以上的 ASR。学习率lr对攻击效果影响也大。恶意客户端用更大的学习率比如 0.05能放大符号翻转的破坏力但后门攻击反而要用小一点的学习率0.005否则触发器特征学不牢。攻击轮次上符号翻转从第 3 轮开始见效后门攻击要到第 8 轮左右 ASR 才明显起来。这些参数在源码里都有注释标出推荐范围你改的时候一次只动一个记录结果。5. 避坑与排查跑联邦对抗攻击时最容易翻车的五件事5.1 现象全局模型精度不降反升原因恶意客户端的梯度取反后如果聚合权重按样本数算而恶意客户端样本数很少取反的更新被正常更新淹没。或者攻击轮次太靠后全局模型已经收敛到很稳的局部最优单轮扰动拉不动。解决把malicious_ratio提到 0.2 以上或者让恶意客户端在每轮都参与不要随机丢。另外检查fed_avg里的权重是不是真的按样本数加权如果所有客户端权重一样恶意更新的影响会被平均掉。5.2 现象后门成功率上不去卡在 30% 左右原因触发器太小或者poison_ratio太低本地模型没学到触发器和目标标签的强关联。另一个常见原因是目标标签选得太偏比如选了一个正常样本里很少出现的类别模型倾向于忽略。解决把触发器从 3x3 加到 5x5poison_ratio提到 0.2。目标标签选正常分布里中等频率的类别比如 MNIST 上选 0 或 1。还可以在恶意客户端本地训练时多跑几轮让触发器特征更牢。5.3 现象训练中途 loss 变成 NaN原因符号翻转后梯度取反如果学习率没调小参数更新步长过大几轮之后权重爆炸。或者fed_avg里用了整型张量做累加溢出后变成 NaN。解决恶意客户端的学习率单独设小一点比如正常客户端 0.01、恶意客户端 0.005。fed_avg里所有累加前先.float()聚合完再转回原类型。如果还出 NaN加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)。5.4 现象不同轮次结果波动极大曲线像心电图原因恶意客户端每轮重新随机选导致攻击时有时无。或者 Dirichlet 切分时alpha太小某些客户端样本数为 0参与聚合时权重为 0实际参与客户端数量每轮都在变。解决恶意客户端集合在训练开始时固定整个训练过程不变。alpha不要低于 0.1切分后打印每个客户端的样本数发现有 0 就换种子或调大alpha。5.5 现象CPU 上跑一轮要十几分钟原因local_train里没有设batch_size默认一次喂整个数据集内存和计算都吃不消。或者evaluate每轮都在全量测试集上跑测试集大了就慢。解决DataLoader设batch_size64shuffleTrue。评估可以每两轮做一次或者只取测试集的一个子集。MNIST 上 10 个客户端、20 轮CPU 跑完大概 5 到 8 分钟超过这个数就是哪里没设对。6. 进阶技巧用检测指标反过来验证你的攻击是否真的有效跑通攻击只是第一步课程大作业想拿高分得证明你的攻击「有效且合理」。我一般会加一个检测环节用余弦相似度衡量恶意更新和正常更新的偏离程度如果偏离度很高但全局模型精度没怎么掉说明攻击隐蔽如果偏离度高且精度掉了说明攻击猛烈但容易被检测。def update_similarity(updates): 计算各客户端更新与平均更新的余弦相似度 flat [torch.cat([v.flatten() for v in u.values()]) for u in updates] mean_update torch.stack(flat).mean(dim0) sims [F.cosine_similarity(f.unsqueeze(0), mean_update.unsqueeze(0)).item() for f in flat] return sims这个函数返回每个客户端更新和平均更新的相似度。正常客户端的相似度一般在 0.9 以上符号翻转的恶意客户端会掉到 -0.5 以下后门注入的恶意客户端可能在 0.7 到 0.8 之间因为它只在触发器样本上偏离整体更新方向还是接近正常的。你可以把相似度画成柱状图恶意客户端一目了然。我自己的习惯是每做完一组攻击实验先看相似度分布再看精度曲线最后看 ASR。三个指标对不上说明实验里有变量没控制住比如恶意客户端集合变了、或者数据切分的种子没固定。这个习惯帮我省了很多次返工——有一次符号翻转精度只掉了 3 个百分点我以为是攻击太弱结果一查相似度发现恶意客户端的更新根本没被聚合进去原因是malicious_ids在循环里被重新赋值了。这种坑光看精度曲线是看不出来的。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网