单机模拟横向联邦学习:PyTorch实现FedAvg与Non-IID实战
2026/9/9 9:11:05 网站建设 项目流程

简介:面向Python开发者与机器学习初学者,提供一份横向联邦学习本地模拟的完整Python工程,演示了多客户端在各自私有数据上训练、仅上传模型参数供服务端聚合的协作流程。RAR压缩包共25个文件,以Python源码、CIFAR-10数据集、配置与项目文件为主,涵盖模型训练、数据加载、项目编译配置等多种用途,整体约302.68MB。已有723人学习下载,适合作为理解联邦学习原理的入门参考。资源内包含可直接运行的模拟代码,按服务端、客户端、模型、数据集等模块拆解,能够帮助读者快速搭建本地多客户端实验环境,深入理解联邦平均算法的实现细节,同时为后续研究隐私保护技术或通信优化提供基础。 横向联邦学习这几年在隐私计算和分布式机器学习里算是个高频词,尤其当“数据孤岛”和“数据合规”被反复强调之后,它的价值更明显了。这个方向的核心诉求讲起来很简单:多个参与方各自持有数据,不把原始数据汇总到一起,而是通过交换模型参数来协作训练一个全局模型。但真上手的时候,你会发现学习成本并不低——论文里的FedAvg、客户端选择、Non-IID数据分布这些概念,听起来都能懂,一到自己动手写代码就到处卡壳。

我最初学横向联邦学习时的困境很典型:没有真实的多机集群,也没有多方参与的数据环境,只能对着论文里的公式干瞪眼。后来我总结出一条实际走通的路径——在单机上用Python把整个横向联邦学习流程完整模拟一遍。这个方案不需要额外硬件,不需要搭分布式通信框架,只需要一台普通笔记本和PyTorch就能把FedAvg跑起来。这篇博文就把我实践中验证过的完整方案分享出来,包括数据划分、客户端训练、服务端聚合的每一步代码和踩坑记录,适合想入门联邦学习的研究者、算法工程师,以及需要快速验证联邦算法效果的学生。

1. 项目背景:为什么要在本地模拟横向联邦学习

1.1 横向联邦学习到底在解决什么问题

横向联邦学习(Horizontal Federated Learning)适用场景是:参与方的数据特征维度相同,但样本ID不同。举几个常见的例子会更直观——多家医院都有患者的年龄、性别、化验指标这些特征,但各自服务不同的患者群体;多家银行都有用户的基本信息和交易流水特征,但各自拥有的客户不同。传统做法是把所有数据集中到一个机房训练,但现实里这根本行不通,数据归属不同机构,直接汇总既涉及合规问题,又牵扯商业机密。

横向联邦学习的思路是“数据不动模型动”。每一方在自己的本地数据上训练模型,只把模型的参数更新(或者梯度)发给一个协调方,协调方把大家的结果聚合起来生成一个全局模型,再把新模型广播回去。原始数据从来没有离开过各自的服务器,但最终模型却能达到接近“数据集中训练”的效果。我最初听到这个思路时第一反应是“这不就是个分布式训练变种吗”,但真正动手实现后才发现,通信轮次、数据分布差异、参数聚合策略这些细节,每一个都值得重新推演。

我选择在本地模拟而不是直接搭集群,主要是三个阶段的不同需求。第一个阶段是理解机制,你需要能看到“服务器分发参数—客户端训练—回传参数—聚合广播”这个闭环每一步到底发生了什么;第二个阶段是验证算法改进,比如你改了聚合权重或者加了差分隐私噪声,需要快速对比效果;第三个阶段才是真正部署前做多机实验。本地模拟恰好覆盖前两个阶段,而且能让你在几十秒内完成一次完整的联邦训练闭环,这是任何分布式环境都给不了的调试效率。

1.2 本地模拟与传统集中式训练的本质差异

这里要说清楚一个问题:本地模拟联邦学习,并不是简单地写一个多客户端循环。集中式训练是“一份数据,一个模型,一次训练”,而联邦学习至少有三个关键差异。

