news 2026/9/24 19:01:08

联邦学习实战:从FedAvg到FedProx的递进实验指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
联邦学习实战:从FedAvg到FedProx的递进实验指南

简介:一套基于Python实现的联邦学习实验项目,包含三个递进式实验,适合人工智能、计科、通信工程等专业的毕设、课程设计或入门进阶。实验围绕Cifar-10、MedMNIST和Chest X-Ray Images三个数据集展开,对比FedAvg、FedPer、FedRep与自研FedOur算法的准确率和目标损失值,并考察10/50/100不同客户端规模及全局/本地模型效果,能帮助读者系统理解联邦学习的核心流程与常用评估口径。资源共43个文件,以14个Python源码脚本和20张图片(png/jpg)为主体,另含少量xml配置、模型说明和需求文档,压缩包仅631KB,目录结构直观,便于按实验模块查阅与复现。当前已有225人学习下载。项目代码经作者充分测试,答辩评审平均分96分,并配有模型与图片演示,适合直接作为项目演示或在此基础上扩展更多联邦学习思路。

1. 联邦学习实验到底在做什么:一个能跑通、能出图、能写进报告的项目

如果你在搜索引擎里敲下“联邦学习”四个字,八成是被课程要求或者论文选题逼到这一步的。联邦学习不是一种新算法,而是一套分布式机器学习框架,它的核心诉求是「数据不动模型动」——各个客户端保留本地数据,只上传模型参数,由服务器聚合后回传。这个思路在数据隐私法规越来越严的今天几乎是唯一可行的跨机构训练方案,但很多人第一次接触它,不是被原理劝退,而是被「不知道怎么跑起来」卡住。这个基于Python实现的联邦学习实验项目,价值恰恰在于它把三个递进式实验连同源代码、训练好的模型和图表一次性给你,从单机模拟FedAvg开始,到非独立同分布数据切分,再到不同聚合策略的对比,跑完你就知道联邦学习的数据流是什么样子、参数聚合到底怎么做的、Non-IID为什么让效果变差。它适合三类人:被实验报告逼疯的学生、想快速上手联邦学习但不想从零搭框架的开发者、以及在传统机器学习里待久了想看看数据隔离场景下怎么训练的工程师。

2. 三个实验的递进逻辑与运行环境:为什么第一个实验选MNIST,第二个实验必须做Non-IID

2.1 实验架构的整体设计:从一个模型、多个客户端到参数服务器的闭环

这个项目不是三个无关的实验堆在一起,而是一条递进链。第一个实验是基础版联邦平均算法,让你理解参数服务器的通信流程;第二个实验把标准数据集切分成Non-IID分布,逼近真实场景;第三个实验在Non-IID基础上引入不同的聚合策略或训练方式,形成对比。用我做过类似项目的经验说,这种「基线 → 问题引入 → 解法对比」的结构,恰恰是实验报告拿高分最稳妥的框架。

整个系统的通信拓扑是经典的中心化联邦学习结构——一个服务器节点 + N个客户端节点。注意,这里的「节点」不是物理意义上的多台机器,而是进程级模拟。常见做法是写一个客户端类,在循环里依次实例化,模拟参与本轮训练的客户端子集。这样做的好处是单机就能跑通全流程,不需要搭分布式环境,对课程实验来说完全够用。

# 联邦学习整体结构示意(伪代码级描述) class FedServer: def __init__(self): self.global_model = init_model() # 初始化全局模型 self.client_list = [] # 客户端实例列表 def train_round(self, round_idx): sampled = random.sample(self.client_list, k=num_clients_per_round) local_weights = [] for client in sampled: w = client.local_train(self.global_model.state_dict()) local_weights.append(w) new_weights = self.aggregate(local_weights) # 加权平均 self.global_model.load_state_dict(new_weights)

