FedAvg 算法原理与收敛性分析
FedAvg(Federated Averaging)是联邦学习的基础算法,由 McMahan 等人在 2017 年提出。它通过加权平均各客户端的本地模型参数来更新全局模型。
算法流程
输入:全局模型 w_0,客户端集合 C,每轮参与数 K,本地训练轮数 E
输出:全局模型 w_T
for t = 0, 1, ..., T-1 do
1. 随机选择 K 个客户端 S_t ⊂ C
2. 将 w_t 广播给 S_t 中的每个客户端
3. for 每个客户端 k ∈ S_t 并行 do
w_k^(t+1) = ClientUpdate(w_t) // 本地训练 E 轮
end
4. w_(t+1) = Σ_{k∈S_t} (n_k/n) × w_k^(t+1) // 加权平均
end实现代码
python
import numpy as np
from copy import deepcopy
class FedAvgServer:
def __init__(self, model, num_clients, fraction=0.1):
self.global_model = model
self.num_clients = num_clients
self.fraction = fraction
def select_clients(self):
"""随机选择部分客户端"""
num_selected = max(1, int(self.fraction * self.num_clients))
return np.random.choice(
self.num_clients, size=num_selected, replace=False
)
def aggregate(self, client_models, client_weights):
"""FedAvg 加权聚合"""
total_weight = sum(client_weights)
aggregated = deepcopy(client_models[0])
for param in aggregated.parameters():
param.data.zero_()
for model, weight in zip(client_models, client_weights):
ratio = weight / total_weight
for agg_param, client_param in zip(
aggregated.parameters(), model.parameters()
):
agg_param.data += ratio * client_param.data
return aggregated收敛性分析
FedAvg 的收敛上界:
F(w_T) - F* ≤ O(1/√(KTE)) + O(E/T) + O(σ/√(KT))
其中:
- K: 每轮参与客户端数
- T: 通信轮数
- E: 本地训练轮数
- σ: 数据异构程度关键结论:
- E 越大:本地训练越充分,但数据偏移越大
- K 越大:每轮参与越多,收敛越快
- σ 越大:数据越异构,收敛越慢
本地训练轮数选择
| E 值 | 通信轮数 | 收敛质量 | 适用场景 |
|---|---|---|---|
| 1 | 多 | 最优 | IID 数据 |
| 5 | 中 | 好 | 轻度异构 |
| 20 | 少 | 较差 | 高度异构 |
FedAvg 调优建议
- IID 数据:E=5-10,C=0.1(10% 参与率)
- Non-IID 数据:E=1-5,C=0.3-0.5(增加参与率)
- 通信受限:增大 E,减少通信轮数
- 数据量差异大:使用加权平均(按样本数加权)
Non-IID 陷阱
当客户端数据分布差异很大时,FedAvg 的本地训练会导致客户端模型偏离全局最优。E 过大时,聚合后的模型可能比上一轮更差。此时应使用 FedProx 或降低 E。