第一,数据是分割的,而且是“不可见的”分割。每个客户端只能看到自己的数据子集,服务器从头到尾看不到任何样本。代码层面这意味着你要人为构造出多个数据子集,并且让各个子集之间“老死不相往来”。第二,模型参数是“旅行”的。全局模型要不断从一个地方拷贝到另一个地方:服务器生成、分发给各客户端,客户端本地更新后再传回服务器,服务器聚合后再分发。整个训练过程就是模型参数在“服务器—客户端”之间反复搬运的过程。第三,数据分布是“有偏”的。真实场景里各参与方的数据几乎不可能是独立同分布(IID)的,这恰恰是联邦学习最核心的挑战,也是本地模拟最值得认真处理的地方。

我把这三种差异用一张对比表列出来,方便大家对照理解:

维度集中式训练本地模拟横向联邦真实联邦集群
数据可见性中心可见全部数据客户端子集相互隔离客户端子集相互隔离
模型状态单份,原地更新参数在服务端与客户端间拷贝参数经网络传输
通信模拟函数调用/进程间传递网络请求
数据分布默认IID可构造IID或Non-IID自然Non-IID
调试难度中等

从这个表能看出,本地模拟的核心是把“通信”这个物理过程抽象成函数调用,把“数据隔离”用内存中的索引切分来实现,其他逻辑和真实联邦完全一致。这意味着你在本地验证过的算法逻辑,迁移到真实集群时不需要改核心代码,只需要把“参数的传递方式”换成真正的网络通信。

2. 核心机制拆解:FedAvg算法与数据分布设计

2.1 联邦平均(FedAvg)的完整计算流程

联邦学习里最经典、最基础也是工业落地最多的算法就是FedAvg,全称Federated Averaging,由McMahan等人在2017年提出。这个算法之所以经典,是因为它足够简单且有效。我用大白话描述它的一个完整round的运转流程。

先初始化一个全局模型参数,记作w₀。第t轮通信开始时,服务器从全部客户端中挑选一部分参与本轮训练(通常是随机抽,模拟真实场景中部分客户端离线或网络不佳的情况),把当前全局参数wₜ分发给被选中的客户端。每个客户端收到wₜ后,用这份参数作为自己本地模型的初始参数,在本地数据上做若干轮梯度下降,得到本地更新后的参数wₜ₊₁ₖ。随后各客户端把wₜ₊₁ₖ传回服务器。服务器收到所有参与客户端的本地模型后,按每个客户端本地样本量占总样本量的比例做加权平均,得到新的全局参数wₜ₊₁。

聚合公式可以写成:

wₜ₊₁ = Σₖ (nₖ / n) · wₜ₊₁ₖ

这里的nₖ是第k个客户端本地样本数,n是所有参与客户端本地样本总数。加权平均的逻辑很好理解:数据量大的参与方,本地模型训练得更充分,理应在全局模型中拥有更大的话语权。这个公式我在代码里直接用向量运算实现,也就是说对模型里的每一层参数张量都做一次加权求和。

在写代码时有一个非常隐蔽的坑——深拷贝问题。客户端从服务器接收全局参数时,必须用copy.deepcopy或者显式调用load_state_dict来创建独立的参数副本,否则客户端本地训练会直接修改全局模型那部分内存。这个坑我刚开始模拟时反复踩,训练一轮后所有客户端的模型参数一模一样,整个联邦学习退化成了一次普通训练,排查了半天才发现是内存共享导致的。

2.2 IID与Non-IID数据划分:模拟真实场景的关键

数据划分方式直接决定了你在模拟什么级别的联邦学习难度。如果所有客户端的数据都是从完整数据集里随机抽出来的子集,那它们的数据分布基本一致,属于IID场景,这对应的是最理想最简单的情况。但实际业务里,参与方的数据几乎必然是Non-IID的,也就是说不同客户端的数据分布差异很大。

最经典的Non-IID构造方式是“按标签切分”。以MNIST数据集为例,你有10个数字类别,如果客户端A只拿到标签为0、1、2的样本,客户端B只拿到标签为3、4、5的样本,那么每个客户端看到的类别分布就是极度偏斜的。这种情况下每个客户端本地训练出来的模型会严重偏向自己看到的类别,聚合时就要靠算法和超参数来“调和”各方矛盾。

我建议代码里同时实现IID和Non-IID两种划分方式,这样你可以通过对比实验亲眼看到Non-IID对联邦训练收敛速度的影响——这对于理解联邦学习的核心挑战非常有帮助。代码实现我刚才提过,IID直接随机打乱索引再均分,Non-IID则要把每个类别的样本先分成若干“分片”(shard),再把分片随机分配给不同客户端。

3. 环境准备与依赖安装

