Skip to content

跨设备联邦学习系统

跨设备联邦学习(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

跨设备联邦最佳实践

  1. 设备选择:充电+WiFi+电量充足
  2. 通信压缩:Top-K + 量化可减少 100x 通信量
  3. 异步聚合:容忍设备延迟
  4. 安全防护:异常检测 + 鲁棒聚合
  5. 隐私保护:DP 是标配

参与偏差

只有满足条件的设备(充电、WiFi)才能参与训练,导致数据分布偏差。低电量设备通常是有特定使用模式(如通勤中),样本可能不代表全体用户。

相关资源

最近更新