联邦学习与隐私保护训练

联邦学习让数据不出域也能联合建模,代价是通信、异构与隐私三座大山。本文从系统视角讲解横向与纵向联邦的架构差异、梯度压缩与量化等通信优化、异步与半异步更新策略、差分隐私与安全聚合的组合方式、非独立同分布下的收敛缓解手段,以及用 Flower 落地的完整工程实践与排错清单。

数据不能出域,但模型想变聪明——联邦学习就是为这个矛盾设计的。它把训练拆成「本地算梯度、中心聚合参数」的循环,原始数据始终留在各参与方手里。听起来是个优雅的分布式训练特例,但真正落地时会发现:通信是瓶颈、客户端是异构的、数据是非独立同分布的、隐私还需要额外的数学保证。这四件事把联邦学习从「分布式 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 与个性化)、隐私(差分隐私与安全聚合的组合)。这三笔账必须一起算——只做通信压缩不做隐私保护,等于在明文管道上跑敏感梯度;只做隐私不做压缩,则受限于带宽而无法扩展到足够多的客户端。把它和 分布式训练 放在一起看,会发现两者共享同一套并行与通信直觉,只是联邦学习把「信任边界」加进了系统设计的第一性约束里。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. 排序学习与搜索召回排序系统
  2. 数据版本控制与血缘:DVC 与 LakeFS
  3. 模型可解释性:SHAP、LIME 与注意力归因