3.1 Python环境与依赖库清单

在本地模拟横向联邦学习,技术选型上我推荐用PyTorch。原因有三点:第一,PyTorch的state_dict机制非常方便存取模型参数,正好契合联邦学习“参数搬运”的核心操作;第二,它自带DataLoader数据加载器,对MNIST、CIFAR这类经典数据集的开箱支持很好,省去大量数据预处理代码;第三,PyTorch在学术和工业界普及率都很高,遇到问题随手一搜就能找到解决方案。

需要安装的依赖很少,核心就是三件套:

pip install torch torchvision pip install numpy

如果你用的是带NVIDIA GPU的机器,torch会自动匹配CUDA版本;如果没有GPU,用CPU跑MNIST这种小型数据集也完全够用,一轮来回也就几秒钟。torchvision用于加载MNIST数据集和做基本的图像变换。

我习惯在动手前先确认torch版本正常:

import torch print(torch.__version__) print(torch.cuda.is_available())

安装本身没有太多坑,唯一要提醒的是PyTorch版本之间API有时候会有细微差异,如果你按照我下面的代码跑报错,优先检查一下是不是版本问题。

3.2 代码逻辑模块划分

工程上我习惯把代码按角色拆分成几个模块,这样逻辑边界清晰,以后扩展也方便。不用太复杂的目录结构,一个单文件加几个函数其实也能跑通全部流程,但我建议至少按逻辑分块组织代码,良好的结构能让后续调试和扩展事半功倍。

完整的实现需要四个核心模块:数据加载与划分模块、模型定义模块、客户端逻辑模块、服务端逻辑模块。数据模块负责加载数据集,并按IID或Non-IID策略把数据索引划分给各个客户端;模型模块定义神经网络结构;客户端模块封装“用全局参数初始化本地模型,在本地数据上训练若干轮,返回更新后的参数”;服务端模块负责分发参数、聚合参数、评估全局模型。

我实际写的代码结构大致是:前两个部分放导入和超参配置,中间三个部分分别实现split_iidsplit_non_iidSimpleNetClientServer这几个核心类和函数,最后是主训练循环和结果输出。下面我按这个顺序逐一展开。

4. 核心代码实现:数据划分与模型定义

4.1 用PyTorch构建一个简单但完整的手写数字识别模型

模型不必复杂,能用来说明联邦学习机制就足够了。我用的是一个三层的全连接网络,输入是28×28=784维的像素向量,中间两个隐藏层各128和64个神经元,激活函数用ReLU,输出层是10个类别。完整定义如下:

import torch import torch.nn as nn class SimpleNet(nn.Module): def __init__(self): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(784, 128) self.fc2 = nn.Linear(128, 64) self.fc3 = nn.Linear(64, 10) def forward(self, x): x = torch.relu(self.fc1(x)) x = torch.relu(self.fc2(x)) return self.fc3(x)

这个网络放在2025年看显得很“古董”,但用来做联邦学习机制验证完全够用。MNIST分类本身不是一个有挑战性的任务,简单的MLP就能达到95%以上的准确率,这反而方便我们把注意力聚焦在联邦机制本身。如果你后续想验证更复杂的联邦算法,比如FedProx、SCAFFOLD,再换成CNN模型也不迟。

这里有一个关键知识点需要说明:PyTorch里所有神经网络层都通过nn.Module管理参数,model.state_dict()返回的是一个有序字典,里面保存了每一层可学习参数的张量。这个字典就是联邦学习里“搬运”的对象——服务器分发的、客户端回传的、服务端聚合的,本质都是这个字典。

4.2 数据预处理脚本

MNIST数据集在torchvision里可以直接下载,但第一次运行前需要设置好transform,把图像转为Tensor并且归一化。我习惯将输入像素范围从[0,255]缩放到[0,1]再减去均值除以标准差,这样可以加速收敛。

import torchvision import torchvision.transforms as transforms transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset = torchvision.datasets.MNIST( root='./data', train=True, download=True, transform=transform ) test_dataset = torchvision.datasets.MNIST( root='./data', train=False, download=True, transform=transform )

download=True只需要第一次运行下载数据,之后就会走本地缓存。如果网络条件差下载总是失败,可以手动从MNIST官网下载四个.gz压缩包放到./data/MNIST/raw/目录下,再重新运行代码。

MNIST原始训练集有60000张图片,测试集10000张。我后面的实验会用到训练集做客户端数据划分,用测试集做全局模型评估。

