关于联邦学习的几点记录

如果只能用一句话说联邦学习:先把失败复现出来。

先说场景

本来是打算搞个跨机构的预测模型,三家医院共同训练。每家都有病人数据,但隐私合规卡得很严,原始数据根本出不来。客户说"你们不是有联邦学习吗?弄一个试试"。

嘴上答应得很爽快,回家一查资料发现,市面上的教程要么是讲概念,要么是跑个官方 demo 就完事。真正要落地的时候,这些问题全都没人提:

  • 节点网络环境不一样,有的在内网,有的在云上,怎么通
  • 训练数据分布不均,有的类别只有几十个样本,怎么平衡
  • 模型更新传输,安全性怎么保证
  • 恶意节点投毒,怎么检测和防御

说干就干,先搭个环境试试。

环境搭建

# Python 3.10 环境
conda create -n federated python=3.10 -y
conda activate federated

# TensorFlow 和 TFF
pip install tensorflow==2.12.0
pip install tensorflow-federated==0.46.0

# 辅助库
pip install numpy==1.24.3
pip install matplotlib==3.7.1
pip install pyopenssl==23.1.1

第一个坑就来了。TFF 对 TensorFlow 版本特别敏感,官方说支持 2.9 到 2.12,但实际跑起来只有 2.12 稳定。2.9 经常报奇怪的错误,2.13 直接不兼容。

血的教训: 安装前先看 TFF 官方的兼容性矩阵,不要想当然用最新版 TensorFlow。

最简 Demo 跑起来

先跑个 MNIST 的官方 demo 确保环境没问题:

import tensorflow_federated as tff
import tensorflow as tf
import numpy as np

# 加载 MNIST 数据
emnist_train, emnist_test = tff.simulation.datasets.emnist.load_data()

def create_keras_model():
    model = tf.keras.models.Sequential([
        tf.keras.layers.Flatten(input_shape=(28, 28, 1)),
        tf.keras.layers.Dense(128, activation='relu'),
        tf.keras.layers.Dense(10, activation='softmax')
    ])
    return model

def model_fn():
    keras_model = create_keras_model()
    return tff.learning.from_keras_model(
        keras_model,
        input_spec=emnist_train.element_spec,
        loss=tf.keras.losses.SparseCategoricalCrossentropy(),
        metrics=[tf.keras.metrics.SparseCategoricalAccuracy()]
    )

# 创建联邦平均算法
iterative_process = tff.learning.build_federated_averaging_process(
    model_fn,
    client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.02),
    server_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=1.0)
)

# 初始化
state = iterative_process.initialize()

# 选择客户端
def make_federated_data(client_data, client_ids):
    return [client_data.create_tf_dataset_for_client(x) for x in client_ids]

client_ids = sorted(emnist_train.client_ids)[:10]
federated_train_data = make_federated_data(emnist_train, client_ids)

# 训练一轮
state, metrics = iterative_process.next(state, federated_train_data)
print(f'Round 1: {metrics}')

跑起来是跑起来了,但这个 demo 跟真实场景差得太远。官方数据集都是预处理好的,真实数据什么乱七八糟的情况都有。

graph TD A[服务器初始化模型] --> B[发送模型到客户端] B --> C[客户端本地训练] C --> D[返回模型更新] D --> E[服务器聚合更新] E --> F[重复训练] F -->|达到收敛条件| G[训练结束]

真实数据怎么处理

真实数据一般长这样:

# 模拟真实场景的数据分布
def create_client_datasets(num_clients=3, samples_per_client=(100, 5000)):
    datasets = []
    for i in range(num_clients):
        n_samples = np.random.randint(*samples_per_client)

        # 这里模拟数据分布不均
        # 有的客户端某类样本特别少
        class_dist = np.random.dirichlet([0.1, 0.1, 0.1, 0.1, 0.1,
                                         0.1, 0.1, 0.1, 0.1, 0.1]) * n_samples
        class_dist = class_dist.astype(int)

        X = np.random.randn(n_samples, 28, 28, 1).astype(np.float32)
        y = np.repeat(np.arange(10), class_dist)[:n_samples]

        dataset = tf.data.Dataset.from_tensor_slices((X, y))
        dataset = dataset.batch(32).prefetch(tf.data.AUTOTUNE)
        datasets.append(dataset)

    return datasets

第二个坑来了: 数据分布不均导致模型在某些类别上表现很差。有个客户端只有 3 个样本属于某个类别,训练出来的更新基本就是噪声。

解决方案之一是调整损失函数,给样本少的类别更高权重:

# 使用加权交叉熵
loss_fn = tf.keras.losses.SparseCategoricalCrossentropy()

def compute_weighted_loss(y_true, y_pred, sample_weights):
    return loss_fn(y_true, y_pred, sample_weight=sample_weights)

但这个方案只能缓解问题,根本解决还是要重新采样或者做数据增强。联邦学习里做增强比较麻烦,因为要改客户端的数据处理逻辑。

网络通信怎么搞

真实场景下,服务器和客户端可能不在同一个网络里。内网的客户端怎么跟云上的服务器通信?

TFF 官方主要提供仿真环境,但支持自定义执行上下文:

# 自定义通信层(简化示例)
class CustomExecutor:
    def __init__(self, server_address):
        self.server_address = server_address

    def send_update(self, client_id, weights):
        # 实际应该用 gRPC 或 HTTPS
        response = requests.post(
            f'{self.server_address}/update',
            json={
                'client_id': client_id,
                'weights': weights.tolist()
            },
            verify=False  # 开发环境临时禁用证书验证
        )
        return response.json()

    def get_global_model(self):
        response = requests.get(
            f'{self.server_address}/model',
            verify=False
        )
        return np.array(response.json()['weights'])

