clustered-federated-learning
联邦学习 + 层次聚类:隐私保护下自动将客户端分组,训练专业化 AI 模型
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
联邦学习 + 层次聚类:隐私保护下自动将客户端分组,训练专业化 AI 模型
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象这样一个场景:一家医院和十家诊所想要联合训练一个医学影像诊断模型,但患者的隐私数据受到严格保护,任何直接的数据共享都触碰法律红线。与此同时,各诊所的数据分布差异巨大——有的诊所主治手写病历(EMNIST),有的则负责处理大量某类特殊字符。在这种情况下,标准联邦学习会让所有客户端被迫学到一个"平均"模型,效果往往不如各自单独训练。
Clustered Federated Learning(CFL) 正是为解决这一矛盾而生的学术框架。2019年,Felix Sattler、Klaus-Robert Müller 和 Wojciech Samek 三位来自柏林工业大学的学者在 arXiv 发表论文(arXiv:1910.01991),提出了一种将联邦学习与层次聚类相结合的方法——在保护隐私的前提下,自动将具有相似数据分布的客户端动态分组,让他们各自训练更适合自己的专业化模型。
本项目 felisat/clustered-federated-learning 是这篇论文的官方参考实现,基于 PyTorch 和 Jupyter Notebook 构建,为研究者和工程师提供了完整的算法复现与可视化验证。
联邦学习(Federated Learning)最早由 Google 在 2017 年提出,核心思想是"数据不动模型动":每个客户端在本地用私有数据训练模型,仅将模型权重更新(而非原始数据)上传到中央服务器,由服务器聚合后分发回各客户端。这种方式天然满足 GDPR 等隐私法规的要求,却带来了新的挑战——当客户端之间的数据分布差异显著(非独立同分布,Non-IID)时,聚合后的全局模型往往对所有客户端都"不够好"。
常见的应对策略包括个性化联邦学习、微调联邦学习等,而 CFL 选择了另一条路:在联邦学习过程中动态识别并聚类相似的客户端,让每个簇内的客户端独立聚合,从而实现"同类客户端协作,异类客户端分道扬镳"的效果。
CFL 的算法设计精巧,分为两个关键阶段:
服务器端持续监测所有客户端权重更新的平均范数(Mean Norm)——当平均范数低于阈值 ε₁ 时,表明联邦学习已经收敛,系统进入"聚类决策"阶段。这一判断背后有清晰的直觉:只有当全局模型稳定后,我们才能可靠地比较不同客户端的更新方向。
当平均范数 < ε₁ 且最大范数 > ε₂ 时,说明客户端之间出现了显著的更新方向分歧——某些客户端想要往不同的方向调整模型参数,此时 CFL 会触发聚类分裂:
项目中 fl_devices.py 的 pairwise_angles() 函数正是实现余弦相似度计算的核心逻辑:
def pairwise_angles(sources):
angles = torch.zeros([len(sources), len(sources)])
for i, source1 in enumerate(sources):
for j, source2 in enumerate(sources):
s1 = flatten(source1)
s2 = flatten(source2)
angles[i,j] = torch.sum(s1*s2)/(torch.norm(s1)*torch.norm(s2))
return angles
值得注意的是,分裂条件中还有 len(idc)>2 and c_round>20 的限制,确保每个簇至少有三个客户端、且经过足够的训练轮次后才做决策,避免过早分裂导致的振荡。
项目代码规模适中(约 130KB 的 Notebook + 5 个辅助 Python 模块),结构清晰:
| 文件 | 职责 |
|---|---|
clustered_federated_learning.ipynb | 主演示 Notebook(15 个 cell),包含完整实验流程 |
fl_devices.py | 联邦学习核心逻辑:Server 类(聚合/聚类)、Client 类(本地训练) |
models.py | 模型定义:ConvNet(两层卷积 + 池化 + 全连接) |
data_utils.py | 数据划分:split_noniid() 实现 Dirichlet 非IID分布划分 |
helper.py | 可视化工具:训练曲线绑制、聚类分裂时机标注 |
在真实场景中,不同医院/诊所的数据分布天然不同。CFL 论文采用了 Dirichlet 分布来模拟这种异质性。data_utils.py 中的 split_noniid() 函数将 MNIST 字符数据按 Dirichlet(α=1.0) 分配给 N 个客户端:α 越小,各客户端的数据分布越倾斜;当 α→∞ 时趋近于 IID 分布。
Notebook 中的实验设计了一个巧妙的异构场景:前 5 个客户端的数据全部旋转 180°,模拟"数据标注风格差异"的实际挑战。如果使用标准联邦学习,所有客户端被迫聚合到一个统一模型,效果会很差;而 CFL 则能自动检测到这 5+5 的分组,在第 20 轮左右触发分裂,分别收敛到两个专业化模型。
class ConvNet(torch.nn.Module):
def __init__(self):
super(ConvNet, self).__init__()
self.conv1 = torch.nn.Conv2d(1, 6, 5) # 1通道 → 6通道
self.pool = torch.nn.MaxPool2d(2, 2) # 2x2 最大池化
self.conv2 = torch.nn.Conv2d(6, 16, 5) # 6通道 → 16通道
self.fc1 = torch.nn.Linear(16 * 4 * 4, 62) # 输出62类(数字+大小写字母)
这是一个轻量级的 CNN,专为 EMNIST(扩展版 MNIST,62 类)设计。模型参数量小,适合快速实验验证算法有效性。
Notebook 的执行流程高度模块化,主要步骤如下:
关键超参数:
COMMUNICATION_ROUNDS = 80:联邦学习通信轮次EPS_1 = 0.4:平均范数阈值,控制何时进入聚类决策EPS_2 = 1.6:最大范数阈值,控制何时触发分裂alpha = 1.0:Dirichlet 分布参数,控制数据异构程度CFL 的最大创新在于无需提前知道客户端的分组结构,系统通过监控梯度更新方向自动发现并适应聚类。这与需要预先指定簇数量的传统聚类方法形成鲜明对比。在论文的实验中,CFL 在 CIFAR-10 和 FEMNIST 数据集上均显著优于标准 FedAvg,尤其在高度异构场景下优势明显。
作为一个研究原型,项目存在以下局限:
CFL 论文发表后,联邦聚类学习逐渐成为 FL 研究的一个重要分支。后续工作如 IFCA(各客户端同时属于多个簇)、FedRep(个性化表示学习)等在 CFL 的基础上进一步深化了个性化联邦学习的探索。
从技术演进角度看,CFL 代表了联邦学习从"一刀切的全局模型"向"动态自适应的专业化模型"过渡的关键节点。在医疗、金融、边缘计算等数据隐私敏感领域,这种"既协作又个性"的范式正在获得越来越广泛的应用。
本项目作为 CFL 的参考实现,代码简洁、注释充足,非常适合作为联邦学习入门研究和算法改进的 baseline。