跨设备联邦学习系统
跨设备联邦学习(Cross-Device FL)是面向海量移动设备的联邦学习场景。Google 的 Gboard 输入法是典型应用,每天数亿设备参与模型更新。
系统挑战
| 挑战 | 原因 | 影响 |
|---|---|---|
| 设备异构 | CPU/内存/电量差异 | 训练速度不一 |
| 网络不稳定 | WiFi/4G/5G 切换 | 上传失败 |
| 设备可用性 | 用户随时关闭 | 参与率低 |
| 数据量小 | 每设备少量数据 | 本地训练不充分 |
| 安全威胁 | 恶意设备混入 | 投毒攻击 |
设备选择策略
python
class DeviceSelector:
"""设备选择策略"""
def __init__(self, target_participants=1000):
self.target = target_participants
def select(self, device_pool):
"""选择适合参与训练的设备"""
eligible = []
for device in device_pool:
# 检查设备条件
if (device.battery_level > 0.5 and # 电量 > 50%
device.on_wifi and # WiFi 连接
device.charging and # 正在充电
device.memory_available > 100 and # 内存充足
device.last_active < 3600): # 1小时内活跃
eligible.append(device)
# 随机采样(避免偏差)
if len(eligible) > self.target:
return random.sample(eligible, self.target)
return eligible通信优化
跨设备场景带宽受限,需要极度压缩通信量:
python
class GradientCompressor:
"""梯度压缩"""
@staticmethod
def top_k_sparsify(gradient, k_ratio=0.01):
"""Top-K 稀疏化:只保留最大的 K% 梯度"""
flat = gradient.flatten()
k = int(len(flat) * k_ratio)
_, indices = torch.topk(flat.abs(), k)
mask = torch.zeros_like(flat)
mask[indices] = flat[indices]
return mask.reshape(gradient.shape)
@staticmethod
def quantize(gradient, num_bits=8):
"""量化到低精度"""
min_val = gradient.min()
max_val = gradient.max()
scale = (max_val - min_val) / (2**num_bits - 1)
quantized = torch.round((gradient - min_val) / scale)
return quantized.to(torch.int8), scale, min_val异步聚合
设备上线时间不同,需要异步聚合:
python
class AsyncFedAvg:
"""异步联邦平均"""
def __init__(self, staleness_weight=0.5):
self.staleness_weight = staleness_weight
self.global_version = 0
def aggregate(self, client_update, client_version):
"""基于陈旧度的加权聚合"""
staleness = self.global_version - client_version
weight = 1.0 / (1.0 + staleness) ** self.staleness_weight
# 按权重混合
for param, update in zip(self.model.parameters(), client_update):
param.data = (1 - weight) * param.data + weight * update跨设备联邦最佳实践
- 设备选择:充电+WiFi+电量充足
- 通信压缩:Top-K + 量化可减少 100x 通信量
- 异步聚合:容忍设备延迟
- 安全防护:异常检测 + 鲁棒聚合
- 隐私保护:DP 是标配
参与偏差
只有满足条件的设备(充电、WiFi)才能参与训练,导致数据分布偏差。低电量设备通常是有特定使用模式(如通勤中),样本可能不代表全体用户。