第三个坑: 证书验证问题。内网环境里经常用自签名证书,HTTPS 连接会报错。生产环境肯定要解决证书问题,但开发和测试阶段可以先临时禁用(仅限内网)。

隐私保护不是万能药

很多人以为联邦学习就等于隐私保护,这是个误区。

联邦学习只是减少了数据传输,但模型更新本身仍然可能泄露信息。攻击者可以通过分析多次更新来反向推导训练数据。

简单的保护措施:

# 添加差分隐私噪声
def add_dp_noise(weights, noise_multiplier=0.1, clip_norm=1.0):
    # 裁剪梯度
    clipped_weights = [tf.clip_by_norm(w, clip_norm) for w in weights]

    # 添加噪声
    noisy_weights = []
    for w in clipped_weights:
        noise = tf.random.normal(
            tf.shape(w),
            stddev=noise_multiplier * clip_norm
        )
        noisy_weights.append(w + noise)

    return noisy_weights

更高级的方案是使用同态加密,但计算开销会显著增加:

# 简化示例,实际应该用成熟的加密库
from cryptography.hazmat.primitives.asymmetric import padding
from cryptography.hazmat.primitives import hashes

def encrypt_weights(public_key, weights):
    # 序列化权重
    weights_bytes = tf.io.serialize_tensor(weights).numpy()

    # 加密
    encrypted = public_key.encrypt(
        weights_bytes,
        padding.OAEP(
            mgf=padding.MGF1(algorithm=hashes.SHA256()),
            algorithm=hashes.SHA256(),
            label=None
        )
    )
    return encrypted

恶意节点怎么办

真实世界里总有人搞事情。恶意节点可能:

  • 发送假的数据集
  • 在本地训练时投毒
  • 返回错误的模型更新

检测恶意节点的方法之一是对比更新异常值:

def detect_malicious_updates(updates, threshold=2.0):
    # 计算每个更新的 L2 范数
    norms = [tf.norm(update).numpy() for update in updates]
    mean_norm = np.mean(norms)
    std_norm = np.std(norms)

    # 标记异常值
    anomalies = []
    for i, norm in enumerate(norms):
        z_score = (norm - mean_norm) / std_norm
        if abs(z_score) > threshold:
            anomalies.append(i)

    return anomalies

更复杂的方案是用声誉系统,根据历史表现给每个客户端打分,然后加权聚合:

def weighted_aggregation(updates, reputation_scores):
    # 归一化声誉分数
    total_reputation = sum(reputation_scores)
    weights = [score / total_reputation for score in reputation_scores]

    # 加权聚合
    aggregated = []
    for param_idx in range(len(updates[0])):
        weighted_params = []
        for client_idx, update in enumerate(updates):
            weighted_params.append(update[param_idx] * weights[client_idx])
        aggregated.append(tf.reduce_sum(weighted_params, axis=0))

    return aggregated

实际踩过的坑

网络超时

一开始没设超时,某个客户端网络不稳定,整个训练流程卡住了。

# 设置超时
response = requests.post(
    url,
    timeout=30.0  # 30 秒超时
)

内存泄漏

多次训练后,TensorFlow 会积攒计算图,内存不断涨。

# 定期清理会话
tf.keras.backend.clear_session()

版本兼容

服务器和客户端用的 TensorFlow 版本不一致,模型更新序列化/反序列化失败。

# 统一版本号,写入 requirements.txt
cat > requirements.txt << EOF
tensorflow==2.12.0
tensorflow-federated==0.46.0
numpy==1.24.3
EOF

性能优化的几点体会

  • 批量大小要调小:联邦学习里每个客户端的数据量有限,批量太大反而影响收敛
  • 学习率要比普通训练低:因为更新频率更高,容易震荡
  • 客户端选择策略要灵活:不用每次都选全部客户端,随机选子集也能收敛得不错
  • 压缩传输数据:模型更新用 fp16 而不是 fp32,能减少一半传输量
# 使用混合精度训练
tf.keras.mixed_precision.set_global_policy('mixed_float16')

# 或者手动转换
def compress_weights(weights):
    return [tf.cast(w, tf.float16) for w in weights]

def decompress_weights(weights):
    return [tf.cast(w, tf.float32) for w in weights]

什么时候该用联邦学习

折腾了几个月,有些感悟。

联邦学习不是万能钥匙。这些场景可以考虑:

  • 数据敏感,不能离开本地
  • 客户端数据量够大,本地训练有意义
  • 网络条件尚可,能支撑频繁通信
  • 对实时性要求不高,训练周期可以拉长

这些场景不推荐:

  • 数据本身就很少,本地训练没什么意义
  • 客户端计算资源有限,跑不动模型
  • 网络不稳定,经常断连
  • 对模型精度要求极高,联邦学习的精度损失不可接受

收尾

联邦学习解决了一个具体问题:数据不能用聚合的方式训练模型。但它引入了新的复杂度和限制。

技术上没有银弹,每个方案都有适用场景。联邦学习不一定是最好的选择,但在某些约束条件下,它可能是唯一的选择。

最后,如果真要上生产,建议先从简单方案开始:本地训练、定期同步、人工审核。等验证了价值,再上自动化的联邦学习。落地比概念重要。


[参考链接]

  • TensorFlow Federated 官方文档
  • “Communication-Efficient Learning of Deep Networks from Decentralized Data” (McMahan et al., 2016)
  • “Differentially Private Federated Learning: A Client Level Perspective” (Agarwal et al., 2018)

版权声明: 本文首发于 指尖魔法屋-关于联邦学习的几点记录https://blog.thinkmoon.cn/post/179-federated-learning-deep-dive-privacy-distributed-training/) 转载或引用必须申明原指尖魔法屋来源及源地址!