Improving Federated Learning Personalization via MAML#
Abstract#
- FL算法与MAML具有很多相似性,可以用元学习算法来对其进行解释
- 微调可以使得gloabl 模型具有更强的准确率,同时更容易做定制化处理
- 通过标准的中心化数据库训练出来的模型相比Fedavg训练的更难进行定制化处理
Introduction#
- 指出了FL与MAML算法的联系,并用MAML算法对FL算法进行解释
- 对FedAvg进行改进,采用两阶段的训练和fine-tune进行优化
- 发现FedAvg其实本质是一种metalearning算法,用于优化个性化定制的效果,而不是全局模型的优化。
下图展现了在FL中应用MAML算法(左侧),Reptile算法(中间)和FL的训练算法FedAvg(右侧)。设L为损失函数,在每一轮的迭代中,MAML会通过随机采样一个batch的任务T来进行训练。对于每个任务T,会有一个内循环,然后在外循环中聚集每个任务所获得的的梯度更新。对于FL算法会随机采样数个client T。对于每个T和其权重,会在local数据上进行数轮的迭代优化,然后将更新的梯度聚集形成一个新的global model。如果我们简化设置,并认为所有的client拥有相同的数据,那么所有的权重就会一样,这个时候reptile和fedavg其实就是同一种算法。

假设在FedAvg中的权重相同为wi。考虑有T个clients,并设置每个相关模型参数为。对于每个cilent i,其损失函数为,记为第local训练过程所计算得到的梯度。
FedSGD的梯度更新函数为:
将设我们将FOMAML用相同的术语来表示。假设Client 学习率为, 每个client的个性化模型经过K步后所获得的梯度更新为
求微分可得到:
在进行K次梯度更新后,对整个模型进行更新:
为了避免二次求导带来的计算量问题,FOMAML应运而生,通过K次的梯度更新后,直接采用第K+1次的梯度更新作为local update。
通过上面的公式,不难看出,其实FedAvg的更新,所有client的更新的平均,其实就是以上两种idea的线性组合。
Personalized FedAvg#

如上图所示,采用算法1中的FedAvg E训练E个local epoch,根据local数据量来对梯度更新进行权衡。然后在FL的环境下采用Retile(K)训练K个local steps,不考虑本地的数据量。
一般来说,就通信轮次的数量而言,FedAvg训练数个local epochs后,可以在数轮通信内就能快速收敛。由于生产环境的复杂性,这种测量方式被用于衡量FL算法的收敛速度。本文发现,采用momentum SGD的方法作为server优化器已经对于personalized model进行了优化,然而initial model相对不稳定。以前的方法是减少本地的训练轮次或者学习率。
本文提出采用Retile(K)的方法进行fintune,然后用Adam作为server优化器,来提升initialmodel的效果。同时可以稳定personalized model。
To be continued