逻辑说明:每一轮先随机采样一定比例的客户端,各客户端拿到全局模型参数后用自己的本地数据训练若干个epoch,然后把更新后的参数返回服务器,服务器做加权平均得到新一轮全局模型。state_dict()是PyTorch里模型参数的载体,聚合操作本质上就是对这些张量做加权平均。

参数说明:k = num_clients_per_round是每轮参与训练的客户端数量。在模拟场景中总客户端数可能在10到100之间,每轮采样比例常见取0.2到0.5。比例太高通信成本大,比例太低每轮信息量不足,模型收敛变慢——这个trade-off在第三个实验的对比里会体现得很明显。

2.2 环境搭建与依赖清单:Python版本、PyTorch、torchvision和matplotlib

先把环境的问题说透。这个项目基于Python实现,建议直接使用Python 3.8到3.10之间的版本,不是越新越好——某些旧版依赖库对新版本Python的支持会有兼容性问题,这会浪费大量时间。PyTorch选用CPU版还是GPU版取决于你的电脑有没有NVIDIA显卡,显存低于4G的话建议直接用CPU版先把流程跑通。

# 创建虚拟环境并安装依赖(Windows / Linux 通用) python -m venv fl_env # Linux: source fl_env/bin/activate Windows: fl_env\Scripts\activate pip install torch torchvision matplotlib numpy scikit-learn

逻辑说明:venv创建隔离环境,避免和系统里其他Python项目打架。依赖库一共五个方向——torch提供张量运算和神经网络模块,torchvision用来下载和加载MNIST/Fashion-MNIST数据集,matplotlib负责画损失曲线和准确率曲线,numpy做数据切分和数组操作,scikit-learn主要用于t-SNE可视化(第三个实验可能会用到)。

参数说明:如果安装缓慢,可以加上-i https://mirrors.aliyun.com/pypi/simple/这类镜像参数加速。如果你安装的是CPU版PyTorch,务必注意别装成CUDA版——启动时不会报错,但训练速度会慢到让你怀疑人生。判断方法是运行python -c "import torch; print(torch.cuda.is_available())",输出False说明在跑CPU。

2.3 数据集准备:MNIST还是Fashion-MNIST,以及如何划分训练集测试集

这个项目用的数据集在标题里没写死,但从「模型+图片演示」这个组合推断,最常见的配套数据集是MNIST手写数字。MNIST是联邦学习实验的事实标准,原因很简单:28x28的灰度图,60000张训练图,类别均衡,模型不需要太复杂就能达到不错的效果,跑一轮的时间可控。Fashion-MNIST是它的Drop-in替代品,难度略高,如果想展示「联邦学习在复杂任务上的效果」可以换用。

from torchvision import datasets, transforms # 下载 MNIST 并转换为 Tensor transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset = datasets.MNIST('./data', train=True, download=True, transform=transform) test_dataset = datasets.MNIST('./data', train=False, download=True, transform=transform)

逻辑说明:ToTensor()把PIL图像转换为0到1之间的张量,Normalize用MNIST的全局均值和标准差做标准化。这两个值不是随便定的,是官方给的标准值。测试集不参与训练,只用来评估每一轮聚合后的全局模型准确率。

参数说明:如果你后续要做Non-IID切分,测试集保持IID不变,只对训练集做划分。这是实验设计上的关键点——测试集IID保证评测公平,训练集Non-IID模拟真实数据孤岛。数据下载失败的话,检查网络,或者手动下载到./data/目录下对应文件夹中。

3. 三个实验的代码实现与运行逻辑:从FedAvg到Non-IID再到聚合策略对比

3.1 实验一:FedAvg算法的最小实现,两个客户端就够跑通

实验一的目标不是刷精度,而是让你看到「模型参数真的在被聚合」。所以我建议先从两个客户端开始,各持有一半数据,跑20轮通信,把准确率曲线画出来。如果一上来就模拟100个客户端,反而看不清核心逻辑。

