数据不能出域,但模型想变聪明——联邦学习就是为这个矛盾设计的。它把训练拆成「本地算梯度、中心聚合参数」的循环,原始数据始终留在各参与方手里。听起来是个优雅的分布式训练特例,但真正落地时会发现:通信是瓶颈、客户端是异构的、数据是非独立同分布的、隐私还需要额外的数学保证。这四件事把联邦学习从「分布式 SGD 换个说法」变成了一个独立的工程学科。
联邦学习最常见的误解是「它天然保护隐私」。事实恰恰相反:原始梯度可以反演出训练数据,明文聚合的联邦学习在隐私上几乎不设防。要真正保护隐私,必须叠加差分隐私或安全聚合——这是两个独立的技术,不是可选项。
联邦学习的三种形态
横向、纵向与迁移
| 形态 | 数据划分方式 | 典型场景 |
|---|---|---|
| 横向联邦 | 样本不同、特征相同 | 多家医院同一套检查指标 |
| 纵向联邦 | 样本相同、特征不同 | 银行 + 电商的同一批用户 |
| 联邦迁移 | 样本与特征都不同 | 跨领域知识迁移 |
横向联邦最常见,也是 FedAvg 的原生场景:各客户端有同类数据的不同样本,训练同一个模型。若想先建立 FedAvg 与 DP-SGD 的数学直觉,可先读 联邦学习与隐私保护机器学习 ,那里从优化目标与隐私会计的角度做了更完整的推导。
纵向联邦的难点完全不同:需要先做样本对齐(Private Set Intersection)——找出共同的用户,且不能泄露各自的用户列表。训练时各方只持有部分特征,需要交换中间表示(梯度或嵌入)而非原始特征。工程复杂度高一个量级。
集中式与去中心化拓扑
- 中心聚合:有一个服务器负责聚合。简单,但服务器是单点与信任焦点。
- 去中心化:客户端组成 P2P 网络,用 gossip 协议交换参数。无单点,但收敛慢、难调试。
生产系统绝大多数用中心聚合——它的可观测性与可控性远超去中心化方案带来的那点去信任收益。
系统架构
一次联邦轮次的生命周期
1. 服务器选择客户端子集(按可用性、数据量、历史贡献)
2. 下发全局模型参数
3. 客户端本地训练 E 个 epoch
4. 上传模型更新(或梯度)
5. 服务器聚合,更新全局模型
6. 记录指标,进入下一轮
客户端选择策略
不是每轮都让所有客户端参与——那既慢又不必要。选择策略直接影响收敛速度与公平性:
import random
def select_clients(clients, k, strategy="weighted", history=None):
if strategy == "random":
return random.sample(clients, k)
if strategy == "weighted":
# 按数据量加权,数据多的客户端贡献更大
weights = [c.n_samples for c in clients]
return random.choices(clients, weights=weights, k=k)
if strategy == "fair":
# 优先选参与次数少的,保证公平性
counts = history or {c.id: 0 for c in clients}
return sorted(clients, key=lambda c: counts[c.id])[:k]
if strategy == "availability":
# 优先选在线且带宽好的
return sorted(clients, key=lambda c: -c.bandwidth)[:k]
客户端选择是联邦学习里被严重低估的设计点。随机选择会让快客户端反复被选中、慢客户端永远边缘化,最终模型偏向快客户端的数据分布。工程上常用「加权随机 + 公平性补偿」的组合。
聚合算法
FedAvg 是最基础的聚合方式,按数据量加权平均:
def fedavg(updates):
"""updates: [(n_samples, state_dict), ...]"""
total = sum(n for n, _ in updates)
agg = {}
for key in updates[0][1].keys():
agg[key] = sum(sd[key] * (n / total) for n, sd in updates)
return agg
它的隐含假设是各客户端数据同分布。一旦分布差异大,加权平均会把不同方向的最优解互相抵消。
通信压缩
通信是联邦学习的第一瓶颈。客户端往往在移动网络或跨地域专线上,上行带宽极其有限。
三种压缩路线
| 路线 | 压缩比 | 是否无损 | 复杂度 |
|---|---|---|---|
| 量化 | 4~32 倍 | 有损 | 低 |
| 稀疏化 | 10~100 倍 | 有损 | 低 |
| 低秩分解 | 5~20 倍 | 有损 | 中 |
| 梯度草图 | 可变 | 有损 | 中 |
量化:从 FP32 到 INT8
最朴素的量化是均匀标量量化:
import torch
def quantize(tensor, bits=8):
qmax = 2 ** bits - 1
t_min, t_max = tensor.min(), tensor.max()
scale = (t_max - t_min) / qmax
zero = torch.round(-t_min / scale)
q = torch.round(tensor / scale + zero).clamp(0, qmax).to(torch.uint8)
return q, scale, zero
def dequantize(q, scale, zero):
return (q.float() - zero) * scale
误差反馈(Error Feedback) 是量化能用的关键:把量化误差累积到下一轮补偿,避免误差持续偏置某一方向。
class ErrorFeedback:
def __init__(self):
self.residual = {}
def compress(self, name, grad):
g = grad + self.residual.get(name, torch.zeros_like(grad))
q, scale, zero = quantize(g)
deq = dequantize(q, scale, zero)
self.residual[name] = g - deq # 把误差留到下一轮
return q, scale, zero
稀疏化:只传 Top-K
只上传绝对值最大的 K 个梯度,其余置零。稀疏度通常 99% 以上仍能保持收敛:
def topk_sparsify(grad, ratio=0.01):
k = max(1, int(grad.numel() * ratio))
values, indices = grad.abs().flatten().topk(k)
mask = torch.zeros_like(grad).flatten()
mask[indices] = 1
mask = mask.view_as(grad)
return grad * mask, mask
同样需要误差反馈,否则被丢掉的梯度永远不参与更新。
压缩与隐私的交互
这里有个反直觉的坑:压缩会削弱差分隐私的保护。因为压缩把梯度投影到低维空间,攻击者可以从压缩后的表示里更高效地反推信息。实践中的正确顺序是先加噪再压缩——在加噪后的梯度上做量化与稀疏化,这样压缩过程本身不会引入额外的信息泄露。
异步与半异步更新
同步聚合要求等所有选中客户端,一个慢客户端就能拖垮整轮。
三种同步模式
| 模式 | 等待行为 | 收敛 | 容错 |
|---|---|---|---|
| 同步 | 等全部 | 稳 | 差 |
| 异步 | 不等 | 快但震荡 | 好 |
| 半异步 | 等 K 个 | 折中 | 好 |
异步的梯度陈旧问题
异步模式下,服务器用某个客户端基于旧参数算出的梯度去更新当前参数,产生陈旧梯度(Stale Gradient)。陈旧度越大,更新方向越偏。
def async_update(global_params, client_update, client_version, cur_version, alpha=0.5):
"""按陈旧度衰减学习率"""
staleness = cur_version - client_version
decay = alpha ** staleness
return global_params - decay * client_update
半异步:有界陈旧
半异步是工程上的甜点:设定一个超时窗口,收集到 K 个更新就聚合,超时的客户端被丢弃:
import asyncio
async def semi_async_round(clients, k, timeout=60.0):
tasks = [asyncio.create_task(c.train_and_upload()) for c in clients]
done, pending = await asyncio.wait(tasks, timeout=timeout,
return_when=asyncio.FIRST_COMPLETED)
updates = [t.result() for t in done if t.exception() is None]
for t in pending:
t.cancel()
if len(updates) < k: # 不足 K 个,降低学习率或跳过
return None
return updates
关键参数是超时窗口与最小客户端数。窗口太短会频繁凑不齐,太长则退化成同步。经验上把窗口设为客户端训练耗时的 P90,能兼顾参与率与效率。
差分隐私
从 DP-SGD 说起
差分隐私的核心是梯度裁剪 + 加噪:
def dp_sgd_step(model, batch, optimizer, clip_norm=1.0, noise_multiplier=1.0):
loss = compute_loss(model, batch)
optimizer.zero_grad()
loss.backward()
# 1. 逐样本裁剪
per_sample_norms = per_sample_grad_norms(model, batch)
scale = (clip_norm / (per_sample_norms + 1e-6)).clamp(max=1.0)
apply_per_sample_scaling(model, scale)
# 2. 加高斯噪声
sigma = noise_multiplier * clip_norm
for p in model.parameters():
if p.grad is not None:
p.grad.add_(torch.randn_like(p.grad) * sigma / batch_size)
optimizer.step()
三个超参决定隐私-效用权衡:
| 参数 | 作用 | 调大后果 |
|---|---|---|
| clip_norm | 限制单样本影响力 | 噪声增大,效用降 |
| noise_multiplier | 噪声强度 | 隐私更强,效用降 |
| batch_size | 摊销噪声 | 显存与算力增 |
联邦场景下的 DP
联邦学习有天然的优势:噪声可以在客户端本地加,隐私预算按客户端累加。这比集中式 DP 更容易论证,因为服务器永远看不到原始梯度。
但要小心隐私预算的会计:每轮每个客户端的参与都会消耗预算。总预算 ε 是跨轮累加的,不能只看单轮。
from opacus.accountants import RDPAccountant
def track_privacy(steps, sample_rate, noise_multiplier, delta=1e-5):
accountant = RDPAccountant()
for _ in range(steps):
accountant.step(noise_multiplier=noise_multiplier, sample_rate=sample_rate)
return accountant.get_epsilon(delta=delta)
客户端级 DP
联邦学习里更贴切的定义是客户端级差分隐私:保护的是「某个客户端是否参与」,而非「某条样本」。这需要先做客户端级裁剪(限制单个客户端的整体贡献),再加噪。代价是噪声量级大得多,因为客户端贡献的方差远大于单样本。
安全聚合
差分隐私解决「加噪后仍可能被反推」,安全聚合解决「服务器不该看到单个客户端的明文更新」。
核心思想:掩码抵消
每个客户端对之间协商一个随机掩码 s_ij,客户端 i 上传 x_i + Σ_j s_ij。所有客户端求和时,两两掩码恰好抵消:
Σ_i (x_i + Σ_j s_ij) = Σ_i x_i
服务器只能看到总和,看不到任何单个 x_i。
import hashlib
def pairwise_mask(client_id, peer_ids, seed, param_shape):
"""基于共享种子的伪随机掩码,无需真实通信"""
mask = torch.zeros(param_shape)
for peer in peer_ids:
key = hashlib.sha256(f"{min(client_id, peer)}:{max(client_id, peer)}:{seed}"
.encode()).digest()
gen = torch.Generator().manual_seed(int.from_bytes(key[:8], "big"))
s = torch.randn(param_shape, generator=gen)
mask += s if client_id < peer else -s # 符号决定加减
return mask
掉线处理
安全聚合的致命弱点:任何一个客户端掉线,掩码就无法抵消,聚合结果全错。工程上的解法是 Shamir 秘密共享——把掩码拆成份额分发给其他客户端,掉线时由幸存者恢复。
def shamir_split(secret, n, t):
"""把 secret 拆成 n 份,任意 t 份可恢复"""
coeffs = [secret] + [random_scalar() for _ in range(t - 1)]
shares = [(i, poly_eval(coeffs, i)) for i in range(1, n + 1)]
return shares
def shamir_recover(shares, t):
return lagrange_interpolate(shares[:t])
DP 与安全聚合的组合
两者互补,应当同时启用:
| 技术 | 防的是谁 | 失效场景 |
|---|---|---|
| 安全聚合 | 好奇的服务器 | 多数客户端合谋 |
| 差分隐私 | 任意后处理与合谋 | 噪声过大效用崩 |
推荐组合:客户端本地做梯度裁剪 + 加噪(DP),再用安全聚合上传(SecAgg)。这样服务器既看不到单个更新,也无法从总和中反推个体。若还需要在多方之间做不泄露输入的联合计算,可进一步了解 隐私计算 中的多方安全计算与可信执行环境方案。
非独立同分布挑战
为什么 Non-IID 致命
假设两个客户端的真实最优参数方向相反,FedAvg 把它们平均,得到的是一个对谁都不好的折中。极端情况下(每个客户端只有单一类别),FedAvg 会震荡不收敛。
缓解手段
| 方法 | 原理 | 代价 |
|---|---|---|
| FedProx | 加近端项约束本地模型别跑太远 | 略降本地拟合 |
| SCAFFOLD | 用控制变量修正梯度方向 | 需额外通信 |
| MOON | 用对比学习对齐全局表示 | 需存两份模型 |
| 数据共享 | 共享少量公共数据 | 需隐私权衡 |
| 个性化 | 每人保留自己的头部 | 失去单一全局模型 |
FedProx:最简单的改动
在本地损失上加一项,约束本地参数不要偏离全局太远:
def fedprox_local_loss(model, global_params, batch, mu=0.01):
task_loss = compute_loss(model, batch)
prox = 0.0
for name, p in model.named_parameters():
prox += ((p - global_params[name]) ** 2).sum()
return task_loss + (mu / 2) * prox
mu 是近端系数:mu=0 退化成 FedAvg,mu 越大越接近全局模型。实践中 0.01~0.1 是常用范围,太小不起作用,太大等于没做本地训练。
个性化:承认异构
当异构太严重时,追求单一全局模型本身就是错的方向。个性化联邦学习给每个客户端保留一个专属头部(或低秩适配器),只共享主干:
class PersonalizedClient:
def __init__(self, global_backbone, d_feat, n_class):
self.backbone = global_backbone # 共享,参与聚合
self.head = torch.nn.Linear(d_feat, n_class) # 私有,不聚合
def local_train(self, data, epochs=3):
# 只训练 backbone + head,但上传时只传 backbone
...
这种做法与参数高效微调(LoRA)的思路一致——共享大部分参数,个性化小部分。
工程实现:用 Flower 搭一个可跑的联邦任务
Flower 是当前最易上手的联邦学习框架。
服务端
import flwr as fl
def weighted_average(metrics):
total = sum(n for n, _ in metrics)
acc = sum(n * m["accuracy"] for n, m in metrics) / total
return {"accuracy": acc}
strategy = fl.server.strategy.FedProx(
fraction_fit=0.5,
min_fit_clients=4,
min_available_clients=8,
on_fit_config_fn=lambda rnd: {"local_epochs": 2, "lr": 1e-3},
proximal_mu=0.05,
evaluate_metrics_aggregation_fn=weighted_average,
)
fl.server.start_server(
server_address="0.0.0.0:8080",
config=fl.server.ServerConfig(num_rounds=50),
strategy=strategy,
)
客户端
import flwr as fl
import torch
class FlowerClient(fl.client.NumPyClient):
def __init__(self, model, train_loader, val_loader):
self.model = model
self.train_loader = train_loader
self.val_loader = val_loader
def get_parameters(self, config):
return [v.cpu().numpy() for v in self.model.state_dict().values()]
def fit(self, parameters, config):
set_params(self.model, parameters)
for _ in range(config.get("local_epochs", 1)):
train_one_epoch(self.model, self.train_loader, lr=config.get("lr", 1e-3))
return self.get_parameters(config), len(self.train_loader.dataset), {}
def evaluate(self, parameters, config):
set_params(self.model, parameters)
loss, acc = evaluate(self.model, self.val_loader)
return float(loss), len(self.val_loader.dataset), {"accuracy": float(acc)}
fl.client.start_client(server_address="127.0.0.1:8080",
client=FlowerClient(build_model(), train_loader, val_loader).to_client())
加入安全聚合
Flower 支持策略层叠加 SecAgg+ 与差分隐私,只需在策略外面包一层修饰器,客户端代码几乎不用改。这正是抽象层的价值——把隐私机制做成可插拔的横切关注点。
生产落地与排错
上线前必须回答的问题
- 谁是服务器:自建还是某方托管?托管方的信任级别决定了要不要强制安全聚合。
- 客户端可信吗:恶意客户端可以上传毒化梯度(投毒攻击),需要鲁棒聚合(Krum、Trimmed Mean)。
- 掉线率多高:掉线率超过 30% 时安全聚合的恢复机制必须可靠。
- 隐私预算是多少:合规要求决定
ε的上限,进而决定噪声与效用损失。 - 如何评估:联邦模型无法在中心侧看到全量数据,需要设计「留出客户端」的评估协议。
常见故障
- 训练不收敛:数据极度 Non-IID。先上 FedProx,再考虑个性化。
- 通信量超预期:未启用压缩,或稀疏化的误差反馈失效导致重复上传。
- 聚合结果异常:安全聚合下客户端掉线未恢复。检查秘密共享的阈值配置。
- 精度远低于集中式:检查客户端选择是否偏向少数节点,或本地 epoch 过多导致客户端漂移。
- 隐私预算耗尽:轮次过多。减少轮数、提高每轮本地 epoch,或放宽
noise_multiplier。 - 毒化攻击:某客户端上传巨大梯度主导聚合。改用 Trimmed Mean 或 Krum 鲁棒聚合。
监控指标
联邦场景的监控比集中式复杂,因为中心看不到数据。关键指标包括:每轮参与客户端数、上传成功率、客户端更新的范数分布(异常大值可能意味着投毒)、全局模型在留出客户端上的精度、以及累计隐私预算。这套监控与 MLOps 治理 中的模型可观测性体系可以共用基础设施。
小结
联邦学习把「数据不动模型动」这件事做成了工程,但它换来的不是免费午餐,而是三笔明确的账单:通信(压缩与异步更新)、异构(FedProx 与个性化)、隐私(差分隐私与安全聚合的组合)。这三笔账必须一起算——只做通信压缩不做隐私保护,等于在明文管道上跑敏感梯度;只做隐私不做压缩,则受限于带宽而无法扩展到足够多的客户端。把它和 分布式训练 放在一起看,会发现两者共享同一套并行与通信直觉,只是联邦学习把「信任边界」加进了系统设计的第一性约束里。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。