4.3 IID与Non-IID数据切分的代码实现

数据切分是整个模拟代码里最不应该出错的环节,因为后续所有训练逻辑都建立在“每个客户端拿到属于自己的数据”这个前提上。IID切分很好理解,把全部60000个样本的索引随机打乱,然后均分成若干份:

def split_iid(dataset, num_clients): data_len = len(dataset) indices = list(range(data_len)) random.shuffle(indices) split_size = data_len // num_clients client_indices = [ indices[i * split_size : (i + 1) * split_size] for i in range(num_clients) ] return client_indices

这段代码有个边界情况:如果num_clients不能整除60000,最后一个客户端会分到少一些的数据,这在实验里影响很小,但如果对样本数量敏感,可以把余数样本均匀分配到前几个客户端。

Non-IID切分要复杂一截。我通常用经典的shard方式:先把每个类别的样本索引单独拿出来,在每个类别内部随机打乱并切分成若干个shard,然后把这些shard随机分配给客户端,每个客户端拿固定数量的shard。比如总共10个类别,要分给10个客户端,每个客户端拿2个shard,那么总共需要20个shard,每个类别就要产出2个shard。

def split_non_iid(dataset, num_clients, num_shards_per_client=2): num_classes = 10 total_shards = num_clients * num_shards_per_client shards_per_class = total_shards // num_classes # 先按类别收集样本索引 class_indices = [[] for _ in range(num_classes)] for idx, (_, label) in enumerate(dataset): class_indices[label].append(idx) # 每个类别内随机打乱并切分成shards shard_indices = [] for cls in range(num_classes): random.shuffle(class_indices[cls]) per_shard_size = len(class_indices[cls]) // shards_per_class for i in range(shards_per_class): shard = class_indices[cls][i * per_shard_size : (i + 1) * per_shard_size] shard_indices.append(shard) # 打乱所有shards并分配给客户端 random.shuffle(shard_indices) client_indices = [] for i in range(num_clients): client_shards = shard_indices[i * num_shards_per_client : (i + 1) * num_shards_per_client] client_indices.append([idx for shard in client_shards for idx in shard]) return client_indices

运行这段代码你会直观看到Non-IID和IID的差别:IID场景每个客户端都能看到全部10个类别;Non-IID场景每个客户端只有2个类别。这意味着每个客户端本地训练时,模型只会被其中一两个类的梯度主导,聚合时全局模型面对的“观点”极其分裂,这也是Non-IID联邦训练收敛慢、准确率下降的根本原因。

5. 核心代码实现:联邦训练完整流程

5.1 Client客户端:本地训练实现

客户端的职责是“在本地数据上训练自己那份模型,并返回更新后的参数”。我在实现时用一个类来封装,构造函数传入模型结构、数据索引和全局数据集,在train方法里接收服务器下发的全局参数,完成本地训练后返回新参数和本地样本量。

import copy import torch.optim as optim class Client: def __init__(self, client_id, dataset, indices, device): self.client_id = client_id self.dataset = dataset self.indices = indices self.device = device self.loader = torch.utils.data.DataLoader( torch.utils.data.Subset(dataset, indices), batch_size=64, shuffle=True ) self.size = len(indices) def train(self, global_state, local_epochs=3, lr=0.01): model = SimpleNet().to(self.device) model.load_state_dict(global_state) optimizer = optim.SGD(model.parameters(), lr=lr, momentum=0.9) criterion = nn.CrossEntropyLoss() model.train() for epoch in range(local_epochs): for images, labels in self.loader: images = images.view(-1, 784).to(self.device) labels = labels.to(self.device) optimizer.zero_grad() outputs = model(images) loss = criterion(outputs, labels) loss.backward() optimizer.step() return model.state_dict(), self.size

这里面有几个值得注意的实现细节。每次调用train都必须基于传入的global_state重新初始化一个SimpleNet实例,保证客户端没有携带上一轮的“记忆”,这正是联邦学习的核心约束——每一轮本地更新都从全局模型出发。load_state_dict要求传入的字典键名与模型结构完全匹配,所以服务器分发出来的参数必须原封不动地传给客户端,中间不要做任何格式转换。

本地epoch数我默认设为3,学习率0.01,这是实践下来比较稳的一组参数。本地epoch太大会加剧本地的模型漂移,太小则模型学不到东西,这个平衡点后面实验结果会体现出来。

