在这篇blog中我们一起来阅读一下 On the convergence of FedAvg on non-iid data 这篇 ICLR 2020 的paper.
主要目的
本文的主要目的是证明联邦学习算法的收敛性。与之前其他工作中的证明不同,本文的证明更贴近于实际联邦学习的场景。特别的,
- 所有用户的数据non-iid分布;
- 每次只有一部分用户参与FedAvg.
系统模型
考虑一个联邦学习系统 with NNN 用户和一个PS. 每用户有一些local data,训练发生在用户处,每隔一段时间用户上传自己学习的模型来做FedAvg.
将第 kkk 个用户的数据记为 x={
xk,1,xk,2,xk,3,...,xk,nk}\bm{x}=\{x_{k,1},x_{k,2},x_{k,3},...,x_{k,n_k}\}x={
xk,1,xk,2,xk,3,...,xk,nk}, 每个人都有一个学习目标,即最小化 loss 函数
Fk(w)=∑j=1nkℓk(w,xk,j)(1)F_k(\bm{w})=\sum_{j=1}^{n_k}\ell_k(\bm{w},x_{k,j}) \tag{1}Fk(w)=j=1∑nkℓk(w,xk,j)(1)
其中 ℓk(w,xk,j)\ell_k(\bm{w},x_{k,j})ℓk(w,xk,j) 是每个训练数据的 loss. Fk(w)F_k(\bm{w})Fk(w) 相当于是每个人所有数据上的loss,如果仅仅做local training, 那最终每个用户会 arrive at
Local minimum: Fk∗=minwFk\text{Local minimum}:~~~~F_k^*=\min_{\bm{w}} F_k Local minimum: Fk∗=wminFk
而FL考虑的是一种分布式的优化,即我们要minimize的目标函数
Global minimum: F∗=minw∑k=1NpkFk(w)\text{Global minimum}:~~~~F^*=\min_{\bm{w}} \sum_{k=1}^{N} p_k F_k(\bm{w}) Global minimum: F∗=wmink=1∑NpkFk(w)
其中 pkp_kpk 是一个distribution用来表示每个用户所占的权重。换句话说,我们最终想找到一个共同的 w\bm{w}w 来最小化每个用户 loss 的一个加权和。
To this end, 本文考虑FedAvg, 并证明其能收敛到 global optimum.
FedAvg 的具体步骤描述如下:首先,我们按单次SGD为一个时间刻度把时间轴分为离散的slot t=1,2,3,...,Tt=1,2,3,...,Tt=1,2,3,...,T, 即总共进行 TTT 次 local SGD, 每次 SGD每个用户从自己的数据集中随机均匀的采样出一个数据来进行训练。特别的,每隔 EEE slots, 所有 active users 把自己的本地参数发送给PS进行 FedAvg,之后PS会把avg后的参数发还给各个用户。以上模型用数学语言可以写为以下两步:
Local training
每个用户在第 ttt 个时刻基于 wtk\bm{w}^k_twtk 进行 SGD, 得到
vt+1k=wtk−ηt∇ℓk(wtk,ξtk)(2)\bm{v}^k_{t+1}=\bm{w}^k_t-\eta_{t}\nabla \ell_k(\bm{w}^k_t,\xi^k_t) \tag{2}vt+1k=wtk−ηt∇ℓk(wtk,ξtk)(2)
其中 ξtk\xi^k_tξtk 是从本地数据中随机采样出的一个sample。注意,这样单步SGD得到的vt+1k\bm{v}^k_{t+1}vt+1k 只是一个中间变量而不是下一时刻的 wt+1k\bm{w}^k_{t+1}wt+1k,因为我们还有可能做 FedAvg。 更具体地说,在 EEE 的非整数倍slot上,
wt+1k=vt+1k, if t+1∉JE={
nE:n=1,2,...}.\bm{w}^k_{t+1}=\bm{v}^k_{t+1},~~~~\text{if}~~t+1\notin\mathcal{J}_E=\{nE:n=1,2,...\}.wt+1k=vt+1k, if t+1∈/JE={
nE:n=1,2,...}.
而在 EEE 的整数倍slot上,我们还得额外做 FedAvg.
FedAvg
若下一时刻是EEE的整数倍周期,即 t+1∈JE={
nE:n=1,2,...}t+1\in\mathcal{J}_E=\{nE:n=1,2,...\}t+1∈JE={
nE:n=1,2,...},我们进行FedAvg,此时
wt+1k=∑k=1Npkvt+1k(3)\bm{w}^k_{t+1}=\sum_{k=1}^N p_k \bm{v}^k_{t+1} \tag{3}wt+1k=k=1∑Npkvt+1k(3)
注意,这里面我们假设每个人都参与更新,稍后我们会release这个条件允许PS按照某种分布采样一部分人进行更新。
小结
如果我们从每个用户的角度看,它的参数变化可以用下图归纳 (E=3E=3E=3)。

几个假设
本文的推导基于以下假设。
Assumption 1 (LLL-smoothness). 所有用户的 loss 函数 { Fk:k=1,2,...,N}\{F^k:k=1,2,...,N\}{ Fk:k=1,2,...,N} 都是 L-smooth.
Fk(x2)−Fk(x1)≤∇f(x1)⊤(x2−x1)+L2∥x2−x1∥2F^k(\bm{x_2})-F^k(\bm{x_1})\leq \nabla f(\bm{x_1})^\top (\bm{x_2-x_1}) + \frac{L}{2}\|\bm{x_2-x_1}\|^2Fk(x2)−Fk(x1)≤∇f(x1)⊤(x2−x1)+2L∥x2−x1∥2
Assumption 2 (μ\muμ-strongly convex). 所有用户的 loss 函数 { Fk:k=1,2,...,N}\{F^k:k=1,2,...,N\}{ Fk:k=1,2,...,N} 都是 μ\muμ-strongly convex.
Fk(x2)−Fk(x1)≥∇f(x1)⊤(x2−x1)+μ2∥x2−x1∥2F^k(\bm{x_2})-F^k(\bm{x_1})\geq \nabla f(\bm{x_1}


3444

被折叠的 条评论
为什么被折叠?



