上一篇读FedAvg论文时,我主要在理解它为什么要用本地计算换通信,以及一轮训练里服务器和客户端分别做什么。读到最后,流程大致能说出来了,但还有一个很实际的问题没有解决:论文里的Algorithm 1,如果真的准备写成PyTorch代码,每一步到底对应什么?

这篇先不跑实验,也不记录任何准确率。我现在只是把论文伪代码和一个最小程序应该具备的结构对应起来,看看自己之后要写哪些部分。真正的MNIST FedAvg复现还没有完成。

从普通训练循环往外再套一层

前面写MNIST的MLP和CNN时,训练过程大概是:

1
DataLoader -> forward -> loss -> backward -> optimizer.step()

FedAvg并没有把这套过程推翻。一个客户端在本地训练时,仍然要读取一批数据、前向传播、计算loss、反向传播,再更新模型参数。

变化主要发生在这段训练循环的外面:训练集要先拆给不同客户端;每轮只选择其中一部分客户端;这些客户端都从同一个全局模型开始;本地训练结束后,服务器还要收集并聚合它们的模型。

所以我现在把FedAvg的难点理解成:它不是让我重新学习一种神经网络,而是在普通训练循环外面,再加一层客户端和服务器的训练组织。

再看一次Algorithm 1

McMahan等人在论文Algorithm 1里,把FedAvg分成服务器端和ClientUpdate(k, w)两部分。里面几个符号先对应清楚:

符号 论文中的含义 放进程序后的直观对应
$K$ 客户端总数 客户端数据分片或DataLoader的数量
$C$ 每轮参与训练的客户端比例 服务器每轮抽样的比例
$m$ 本轮选中的客户端数量 论文写作$m=\max(CK,1)$,实际代码中要转换为整数
$E$ 每个客户端的本地epoch数 local_train()遍历本地数据的次数
$B$ 客户端本地minibatch大小 本地DataLoader的batch_size
$\eta$ 本地学习率 本地optimizer使用的学习率
$w_t$ 第$t$轮开始时的全局模型参数 当前全局模型的state_dict

服务器先初始化全局模型$w_0$。进入第$t$轮以后,它根据$C$和$K$算出$m$,随机选出客户端集合$S_t$,再把当前全局参数$w_t$交给这些客户端。论文伪代码把$m$写成$\max(CK,1)$,实际选择客户端时还需要在代码中把数量转换为整数。

每个被选中的客户端执行ClientUpdate(k, w_t)。它从同一个$wt$开始,在自己的数据上训练$E$个epoch,然后把本地训练后的模型参数$w{t+1}^k$返回服务器。服务器聚合这些结果,生成下一轮使用的$w_{t+1}$。

如果先不考虑并行、网络和掉线,主流程可以缩成下面这个骨架:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
global_model = create_model()

for round_index in range(num_rounds):
selected_clients = select_clients(client_loaders, client_fraction)
local_results = []

for client_loader in selected_clients:
local_model = copy_global_model(global_model)
local_state = local_train(local_model, client_loader)
local_results.append((local_state, len(client_loader.dataset)))

new_state = fedavg(local_results)
global_model.load_state_dict(new_state)
evaluate(global_model, test_loader)

这不是完整实现,只是把Algorithm 1里的角色摆到代码里。这里的round_index对应communication round,local_train()里面才会出现本地epoch。

普通PyTorch训练循环其实还在

单机MNIST训练时,我有一个模型、一个训练DataLoader和一个optimizer。到了FedAvg客户端,本质上变成:一份全局模型的副本、这个客户端自己的DataLoader,以及只负责这次本地训练的optimizer。

本地更新的核心仍然是以前写过的几行:

1
2
3
4
5
6
7
8
for _ in range(local_epochs):
for images, labels in client_loader:
predictions = model(images)
loss = loss_fn(predictions, labels)

optimizer.zero_grad()
loss.backward()
optimizer.step()

这里的local_epochs就是论文里的$E$,client_loader使用的batch size就是$B$。模型内部可以是MLP,也可以是CNN。loss、batch、epoch和梯度下降都没有失效,只是每个客户端现在只看自己的那部分数据。

这也解释了为什么我前面先跑单机MLP和CNN并不是绕路。FedAvg的ClientUpdate里面,做的仍然是同一套PyTorch训练。

第一次复现里的“客户端”是什么

刚接触联邦学习时,我很容易把client和server联想到传统的C/S架构,好像要先开很多台机器、建立网络连接,再写socket通信。

真实联邦学习当然可能涉及大量设备或多个机构,但第一次本地复现没有必要把系统做成这样。完全可以在一台电脑、一个Python进程里模拟多个客户端:

  • client 0保存一组MNIST样本索引;
  • client 1保存另一组索引;
  • 每组索引各自生成一个DataLoader;
  • server暂时就是最外层的训练循环。

这样做没有模拟真实网络,却保留了算法最重要的边界:每个客户端只能用自己的本地数据训练,服务器负责选择客户端、分发模型和聚合参数。

所以,算法角色上的client/server,并不等于第一次复现就一定要部署真正的网络服务。FedAvg也不是普通客户端向服务器请求一个结果,而是服务器协调多个客户端共同训练一份模型。

为什么要从同一个全局模型开始

一轮开始时,被选中的客户端都应该拿到同一个$w_t$。放到PyTorch里,可以理解成先复制模型,或者新建相同结构的模型,再用load_state_dict()加载当前全局权重。

如果client A、B、C一开始拿到的模型状态就不同,那么本地训练后再平均,就不再是在比较“从同一个起点出发,被不同本地数据分别推向了哪里”。论文也专门讨论过共同初始化对参数平均的重要性。