5.2 Server服务端:FedAvg加权聚合实现

服务端的职责是维护全局模型,分发参数、接收回传、完成聚合。聚合公式前面分析过,按每个客户端本地样本量的比例对模型参数做加权平均,代码实现起来非常直观:

class Server: def __init__(self, device): self.device = device self.global_model = SimpleNet().to(self.device) self.global_state = copy.deepcopy(self.global_model.state_dict()) def aggregate(self, client_states, client_sizes): total_size = sum(client_sizes) weights = [size / total_size for size in client_sizes] new_state = {} for key in self.global_state: new_state[key] = torch.zeros_like(self.global_state[key]) for state, weight in zip(client_states, weights): for key in state: new_state[key] += state[key] * weight self.global_model.load_state_dict(new_state) self.global_state = copy.deepcopy(new_state) return self.global_state

这段代码里我做了两次deepcopy。第一次在初始化时,保存全局模型参数的独立副本——如果不深拷贝,后续所有操作都直接作用于global_model内部的参数,就没法区分“当前全局状态”和“被修改过程中的状态”。第二次在聚合完成后,将new_state深拷贝给self.global_state,这样即使外部逻辑后续修改了new_state的某些键,也不会污染服务端维护的全局状态。这种防御性编程看起来啰嗦,但在联邦模拟里真的能帮你少排查很多诡异bug。

权重是每个客户端的样本数除以总样本数,这个比例计算不需要额外处理数据均匀性的问题,FedAvg本身已经考虑了数据量对聚合贡献的权重。如果某个客户端数据量特别大,它在聚合时自动获得更大的话语权,这是符合直觉的——数据多,本地模型更可靠。

5.3 主训练循环:R轮通信+客户端选择

主循环是整个模拟的骨架。每一轮通信要做的事情是:从全部客户端里随机选取一部分参与本轮训练,把当前全局参数分发给被选中的客户端集合,让它们各自在本地训练,收集训练后的状态,调用Server聚合,然后在测试集上评估全局模型准确率。

def evaluate(model, dataloader, device): model.eval() correct, total = 0, 0 with torch.no_grad(): for images, labels in dataloader: images = images.view(-1, 784).to(device) labels = labels.to(device) outputs = model(images) preds = outputs.argmax(dim=1) correct += (preds == labels).sum().item() total += labels.size(0) return correct / total def run_fedavg(client_list, server, test_loader, num_rounds=20, device='cpu', sample_ratio=1.0): num_selected = max(1, int(len(client_list) * sample_ratio)) for round_idx in range(num_rounds): selected_clients = random.sample(client_list, num_selected) client_states = [] client_sizes = [] for client in selected_clients: state, size = client.train(server.global_state) client_states.append(state) client_sizes.append(size) server.aggregate(client_states, client_sizes) # 每轮结束做一次全局模型评估 acc = evaluate(server.global_model, test_loader, device) print(f'Round {round_idx + 1}/{num_rounds}, Test Acc: {acc:.4f}', flush=True) return server.global_model

这里我特意用random.sample做客户端选择,是为了模拟真实联邦环境里客户端不可能全部稳定在线的情况。如果sample_ratio=1.0就是全量参与,这是最理想的模拟设定;调到0.5或者0.3就能明显看到模型收敛速度变慢、准确率波动的现象,很适合观察“部分客户端离线”对联邦训练的影响。

评估时必须切换成model.eval()模式并包裹在torch.no_grad()里,否则Dropout和BatchNorm这些层的训练/推理行为不一致会污染评估结果,虽然我们这个简单MLP没有Dropout和BN,但这个好习惯值得养成。

6. 实验运行与结果分析

6.1 实验配置与核心参数

我在本地笔记本上跑了三次实验,对照集中式训练、联邦IID和联邦Non-IID三种场景。核心参数如下表格:

参数取值
数据集MNIST
客户端总数10
每轮参与率100%
通信轮数20
本地epoch数3
批量大小64
优化器SGD momentum=0.9
学习率0.01
本地数据划分IID / Non-IID(每客户端2个类别)

集中式训练就是按常规方式把全部60000样本喂给同样的SimpleNet训练15个epoch,用来作为准确率上限的参考。

训练结束后记录测试集准确率的变化曲线,对比这几种场景的收敛速度和最终精度。

6.2 三种训练模式的对比分析

