Skip to content

联邦学习通信效率优化

通信瓶颈是联邦学习的主要挑战。本文介绍减少通信轮次、压缩通信数据量和优化通信时序的方法。

联邦学习通信效率

通信瓶颈分析

场景模型大小每轮通信量100轮总量
小模型 (10M)40MB40MB4GB
中模型 (100M)400MB400MB40GB
大模型 (1B)4GB4GB400GB

梯度压缩

Top-K 稀疏化

python
class TopKCompressor:
    """Top-K 梯度稀疏化"""

    def __init__(self, k_ratio=0.01):
        self.k_ratio = k_ratio

    def compress(self, gradient):
        """只保留绝对值最大的 K% 梯度"""
        flat = gradient.flatten()
        k = max(1, int(len(flat) * self.k_ratio))

        values, indices = torch.topk(flat.abs(), k)
        mask = torch.zeros_like(flat)
        mask[indices] = flat[indices]

        return mask.reshape(gradient.shape), indices

    def decompress(self, sparse_grad, shape):
        """解压缩"""
        return sparse_grad.reshape(shape)

随机量化

python
class StochasticQuantizer:
    """随机量化"""

    def __init__(self, num_bits=8):
        self.num_bits = num_bits
        self.num_levels = 2 ** num_bits

    def quantize(self, gradient):
        """量化到 num_bits 位"""
        flat = gradient.flatten()
        min_val, max_val = flat.min(), flat.max()
        scale = (max_val - min_val) / (self.num_levels - 1)

        # 随机量化(无偏)
        normalized = (flat - min_val) / scale
        floor = torch.floor(normalized)
        prob = normalized - floor
        quantized = floor + (torch.rand_like(prob) < prob).float()

        return quantized.to(torch.int8), scale, min_val

减少通信轮次

本地多轮训练

增加本地训练轮数 E,减少通信轮数:

python
# E=1: 100 通信轮, 每轮 1 epoch
# E=5: 20 通信轮, 每轮 5 epochs
# E=20: 5 通信轮, 每轮 20 epochs

config = {"local_epochs": 5, "num_rounds": 20}
# 总计算量相同,但通信量减少 5x

周期性聚合

python
class PeriodicAggregation:
    """周期性聚合:每 K 步聚合一次"""

    def __init__(self, aggregation_interval=10):
        self.interval = aggregation_interval
        self.step_count = 0

    def should_aggregate(self):
        self.step_count += 1
        return self.step_count % self.interval == 0

通信与计算重叠

python
class OverlappedFL:
    """通信与计算重叠"""

    def train_step(self):
        # 1. 启动异步聚合(非阻塞)
        future = self.async_aggregate(self.pending_updates)

        # 2. 同时进行本地计算
        local_update = self.local_train(self.model)

        # 3. 等待聚合完成
        global_update = future.result()

        # 4. 应用全局更新
        self.apply_update(global_update)

        # 5. 将本地更新加入待聚合队列
        self.pending_updates.append(local_update)

压缩方法对比

方法压缩比精度损失无偏计算开销
Top-K10-100x
随机量化4-32x
随机旋转4x
哈希压缩100x+

通信优化组合

最佳实践是组合多种方法:1) Top-K 稀疏化(100x 压缩);2) INT8 量化(4x 压缩);3) 本地多轮训练(减少轮次);4) 通信计算重叠(隐藏延迟)。组合可达 1000x+ 通信节省。

压缩与收敛

过度压缩会降低收敛速度甚至导致发散。经验法则:总压缩比不超过 1000x,且 Top-K 的 K 不低于 0.1%。

相关资源

最近更新