这里我目前只需要记住两件事:每个客户端要有独立的本地模型,避免一个客户端的训练直接改到另一个客户端;这些本地模型在同一轮开始时又必须共享同一份全局权重。

PyTorch的state_dict()可以用来保存和加载模型状态。它与optimizer自己的状态不是一回事。对第一次使用简单MLP或不带BatchNorm的CNN复现来说,我准备先把重点放在模型权重的复制和聚合上,不提前展开更复杂的buffer处理。

Local Update应该怎样理解

论文里的ClientUpdate(k, w)先把客户端$k$的数据$P_k$按大小$B$拆成多个batch,再遍历$E$个本地epoch。每个batch都执行一次学习率为$\eta$的SGD更新,最后返回本地训练后的$w$。

这里最容易混的是两种“轮”:

Communication round(通信轮次)是服务器下发全局模型、客户端本地训练、上传结果、服务器聚合的完整过程。

Local epoch(本地训练轮数)是某个客户端在一次communication round中,把自己的本地数据遍历几遍。

例如全局训练进行50个communication round,并且$E=5$,意思不是总共只训练5轮,而是每个被选中的客户端在每一次参与时,都对自己的数据训练5个epoch。

论文Algorithm 1返回的是本地训练后的模型权重$w_{t+1}^k$。平时把它泛称为“客户端上传模型更新”没有问题,但写代码时还是要分清:这里准备收集的是训练后的模型参数,不是直接把每个batch算出的梯度上传,也不是optimizer的内部状态。

Server Aggregation平均的是什么

服务器拿到的不是准确率,也不是每个客户端loss的平均值,而是各个本地模型中对应的参数张量。论文Algorithm 1的记号和这里的最小实现写法并不完全一样。实际代码只收集本轮被选中客户端的结果时,我准备先按照这些客户端的本地样本数重新归一化权重:

$n_k$是客户端$k$参与本地训练的样本数。在第一次使用简单MLP的实现里,可以遍历各个本地模型对应的参数张量,对同一个参数位置按照客户端权重加权求和,再得到新的全局参数。

比如client A有100个样本,client B有300个样本,本轮只有它们参加,那么两个权重分别是:

聚合结果是:

只有两边样本数相同时,它才等价于$(w_A+w_B)/2$。这也是上一篇里全局目标按样本数加权,在代码层面最直接的对应。

一个最小程序至少需要哪些部分

如果准备自己写,我现在认为最少需要下面几部分:

  • create_model():创建结构一致的MLP或CNN;
  • split_dataset():把MNIST训练集划给多个模拟客户端;
  • select_clients():每轮随机选择参与训练的客户端;
  • local_train():从全局权重开始,在某个客户端数据上训练;
  • fedavg():根据本地样本数聚合模型参数;
  • evaluate():在统一测试集上评估新的全局模型;
  • server training loop:重复客户端选择、本地训练和聚合。

把它们连起来,数据流就是:

1
2
3
4
5
6
7
8
9
全局模型
-> 选择客户端
-> 为每个客户端复制全局模型
-> 本地训练
-> 收集模型参数和样本数
-> 加权聚合
-> 更新全局模型
-> 在统一测试集评估
-> 进入下一轮

以后看别人的FedAvg复现时,我也准备先找这几个位置,而不是一上来从入口文件逐行读到底:数据在哪里划给客户端,客户端在哪里拿到全局模型,local update在哪里,weighted aggregation在哪里,新全局模型又在哪里评估并进入下一轮。

只要这五个位置能和Algorithm 1对上,代码的主线就不会完全丢掉。其他命令行参数、日志和绘图可以等主流程看懂以后再补。

第一次复现准备缩到多小

之前我的学习计划很容易越列越多,最后从一个FedAvg实验扩展到隐私保护、攻击防御和真实网络通信。现在看,这样反而不容易判断自己到底有没有把基础算法写对。

第一次复现我准备把范围缩到:MNIST、一个已经理解过的小型MLP、单机单进程模拟多个客户端、先做IID、不使用现成联邦学习框架,只实现最基础的FedAvg。

差分隐私、同态加密、安全多方计算、真实网络通信和复杂攻击防御先不加入。Non-IID也等基础版本跑通以后再做。这样做只是想先验证自己能不能把论文流程翻译成代码,而不是一次解决联邦学习系统的所有问题。

我现在还没有弄清楚什么

代码结构虽然比之前清楚了一点,但真正动手时还有几个问题要验证:

  • IID数据切分怎样写,才能保证客户端之间不重叠且覆盖完整训练集;
  • 每轮客户端如何随机选择,并保证实验可以复现;
  • 本地optimizer是否应该在客户端每次参与时重新创建;
  • 参数聚合怎样处理得既准确又不意外修改原模型;
  • centralized baseline和FedAvg应该怎样设置,比较才有意义;
  • 怎样确认实现不是“代码能跑”,但客户端初始化或聚合权重其实写错了。

这些问题靠继续画流程图解决不了,下一步还是要看可靠实现怎样组织这些位置,再自己写出最小版本并运行。

现在这篇只完成了一件事:把Algorithm 1和最小PyTorch程序的结构对应起来。真正能不能理解FedAvg,还得等我把它写出来、跑起来,再检查每一轮发生了什么。

论文信息

  • 论文:Communication-Efficient Learning of Deep Networks from Decentralized Data
  • 作者:H. Brendan McMahan、Eider Moore、Daniel Ramage、Seth Hampson、Blaise Agüera y Arcas
  • 会议:AISTATS 2017
  • arXiv:1602.05629v4