在我的环境下跑出来的结果符合预期:集中式训练在15个epoch后达到约97%的测试准确率;联邦IID场景在20轮通信后达到约96%左右,和集中式差距很小;联邦Non-IID场景只有约88%~91%,明显偏低,且收敛过程中准确率波动更大。

这个对比充分说明两个问题。第一,IID联邦训练的效果接近集中式,“数据不动模型动”这个思路是可行的;第二,Non-IID数据分布是联邦学习准确率损失的主要来源——每个客户端只在极少数类别上学过,聚合出的全局模型在面对所有类别时表现必然打折。这让我当时特别直观地体会到了Non-IID这个核心难题的分量,也理解了为什么后续会有FedProx、MOON这类针对Non-IID设计的联邦算法。

如果你实验时发现IID联邦结果和集中式差距过大,我建议优先检查客户端本地epoch和学习率是否设置合理。它们分别控制客户端在每轮通信之间的“训练强度”,如果设置得太弱,每轮模型变化太小,会拖慢全局收敛;设置太强,又会加剧本地的参数偏移。这个平衡需要通过实验来调优。

7. 常见问题与排查技巧实录

7.1 本地模拟典型问题速查表

整个模拟过程中我遇到过的坑不少,下面按出现频率排个序,整理成一张速查表:

问题现象可能原因排查方法
所有客户端最终模型完全一样客户端直接持有了全局模型引用,没有深拷贝检查是否用deepcopy或显式load_state_dict
测试准确率极低或直接崩溃数据切分出现类别缺失/数据集划分错误打印每个客户端索引对应的标签分布
Non-IID下模型不收敛本地epoch过多或学习率过大调低学习率到0.005以下,减少本地epoch重跑
聚合后模型参数包含NaN优化器发散导致参数溢出检查学习率、检查数据是否包含异常值
每轮评估准确率剧烈震荡客户端选择随机性太强增加参与率、增大客户端本地样本量

这里面最隐蔽的是第一个问题。通常模拟里客户端类会接收一个全局模型的state_dict,如果在构造函数里直接把这个state_dict赋值给本地模型,两个客户端训练时就会共享底层的内存张量,导致A客户端改了参数,B客户端的模型也跟着变。这个bug非常难查,因为它不会报错,只是结果对不上。

7.2 本地模拟的几个重要边界认知

写这篇文章的时候,我也想强调一个认知边界:本地模拟和真实分布式联邦学习之间,毕竟还有一道跨越鸿沟。本地模拟把通信过程抽象成了函数调用,这就意味着你完全忽略了网络延迟、带宽限制、客户端掉线重连、参数传输的序列化与反序列化开销这些真实分布式系统中的核心代价。

如果你的目标只是理解算法机制、验证改进思路,本地模拟完全是够用的。但如果你要做真正面向生产环境的联邦系统设计,那还需要把思路切换到真实的分布式通信框架上去,比如用Ray、PySyft这种支持多进程或多机的框架,把“参数搬运”改成真实网络传输,把“进程内函数调用”改成消息传递。我个人的建议是:先在本地把算法逻辑跑通跑稳,再迁移到分布式框架,这个顺序能帮你把算法层面的问题和工程层面的问题分离,定位效率翻倍。

8. 实验总结与经验体会

写到这里,整个本地模拟横向联邦学习的实现已经完整呈现了。复盘来看,这个项目最大的价值在于用最小成本打通了联邦学习的核心闭环——从数据切分到客户端训练、从参数聚合到模型评估,每一步都可以在单机上直观观察、反复调试。对我来说,亲手实现一遍FedAvg之后,再回去看那些理论论文就顺畅了很多,因为代码里每一个变量、每一次传参都对应着论文里的一个概念步骤。

最后再分享一个小技巧:如果你想在这个基础上做扩展,可以先从“把MNIST换成CIFAR-10”或者“把SimpleNet换成CNN”开始。这看起来改动简单,但当你把模型从MLP换成CNN,再在Non-IID数据分布下跑一轮,你会更加深刻地体会到联邦学习在真实场景里遇到的计算偏差和收敛困难。也可以继续加模拟其他人算法,比如在客户端本地加差分隐私噪声来观察对模型精度的影响,这是目前隐私计算和联邦学习结合得最紧密的方向之一。本地模拟的边界就在这里——它是你通往更复杂联邦系统的起点,而不是终点。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询