import copy import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Subset # 定义一个简单的CNN模型(两个卷积层 + 两个全连接层) class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 16, 5, stride=1, padding=2) self.conv2 = nn.Conv2d(16, 32, 5, stride=1, padding=2) self.fc1 = nn.Linear(32*7*7, 128) self.fc2 = nn.Linear(128, 10) def forward(self, x): x = torch.relu(self.conv1(x)) x = torch.max_pool2d(x, 2) x = torch.relu(self.conv2(x)) x = torch.max_pool2d(x, 2) x = x.view(x.size(0), -1) x = torch.relu(self.fc1(x)) return self.fc2(x) def fed_avg(global_model, client_weights): """对客户端权重做加权平均""" avg_weights = copy.deepcopy(global_model.state_dict()) for key in avg_weights.keys(): avg_weights[key] = torch.mean( torch.stack([w[key].float() for w in client_weights]), dim=0) return avg_weights # 将数据平均分给2个客户端 split_size = len(train_dataset) // 2 client1_data = Subset(train_dataset, range(0, split_size)) client2_data = Subset(train_dataset, range(split_size, split_size * 2)) clients_data = [client1_data, client2_data] global_model = SimpleCNN() for round_idx in range(20): client_weights = [] for client_data in clients_data: local_model = copy.deepcopy(global_model) loader = DataLoader(client_data, batch_size=32, shuffle=True) optimizer = optim.SGD(local_model.parameters(), lr=0.01, momentum=0.9) local_model.train() for epoch in range(3): # 本地训练3个epoch for batch_x, batch_y in loader: optimizer.zero_grad() out = local_model(batch_x) loss = nn.functional.cross_entropy(out, batch_y) loss.backward() optimizer.step() client_weights.append(local_model.state_dict()) global_model.load_state_dict(fed_avg(global_model, client_weights))

逻辑说明:每个客户端拿到全局模型的深拷贝,在本地数据上训练3个epoch后返回权重。fed_avg函数对每个参数张量在客户端维度上求均值。注意这里用的是torch.mean而不是简单的算术平均——如果客户端数据量不同,你应该用数据量加权,这就是后续实验要改的聚合策略。

参数说明:lr=0.01momentum=0.9是CNN配SGD的经典组合。本地epoch设3是因为联邦学习里每一轮通信的成本远高于本地计算,本地训练太多会导致「模型漂移」,后面避坑章会专门讲。batch_size设32对MNIST来说是稳妥选择,显存不够可以降到16。

3.2 实验二:Non-IID数据切分,用Dirichlet分布模拟真实数据分布

第二个实验是很多人第一次接触联邦学习会忽略、但真实场景必须面对的问题:不同客户端的数据分布差异极大。一台手机上的照片几乎没有猫,另一台全是猫的照片,这种情况下直接用FedAvg,精度会掉得惨不忍睹。这个实验的核心就是把IID数据切成Non-IID,观察精度下降幅度。

import numpy as np from torch.utils.data import ConcatDataset def dirichlet_split(dataset, num_clients, alpha=0.5): """用 Dirichlet 分布将数据集划分为 Non-IID 的 N 份""" labels = np.array([dataset[i][1] for i in range(len(dataset))]) num_classes = np.max(labels) + 1 client_indices = [[] for _ in range(num_clients)] for c in range(num_classes): idx_c = np.where(labels == c)[0] np.random.shuffle(idx_c) # 从 Dirichlet 分布采样每个客户端分到的比例 proportions = np.random.dirichlet([alpha] * num_clients) proportions = (proportions * len(idx_c)).astype(int) proportions[-1] = len(idx_c) - np.sum(proportions[:-1]) start = 0 for i, size in enumerate(proportions): client_indices[i].extend(idx_c[start:start + size]) start += size return [Subset(dataset, idx) for idx in client_indices] clients_data_non_iid = dirichlet_split(train_dataset, num_clients=10, alpha=0.5)

