关于联邦学习的几点记录
如果只能用一句话说联邦学习:先把失败复现出来。
先说场景
本来是打算搞个跨机构的预测模型,三家医院共同训练。每家都有病人数据,但隐私合规卡得很严,原始数据根本出不来。客户说"你们不是有联邦学习吗?弄一个试试"。
嘴上答应得很爽快,回家一查资料发现,市面上的教程要么是讲概念,要么是跑个官方 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 跟真实场景差得太远。官方数据集都是预处理好的,真实数据什么乱七八糟的情况都有。
真实数据怎么处理
真实数据一般长这样:
# 模拟真实场景的数据分布
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/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。