联邦学习概述与核心概念
联邦学习(Federated Learning)由 Google 于 2016 年提出,是一种分布式机器学习范式,允许多个参与方在不暴露原始数据的情况下协同训练全局模型。
核心思想
传统集中式训练需要收集所有数据到中心服务器,而联邦学习将训练过程推送到数据所在位置:
传统:数据 → 中心服务器 → 训练模型
联邦:模型 → 各参与方 → 本地训练 → 上传梯度/参数 → 聚合联邦学习分类
按数据分布划分
| 类型 | 特征重叠 | 样本重叠 | 典型场景 |
|---|---|---|---|
| 横向联邦 | ✓ | ✗ | 同行业不同用户群 |
| 纵向联邦 | ✗ | ✓ | 同用户不同特征 |
| 联邦迁移学习 | ✗ | ✗ | 跨领域迁移 |
按参与方规模划分
| 类型 | 参与方数量 | 网络条件 | 典型场景 |
|---|---|---|---|
| 跨设备 (Cross-Device) | 百万级 | 不稳定 | 移动端输入法 |
| 跨机构 (Cross-Silo) | 数十级 | 稳定 | 医院联合建模 |
基本流程
python
# 联邦学习基本流程伪代码
def federated_learning(server, clients, num_rounds):
global_model = server.initialize_model()
for round in range(num_rounds):
# 1. 服务端广播全局模型
server.broadcast(global_model)
# 2. 各客户端本地训练
local_updates = []
for client in clients:
local_model = client.receive(global_model)
update = client.local_train(local_model, num_epochs=5)
local_updates.append(update)
# 3. 服务端聚合更新
global_model = server.aggregate(local_updates)
return global_model关键挑战
| 挑战 | 描述 | 影响 |
|---|---|---|
| 数据异构性 | Non-IID 数据分布 | 收敛速度下降 |
| 通信效率 | 带宽受限 | 训练轮次增加 |
| 隐私保护 | 梯度可能泄露信息 | 需要额外机制 |
| 系统异构性 | 设备算力差异 | 延迟和掉线 |
| 恶意参与方 | 投毒攻击 | 模型质量下降 |
联邦学习适用条件
联邦学习适合以下场景:1) 数据无法集中(法规/商业限制);2) 数据量大但每方数据量小;3) 需要保护数据隐私;4) 参与方有训练能力。
联邦学习不是万能的
联邦学习的通信开销和收敛速度通常不如集中式训练。如果数据可以合法集中,集中式训练通常是更好的选择。