逻辑说明:alpha是Dirichlet分布的浓度参数,控制数据分布的不均衡程度。alpha越小,分布越不均匀——某个类别的数据可能集中到少数几个客户端上;alpha趋近无穷大时退化为IID。按类别循环切分的逻辑是:先取出所有标签为c的样本索引,用Dirichlet采样比例切割,再分配到各客户端。最后一行代码把proportions最后一个元素补足到总数,避免浮点误差导致的元素丢失。

参数说明:alpha=0.5是实验中常用的中间值,能明显看到精度下降但又不至于完全无法收敛。如果你想展示极端Non-IID,用alpha=0.1——这时候某些客户端可能完全没有某个数字的训练样本,模型的类间偏差会暴露得很彻底。

3.3 实验三:聚合策略对比,看FedAvg、FedProx和加权的差别

第三个实验是拉开分数差距的地方。FedAvg在Non-IID下效果差的原因是本地优化方向偏离了全局最优,FedProx针对这个问题在本地损失函数中加了近端项,约束本地模型不要偏离全局模型太远。这个实验不需要你重新写一套框架,改动只在客户端训练函数和服务器聚合函数上。

def local_train_prox(global_model, local_data, mu=0.01, local_epochs=3): """FedProx 客户端训练:加了近端项的本地训练""" local_model = copy.deepcopy(global_model) optimizer = optim.SGD(local_model.parameters(), lr=0.01) loader = DataLoader(local_data, batch_size=32, shuffle=True) global_params = {k: v.clone() for k, v in global_model.state_dict().items()} for epoch in range(local_epochs): for batch_x, batch_y in loader: optimizer.zero_grad() out = local_model(batch_x) loss = nn.functional.cross_entropy(out, batch_y) # 近端项:计算与全局参数的 L2 距离 prox_term = 0.0 for k, v in local_model.state_dict().items(): prox_term += torch.sum((v - global_params[k]) ** 2) loss = loss + (mu / 2) * prox_term loss.backward() optimizer.step() return local_model.state_dict() def weighted_aggregate(client_weights, client_sizes): """按数据量加权的聚合函数""" total_size = sum(client_sizes) avg_weights = copy.deepcopy(client_weights[0]) for key in avg_weights.keys(): avg_weights[key] = sum( w[key] * size for w, size in zip(client_weights, client_sizes) ) / total_size return avg_weights

逻辑说明:FedProx的核心区别在loss的计算——原本只有交叉熵损失,现在加了mu/2 * ||w - w_global||^2这一项。这个约束项让本地模型在拟合本地数据的同时不敢跑太远,相当于给本地训练加了一根「风筝线」。加权聚合函数根据客户端数据量按比例分配权重,客户端数据越多,在全局模型中说话的分量越重。

参数说明:mu=0.01是FedProx作者论文中效果较好的默认值。mu太大,本地模型几乎不更新,收敛慢;mu太小,约束失效,等价于FedAvg。对比实验的关键是控制变量——跑三组:FedAvg上Non-IID、FedProx上Non-IID、FedAvg上IID(作为性能上限参照),画在同一张图上,报告里这张图就是最有力的论据。

3.4 图片演示:把损失曲线、准确率曲线和类别混淆矩阵画出来

有了实验结果,怎么把「图片演示」做成实验报告里的亮点?两个技巧:一是把多个实验的曲线画在同一张图里对比,二是用混淆矩阵直观展示模型在哪些类别上容易搞混。

def plot_compare(rounds_list, accs_dict, save_path="compare.png"): """把多个实验的准确率曲线画在一起对比""" plt.figure(figsize=(10, 6)) for name, accs in accs_dict.items(): plt.plot(range(1, rounds_list+1), accs, marker='o', markersize=4, label=name) plt.xlabel("Communication Rounds") plt.ylabel("Test Accuracy (%)") plt.title("FedAvg vs FedProx on Non-IID MNIST") plt.legend() plt.grid(True, linestyle='--', alpha=0.3) plt.tight_layout() plt.savefig(save_path, dpi=200) plt.close()

