联邦学习通信效率优化
通信瓶颈是联邦学习的主要挑战。本文介绍减少通信轮次、压缩通信数据量和优化通信时序的方法。
通信瓶颈分析
| 场景 | 模型大小 | 每轮通信量 | 100轮总量 |
|---|---|---|---|
| 小模型 (10M) | 40MB | 40MB | 4GB |
| 中模型 (100M) | 400MB | 400MB | 40GB |
| 大模型 (1B) | 4GB | 4GB | 400GB |
梯度压缩
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-K | 10-100x | 小 | 否 | 低 |
| 随机量化 | 4-32x | 小 | 是 | 低 |
| 随机旋转 | 4x | 小 | 是 | 中 |
| 哈希压缩 | 100x+ | 中 | 否 | 低 |
通信优化组合
最佳实践是组合多种方法:1) Top-K 稀疏化(100x 压缩);2) INT8 量化(4x 压缩);3) 本地多轮训练(减少轮次);4) 通信计算重叠(隐藏延迟)。组合可达 1000x+ 通信节省。
压缩与收敛
过度压缩会降低收敛速度甚至导致发散。经验法则:总压缩比不超过 1000x,且 Top-K 的 K 不低于 0.1%。