引言
机器学习的效果依赖数据,但最有价值的数据往往不能集中——医院的病历、银行的交易、手机上的输入记录,出于隐私、合规、商业竞争的考虑都不能汇总到一个中心。联邦学习(Federated Learning)给出的答案是:数据不动、模型动——把模型下发到各参与方本地训练,只上传梯度或模型更新,在中心聚合成全局模型。
这条路听起来很美,落地却充满陷阱:各方的数据分布不同(非 IID)会让聚合出的模型变差、上传的梯度可能反推出原始数据、通信成本高昂、参与方可能掉线。本文按「原理 → 挑战 → 隐私加固 → 协作模式 → 落地」的顺序,把联邦学习从 FedAvg 讲到 DP-SGD 与安全聚合,最后给出框架选型与合规要点。
前置:非 IID 数据带来的分布偏移与监控见 https://plumephp.com/ml-model-monitoring-drift/;特征对齐与纵向联邦的特征工程见 https://plumephp.com/ml-feature-store/;模型上线与版本管理见 https://plumephp.com/ml-model-deployment/;评估协议设计见 https://plumephp.com/ml-model-evaluation/。
目录
- 1. 为什么需要联邦学习
- 2. FedAvg 与联邦优化
- 3. 非 IID 挑战
- 4. 差分隐私与 DP-SGD
- 5. 安全聚合
- 6. 横向、纵向与联邦迁移
- 7. 隐私-效用权衡
- 8. 框架选型与合规落地
- 9. 总结
- 延伸阅读
1. 为什么需要联邦学习
1.1 数据孤岛的三重成因
| 成因 | 例子 | 约束来源 |
|---|---|---|
| 隐私法规 | 医疗、金融数据 | GDPR、个人信息保护法 |
| 商业竞争 | 各平台用户行为 | 数据是核心资产 |
| 工程现实 | 海量终端数据 | 上传成本、带宽 |
联邦学习让各方在不共享原始数据的前提下共同训练一个模型,各方都能受益于更大的「数据规模」。
1.2 两种典型形态
- 跨设备(Cross-device):手机、IoT 等海量终端,样本多、单设备数据少、连接不稳定、掉线频繁;
- 跨机构(Cross-silo):医院、银行等少数机构,样本少但每个机构数据量大、连接稳定、可信度较高。
1.3 联邦学习不等于隐私保护
一个常见误解:联邦学习本身不保证隐私。上传的梯度可能被反演出原始数据(梯度泄露攻击、DLG 攻击)。真正的隐私保证需要叠加差分隐私(加噪)或安全聚合(加密)。联邦学习是「数据不出域」的工程架构,隐私加固是它上面的一层。
一句话:联邦学习解决的是「数据不能集中」的工程问题,它本身不是隐私技术——要防梯度反演,必须再叠加差分隐私或安全聚合。
2. FedAvg 与联邦优化
2.1 FedAvg 的流程
FedAvg(Federated Averaging)是最基础的联邦算法,一轮训练分四步:
1. 服务器下发全局模型 w_t 给选中的 K 个客户端
2. 每个客户端用本地数据做 E 轮 SGD,得到 w_t^k
3. 客户端上传更新(或模型)
4. 服务器按样本量加权平均:w_{t+1} = Σ (n_k / n) · w_t^k
加权平均是关键:数据多的客户端权重更大。相比「每轮只做一次本地 SGD 再平均」(FedSGD),FedAvg 让客户端多跑几轮,大幅减少通信轮数。
2.2 手写一个 FedAvg
import copy
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
def local_train(model, dataloader, epochs, lr):
"""客户端本地训练,返回更新后的 state_dict"""
local_model = copy.deepcopy(model)
opt = torch.optim.SGD(local_model.parameters(), lr=lr, momentum=0.9)
local_model.train()
for _ in range(epochs):
for x, y in dataloader:
opt.zero_grad()
loss = nn.functional.cross_entropy(local_model(x), y)
loss.backward()
opt.step()
return local_model.state_dict()
def fed_avg(global_state, client_states, client_sizes):
"""按样本量加权平均所有客户端的模型参数"""
total = sum(client_sizes)
new_state = copy.deepcopy(global_state)
for key in global_state.keys():
new_state[key] = sum(
cs[key] * (n / total)
for cs, n in zip(client_states, client_sizes)
)
return new_state
# 主循环
for round_t in range(num_rounds):
selected = sample_clients(clients, fraction=0.1) # 每轮采样部分客户端
states, sizes = [], []
for c in selected:
sd = local_train(global_model, c.loader, epochs=3, lr=0.01)
states.append(sd)
sizes.append(len(c.dataset))
new_sd = fed_avg(global_model.state_dict(), states, sizes)
global_model.load_state_dict(new_sd)
2.3 通信是主要瓶颈
联邦学习的成本几乎全在通信:模型越大、轮数越多、客户端越多,通信越贵。优化手段包括:
| 手段 | 做法 | 收益 |
|---|---|---|
| 增加本地轮数 E | 客户端多跑几步 | 减少通信轮数 |
| 梯度压缩 | 量化 / 稀疏化 | 每轮传输量↓ |
| 客户端采样 | 每轮只用一部分 | 单轮成本↓ |
| 异步更新 | 不等慢客户端 | 避免拖尾 |
一句话:FedAvg = 「本地多训 + 加权平均」,加权体现数据量、本地轮数换来通信节省;联邦优化的核心 KPI 是「达到目标精度需要多少通信轮数」。
3. 非 IID 挑战
3.1 什么是非 IID
IID(独立同分布)假设各方数据分布一致。但现实中:某医院只见过本地病种、某用户只打某类字——各方数据分布差异巨大。这会导致客户端漂移(client drift):每个本地模型都朝自己数据的最优方向跑偏,平均后反而离全局最优更远。
IID :各方数据分布一致 → 平均即全局最优
非IID:各方分布不同 → 各自跑偏 → 平均后震荡、收敛慢、精度掉
3.2 非 IID 的三种形态
| 形态 | 描述 | 例子 |
|---|---|---|
| 标签分布偏移 | 各方类别比例不同 | 某医院只有某几种病 |
| 特征分布偏移 | 同样的类但特征不同 | 不同地区的口音 |
| 数量不平衡 | 样本量差异大 | 大机构 vs 小机构 |
3.3 缓解方法
| 方法 | 思路 |
|---|---|
| FedProx | 本地损失加近端项,限制偏离全局模型 |
| SCAFFOLD | 用控制变量校正客户端漂移 |
| FedNova | 归一化各方更新步数 |
| 数据共享 | 各方共享少量公共数据(小代价大收益) |
| 个性化 FL | 全局模型 + 本地微调(pFL) |
# FedProx:本地损失加入近端项,mu 控制约束强度
def fedprox_local_loss(model, global_params, x, y, mu=0.01):
ce = nn.functional.cross_entropy(model(x), y)
prox = sum(
((p - gp) ** 2).sum()
for p, gp in zip(model.parameters(), global_params)
)
return ce + (mu / 2) * prox
一句话:非 IID 是联邦学习区别于分布式训练的核心难题——各方数据分布不同导致客户端漂移,靠近端约束(FedProx)、控制变量(SCAFFOLD)或少量共享数据来缓解。
4. 差分隐私与 DP-SGD
4.1 差分隐私的定义
差分隐私(DP)给出可证明的隐私保证:改变数据集中的一条记录,算法输出的分布变化不超过某个界限。
(ε, δ)-DP:对任意相邻数据集 D、D',Pr[M(D) ∈ S] ≤ e^ε · Pr[M(D') ∈ S] + δ
ε 越小隐私越强(噪声越大),δ 是失败概率(通常 1e-5)
4.2 DP-SGD:给梯度加噪
DP-SGD 的两个核心操作:逐样本梯度裁剪(限制单个样本影响)+ 加高斯噪声(掩盖个体贡献)。
1. 对每个样本单独算梯度
2. 裁剪到 L2 范数 ≤ C
3. 求和后加 N(0, σ²C²) 噪声,再用加噪梯度更新
import torch
def dp_sgd_step(model, batch, optimizer, clip_norm=1.0, noise_multiplier=1.1):
"""简化版 DP-SGD 单步"""
x, y = batch
optimizer.zero_grad()
# 逐样本梯度 + 裁剪(实际用 torch.func / opacus 实现)
per_sample_grads = compute_per_sample_grads(model, x, y)
clipped = []
for g in per_sample_grads:
norm = g.norm(2)
factor = min(1.0, clip_norm / (norm + 1e-6))
clipped.append(g * factor)
# 求和 + 加噪
summed = torch.stack(clipped).sum(dim=0)
noise = torch.randn_like(summed) * (noise_multiplier * clip_norm)
summed = summed + noise
# 把加噪梯度写回参数并更新
assign_grads(model, summed)
optimizer.step()
4.3 用 Opacus 一键加 DP
手写 DP-SGD 容易出错,生产用 Opacus(PyTorch 官方推荐):
from opacus import PrivacyEngine
model = MyModel()
optimizer = torch.optim.SGD(model.parameters(), lr=0.05)
privacy_engine = PrivacyEngine()
model, optimizer, loader = privacy_engine.make_private(
module=model,
optimizer=optimizer,
data_loader=loader,
noise_multiplier=1.1,
max_grad_norm=1.0,
)
for epoch in range(epochs):
for x, y in loader:
optimizer.zero_grad()
loss = nn.functional.cross_entropy(model(x), y)
loss.backward()
optimizer.step()
eps = privacy_engine.get_epsilon(delta=1e-5) # 实时查看已消耗的 ε
print(f"epoch {epoch}, ε = {eps:.2f}")
4.4 隐私预算与组合
ε 是累积的:训练步数越多、ε 消耗越多。典型目标:ε ≤ 8(较弱)、ε ≤ 3(中等)、ε ≤ 1(强)。要降低 ε,可以减小噪声倍数、减少步数、增大 batch。
| 参数 | 调大 | 影响 |
|---|---|---|
| noise_multiplier | 隐私↑ | 精度↓ |
| max_grad_norm | 精度↑ | 隐私↓(单样本影响大) |
| batch size | 隐私↑、精度↑ | 显存↑ |
| 训练步数 | 隐私↓ | 精度↑ |
一句话:DP-SGD 用「逐样本裁剪 + 高斯噪声」把隐私保证变成可计算的 ε;ε 随训练步数累积,调参就是在隐私预算和精度之间做取舍。
5. 安全聚合
5.1 服务器不该看到单个更新
即使加了 DP,服务器仍能看到每个客户端的梯度。安全聚合(Secure Aggregation) 保证服务器只能看到梯度的和,看不到任何单个客户端的更新——因为聚合才是训练真正需要的,单个梯度是多余的隐私暴露。
5.2 掩码的核心技巧
安全聚合用成对掩码抵消:客户端 i 和 j 协商一个随机掩码,i 加 +r、j 加 −r,求和时自动抵消,但服务器无法单独还原任何一个:
客户端1 上传:g1 + r12
客户端2 上传:g2 - r12
服务器求和 :(g1 + g2) + (r12 - r12) = g1 + g2 ← 掩码抵消,只见和
实际协议(如 Bonawitz 等人的方案)还要处理掉线(用 Shamir 秘密共享恢复掩码)和验证(防投毒)。
5.3 与其他隐私技术的组合
| 技术 | 保护对象 | 代价 |
|---|---|---|
| 安全聚合 | 单客户端更新 | 通信 2-3 倍、需多方协议 |
| 差分隐私 | 个体样本贡献 | 精度损失 |
| 同态加密 | 计算全程 | 计算开销巨大 |
| 可信执行环境 TEE | 硬件隔离 | 依赖硬件信任 |
DP + 安全聚合是当前最实用的组合:安全聚合防止服务器看到单个更新,DP 防止即使聚合结果也泄露个体信息。
一句话:安全聚合用成对掩码让服务器「只见和、不见单」,是联邦学习的加密底座;它防的是服务器窥探,与差分隐私(防个体反演)互补而非替代。
6. 横向、纵向与联邦迁移
6.1 三种协作模式
| 模式 | 数据划分 | 场景 | 例子 |
|---|---|---|---|
| 横向联邦 | 样本不同、特征相同 | 同类机构 | 多家医院 |
| 纵向联邦 | 样本相同、特征不同 | 异业合作 | 银行 + 电商 |
| 联邦迁移 | 样本、特征都不同 | 跨域 | 跨语言、跨场景 |
6.2 纵向联邦:样本对齐 + 拆分训练
纵向联邦最典型的场景是「银行有信用记录、电商有消费记录,同一批用户」——双方各自持有一部分特征,联合训练。
1. 隐私求交(PSI):双方在不暴露各自 ID 全集的前提下找到共同用户
2. 拆分建模:各持一部分特征的模型,中间层加密交互
3. 梯度加密传输:用同态加密或秘密共享保护中间结果
难点在于样本对齐(PSI)和加密下的梯度交互,工程复杂度远高于横向联邦。
6.3 联邦迁移
当各方数据连特征都对不齐时,用迁移学习搭桥:各方在本地学一个表征,通过共享的隐空间对齐(如 FedMD、基于对抗的对齐),实现知识迁移。
6.4 个性化联邦
全局模型对每个客户端未必最优。个性化联邦(pFL) 让客户端在全局模型基础上做本地微调,或把模型拆成「共享层 + 个性化层」,兼顾协作收益与个体适配。
# 个性化 FL:全局模型 + 本地微调
global_model = receive_global_model()
local_model = copy.deepcopy(global_model)
for epoch in range(finetune_epochs):
train_one_epoch(local_model, local_loader) # 只用本地数据微调
一句话:横向联邦解决「同样特征、不同样本」,纵向联邦解决「同样样本、不同特征」,联邦迁移解决「都不同」;越往后,样本对齐与加密交互的工程难度越高。
7. 隐私-效用权衡
7.1 权衡曲线
隐私越强(ε 越小、噪声越大),模型精度越低。这是一条不可避免的权衡曲线,目标不是消除它,而是找到业务可接受的点。
| ε 范围 | 隐私强度 | 精度损失(典型) | 适用 |
|---|---|---|---|
| ε > 10 | 弱 | < 1% | 合规底线低 |
| 3 < ε ≤ 10 | 中 | 1-3% | 多数业务 |
| 1 < ε ≤ 3 | 强 | 3-8% | 敏感数据 |
| ε ≤ 1 | 极强 | > 8% | 极敏感/研究 |
7.2 提升效用的技巧
1. 增大 batch:噪声摊薄到更多样本,效用↑
2. 预训练 + 微调:先用公开数据预训练,DP 微调步数少 → ε 小
3. 更小的 max_grad_norm + 更多步:有时比大裁剪更好
4. 参数高效微调(LoRA):只对少量参数加 DP,噪声影响小
5. 分组隐私:对不同敏感度的字段用不同 ε
7.3 评估隐私的真实性
警惕「伪隐私」:
- 只做联邦、不加 DP → 梯度可被反演,没有隐私保证;
- DP 但 ε 报得极小却精度无损 → 大概率实现有误或未正确累积;
- 只测聚合模型精度 → 要同时测**成员推断攻击(MIA)**的成功率来验证隐私。
一句话:隐私与效用是硬权衡,没有「免费」的隐私;提升效用的正道是「预训练 + 大 batch + 参数高效微调」,验证隐私的正道是跑成员推断攻击。
8. 框架选型与合规落地
8.1 主流框架
| 框架 | 出身 | 特点 |
|---|---|---|
| Flower | 开源社区 | 框架无关、易上手 |
| FedML | 学术+工业 | 算法全、支持多种场景 |
| FATE | 微众银行 | 工业级纵向联邦、合规 |
| TensorFlow Federated | TFF 研究向 | |
| NVIDIA FLARE | NVIDIA | 医疗影像、生产级 |
# Flower:极简的联邦客户端/服务器
import flwr as fl
class FlowerClient(fl.client.NumPyClient):
def get_parameters(self, config):
return [v.cpu().numpy() for v in model.state_dict().values()]
def fit(self, parameters, config):
set_parameters(model, parameters)
train(model, train_loader, epochs=1)
return self.get_parameters(config), len(train_loader.dataset), {}
def evaluate(self, parameters, config):
set_parameters(model, parameters)
loss, acc = test(model, test_loader)
return float(loss), len(test_loader.dataset), {"accuracy": float(acc)}
fl.client.start_numpy_client(server_address="127.0.0.1:8080", client=FlowerClient())
8.2 合规要点
1. 数据最小化:只传模型更新,不传原始数据
2. 明确告知:用户/机构知情并同意参与
3. 隐私预算记录:每次训练消耗的 ε 要可审计
4. 退出机制:参与方可随时退出,已贡献数据的影响可清除
5. 跨境合规:数据不出境,联邦天然契合
6. 安全评估:上线前做攻击测试(MIA、梯度反演)
8.3 落地中的现实问题
| 问题 | 现实 | 对策 |
|---|---|---|
| 参与方积极性 | 只有一方数据多 | 激励机制、收益分配 |
| 通信成本 | 跨机构带宽贵 | 压缩、异步、少轮次 |
| 掉线 | 移动端不稳定 | 容错聚合、超时剔除 |
| 异构 | 各方法律/算力不同 | 异步、分层聚合 |
| 模型知识产权 | 谁拥有全局模型 | 合同约定 |
一句话:联邦学习落地一半是技术、一半是协作机制——框架选 Flower/FATE,合规上抓「数据最小化 + 隐私预算可审计 + 退出机制」,技术上抓「通信压缩 + 容错聚合」。
9. 总结
9.1 技术栈全景
架构层:客户端本地训练 + 服务器聚合(FedAvg)
优化层:FedProx / SCAFFOLD 对抗非 IID 漂移
隐私层:DP-SGD(加噪)+ 安全聚合(掩码)
模式层:横向 / 纵向 / 迁移 / 个性化
工程层:Flower / FATE + 通信压缩 + 合规审计
9.2 关键决策点
| 问题 | 选择 |
|---|---|
| 同类机构、特征相同 | 横向联邦 + FedAvg |
| 异业合作、样本相同 | 纵向联邦 + PSI |
| 需要可证明隐私 | DP-SGD,ε 定在 3-8 |
| 防服务器窥探单更新 | 安全聚合 |
| 各方数据分布差异大 | FedProx / SCAFFOLD |
| 全局模型不够个性化 | 个性化 FL(本地微调) |
| 快速验证原型 | Flower |
9.3 一句话心法
联邦学习是「数据不出域」的架构,隐私保护是叠加在它上面的加固层——没有 DP 或安全聚合的联邦只是「数据不集中」,不是「隐私安全」;落地的胜负手在非 IID 优化与协作机制,而不在算法本身。
延伸阅读
- https://plumephp.com/ml-model-monitoring-drift/ — 非 IID 分布偏移的监控与检测
- https://plumephp.com/ml-feature-store/ — 特征对齐与纵向联邦的特征工程
- https://plumephp.com/ml-model-deployment/ — 模型上线、版本与灰度管理
- https://plumephp.com/ml-model-evaluation/ — 评估协议与隐私攻击测试设计
- AI/ML 专题 — 隐私计算与分布式系统深度文章
- Flower 官方文档
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。