逻辑说明:把不同实验的准确率历史记录存成字典,统一画图。savefig之前一定要plt.close(),否则在循环或多次调用画图函数时会内存累积,特别是Jupyter里会越来越卡。dpi=200保证插入报告后文字依旧清晰。

参数说明:画混淆矩阵用sklearn.metrics.confusion_matrix配合imshow即可。注意联邦学习场景下混淆矩阵要基于聚合后的全局模型做推理,不是随便挑一个客户端模型——这才能反映系统整体效果。

4. 参数调优与避坑:五个让实验翻车的细节,每个都踩过

4.1 随机种子没固定,实验不可复现

现象:同一套代码连续跑两次,准确率曲线不一样,甚至相差几个百分点。

原因:PyTorch初始化权重、数据加载打乱顺序、NumPy切分数据都有随机性。没固定随机种子,别人的实验结果你复现不了,你自己也复现不了自己。

解决:在代码最前面统一设置随机种子,固定所有能固定的随机源。

import random import numpy as np import torch def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False set_seed(42)

这是血泪教训:不设种子,后面做聚合策略对比实验时,你分不清精度差异是策略带来的还是随机波动带来的。设了种子,FedAvgFedProx跑在同一数据切分下,对比才有说服力。

4.2 本地epoch设得太大,聚合后精度不升反降

现象:本地训练设10个epoch,比设3个epoch的全局收敛精度低,而且损失曲线震荡剧烈。

原因:这就是联邦学习里常说的「客户端漂移」。本地训练太充分,模型被拉到只适配本地数据分布的方向。服务器做平均时,不同客户端的方向差异太大,平均结果反而远离全局最优点。

解决:三个实验统一用local_epochs=3左右,最多别超过5。如果发现聚合后精度震荡,优先降低本地epoch而不是调整学习率。这个参数是联邦学习和普通深度学习在调参上最大的区别——你不能用单机训练的思路想当然。

4.3 Non-IID切分时索引丢失,客户端数据量不均衡导致加权聚合失效

现象:用Dirichlet切分后,某个客户端分到的数据为0,或者总量和原始数据集对不上。

原因:浮点运算误差导致proportions数组加起来不等于len(idx_c),最后一段数据没分给任何客户端。直接把astype(int)的结果用来切分,精度损失可能造成样本被漏掉。

解决:用4.2节代码里的补足方案——计算完proportions后,最后一项用len(idx_c) - sum(前面的)来修正。另外建议在切分后打印每个客户端各类别样本数,用代码验证切分正确性再开跑。

4.4 聚合权重用了torch.mean但客户端数据量差异极大

现象:实验一里两个客户端数据量一致,结果正常;到了实验二Non-IID切分后,准确率暴跌,比IID实验掉了十几个点。

原因:数据量少的客户端和多的客户端平权参与聚合,小客户端本地的噪声被放大——它的梯度方向可能被少量异常样本主导,却和大量数据客户端有同样的话语权。

解决:在第三个实验里已经给出了加权聚合的实现。服务器聚合前拿到各客户端本地数据量,按比例加权。注意要先计算total_size再循环,别在循环里频繁除法。

4.5 图片演示中保存了图但报告里看不清文字

现象:把savefig出来的图片直接插进报告,图例文字糊成一团,坐标轴数字重叠。

原因:默认的画布尺寸偏小,图例过多时互相挤压;dpi太低,放大后模糊。

解决:训练模拟时每隔5轮记录一次准确率,画图时坐标点少而清晰。figsize=(12, 6)dpi=200以上,图例位置用plt.legend(loc='lower right')手动指定——默认的best位置在曲线密集时会乱跳。如果要进一步对比不同round数下的效果,把多个实验画在同一坐标里时,调色板用plt.cm.tab10,颜色区分度最高。

5. 把实验做深:特征可视化与模型鲁棒性评估的进阶手法

如果三个实验已经跑通,报告的基本盘已经稳了。但想拿高分或者真正理解联邦学习的特性,建议再做两个延伸:一是用t-SNE把全局模型在测试集上的特征降维可视化,观察不同类别的特征分布是否清晰可分;二是评估模型在Non-IID分布下的类别偏差——挑出那些在某个客户端占比极低的类别,单独算它的类别准确率。

t-SNE的落地做法是:取全局模型的倒数第二层输出(即fc1的输出,128维),喂给sklearn.manifold.TSNE降到2维,按数字类别着色画散点图。这一步能直观展示联邦聚合后模型是否学到了一个类间可分、类内紧凑的特征空间——这是普通准确率曲线给不了的直观证据。

类别偏差评估则更有意思:统计每个客户端各类别的样本数,找出「全局占比高但客户端占比极低」的类别,单独计算并打印它们的精确率和召回率。你会发现FedAvg在这些类别上尤其脆弱,而FedProx在mu=0.01时有一定缓解。这个观察点可以直接写进报告的价值分析部分。

我当时做这个项目时,在Non-IID切分上卡了两天。第一次在FedProx上跑了一个惊为天人的结果,一度以为是调参玄学,后来发现是对比实验的随机种子没固定,两个实验压根不在同一数据分布上对比。教训很简单:联邦学习实验里,数据切分和训练过程有太多层次的随机性,每加一层都必须固定一个种子,否则你看到的精度差异全部是统计噪声。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/24 19:00:59

Drift Loss生成模型MNIST复现:从原理到代码的完整实践

最近在折腾生成模型,看到Generative Modeling via Drifting这套框架,训练目标简洁到只有一个Drift Loss,就很想拿MNIST完整复现一遍。这套方法的核心思想非常直接:把生成过程看作粒子在数据空间里做漂移,网络只需要学会…

作者头像 李华
网站建设 2026/9/24 19:00:58

WPF 嵌入 HTML 页面实战:WebBrowser 控件从内核配置到性能优化

先说下这次项目的背景:要在 WPF 主界面里嵌两个 HTML 页面,一个是数据看板,一个是报表展示页,开发周期非常紧,团队里也没有专门的前端配合,最省事的方案就是直接用 WPF 自带的 WebBrowser 控件。结果真用起…

作者头像 李华
网站建设 2026/9/24 19:00:09

Flutter鸿蒙化适配实战:simple_json库迁移全流程复盘

一个做了三四年 Flutter 的老手,第一次把项目往鸿蒙(HarmonyOS)侧迁移时,最先崩溃的往往不是页面,而是各种三方库。UI 层还好,最麻烦的是底层依赖,尤其是 JSON 序列化这种全局都得用的基础设施。…

作者头像 李华
网站建设 2026/9/24 18:59:49

AI编程工具选型指南:Cursor、Claude Code、Codex、Copilot深度对比

AI 编程工具这两年更新得快,快到什么程度?我上个月刚把某个工具的快捷键肌肉记忆练熟,这个月它就改了交互逻辑。Cursor、Claude Code、Codex、GitHub Copilot 这几个名字,几乎每隔几天就会在技术群里被拉出来对比一轮。但说实话&a…

作者头像 李华
网站建设 2026/9/24 18:59:44

水果图像分类数据集8分类实战:从数据预处理到模型调优的完整指南

简介:这份资源是面向深度学习入门与图像分类实践者的水果图像分类数据集,覆盖苹果、香蕉、樱桃、火龙果、芒果、橘子、菠萝、木瓜共8个类别,可直接用于模型训练与验证,省去自行采集与清洗图像的环节。压缩包内共约2000个文件&…

作者头像 李华
网站建设 2026/9/24 18:59:24

Flask+深度学习中文情感分析系统实战:从模型推理到Web部署

简介:本资源为基于Python与深度学习的中文情感分析系统毕业设计完整资料包,面向计算机相关专业需要完成毕业设计的学生及希望学习Flask Web开发与文本分类的开发者。系统采用Flask框架搭配MySQL数据库,实现用户注册登录、后台数据统计首页以及…

作者头像 李华