Flower 联邦学习策略实战教程:从 FedAvg 到自定义策略(PyTorch)
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
本教程是 Flower 协作式 AI 教程系列的第 4 部分,基于@flwrlabs/quickstart-pytorch示例应用,逐步演示如何在 PyTorch 联邦学习应用中完成四件关键事情:将默认的FedAvg策略切换为FedAdagrad、在服务端对聚合后的全局模型做集中式评估、通过覆盖configure_train方法向客户端下发动态配置(学习率衰减),以及理解策略与ServerApp之间的完整调用链。读完本文,你将掌握 Flower 策略抽象的核心接口、Strategy.start()的执行流程,以及"用最少的代码定制联邦学习算法"的完整路径。
本文对应的原文档为 tutorial-series-use-a-federated-learning-strategy-pytorch.rst,正文中的源码分析均基于当前仓库
framework/py/flwr/serverapp/strategy/下的真实实现。
前置准备
本教程承接系列前一篇《编写你的第一个 Flower 应用(PyTorch)》。如果你已完成上一篇,直接打开已有的quickstart-pytorch目录继续即可;如果你是直接从本部分开始,请安装 Flower 并从 Flower Hub 拉取同样的应用:
# 安装 Flower $ pip install -U flwr # 从 Flower Hub 获取应用 $ flwr new @flwrlabs/quickstart-pytorch # 进入应用目录 $ cd quickstart-pytorch应用目录中的核心文件位于当前仓库的 examples/quickstart-pytorch/pytorchexample/:
server_app.py:定义ServerApp及入口main(),负责初始化策略并驱动整个联邦学习循环;client_app.py:定义ClientApp,实现@app.train()与@app.evaluate()两个入口;task.py:包含模型Net、数据加载(load_data/load_centralized_dataset)以及train/test函数。
策略是什么:联邦学习算法的"心脏"
在动手改代码之前,先明确一个概念:策略(Strategy)封装了联邦学习的具体算法。它决定每一轮联邦学习中:
- 选取哪些客户端节点参与训练 / 评估(采样逻辑);
- 向客户端发送什么模型参数和配置(
configure_train/configure_evaluate); - 如何聚合客户端返回的结果(
aggregate_train/aggregate_evaluate); - 聚合后的全局参数如何更新(例如
FedAvg的加权平均,或FedAdagrad的自适应更新)。
在 Flower 1.x 的消息式架构中,所有策略都继承自Strategy抽象基类(见 framework/py/flwr/serverapp/strategy/strategy.py),它声明了四个必须实现的方法:
| 方法 | 签名要点 | 作用 |
|---|---|---|
configure_train | (server_round, arrays, config, grid) -> Iterable[Message] | 生成发往客户端节点的训练消息 |
aggregate_train | (server_round, replies) -> (ArrayRecord, MetricRecord) | 聚合训练返回的模型参数与指标 |
configure_evaluate | (server_round, arrays, config, grid) -> Iterable[Message] | 生成发往客户端节点的评估消息 |
aggregate_evaluate | (server_round, replies) -> MetricRecord | 聚合客户端评估指标 |
Strategy.start()(见 strategy.py)则是驱动整个训练循环的引擎:它在第 0 轮先调用一次evaluate_fn评估初始参数,然后逐轮执行"配置训练 →grid.send_and_receive收发消息 → 聚合训练结果 → 配置评估 → 聚合评估指标 → 服务端集中评估"的完整流程,最后返回一个包含最终模型与各轮指标的Result。
当前quickstart-pytorch应用的server_app.py默认使用FedAvg(对应论文 McMahan et al., 2017, "Communication-Efficient Learning of Deep Networks from Decentralized Data",实现见 fedavg.py)。现在,让我们尝试更换一种策略。
更换策略:从 FedAvg 切换到 FedAdagrad
修改 server_app.py
FedAdagrad是自适应联邦优化(Adaptive Federated Optimization)家族的一员,源自 Reddi et al., 2020 的论文(实现见 fedadagrad.py)。与FedAvg每次简单做加权平均不同,FedAdagrad在服务端维护一个二阶矩累积量v_t,用类似 Adagrad 的方式逐参数调整聚合步长:
self.v_t = {k: v + (delta_t[k] ** 2) for k, v in self.v_t.items()} new_arrays = { k: x + self.eta * m_t[k] / (np.sqrt(self.v_t[k]) + self.tau) for k, x in self.current_arrays.items() }修改server_app.py中的以下两处,把FedAvg换成FedAdagrad:
# ... 其余代码不变 # 在导入部分添加这一行 from flwr.serverapp.strategy import FedAdagrad # ... 其余代码不变 @app.main() def main(grid: Grid, context: Context) -> None: """ServerApp 主入口。""" # 读取运行配置 fraction_evaluate: float = context.run_config["fraction-evaluate"] num_rounds: int = context.run_config["num-server-rounds"] lr: float = context.run_config["learning-rate"] # 加载全局模型 global_model = Net() arrays = ArrayRecord(global_model.state_dict()) # 初始化 FedAdagrad 策略 strategy = FedAdagrad(fraction_evaluate=fraction_evaluate) # 启动策略,运行 FedAdagrad 共 num_rounds 轮 result = strategy.start( grid=grid, initial_arrays=arrays, train_config=ConfigRecord({"lr": lr}), num_rounds=num_rounds, evaluate_fn=global_evaluate, )在 SuperGrid 上运行
登录 SuperGrid 并运行应用,确认新策略已生效:
# 如果尚未登录,先登录 $ flwr login supergrid # 在 SuperGrid 上运行应用 $ flwr run . supergrid打开 SuperGrid 控制台,选中你的联邦,查看最新一次运行的日志。你应该能看到 Flower 启动的是FedAdagrad策略而非FedAvg——具体表现为日志中Starting FedAdagrad strategy:以及FedOpt settings:相关的参数摘要(eta、eta_l、beta_1、beta_2、tau),这些摘要由FedOpt.summary()与FedAvg.summary()输出。
本地运行
开发和调试阶段同样可以本地运行:
$ flwr run . local --streamFedAdagrad / FedOpt 家族参数说明
FedAdagrad继承自FedOpt抽象类(见 fedopt.py),其构造函数支持的参数如下(均为关键字参数,默认值取自当前仓库源码):
| 参数 | 默认值 | 含义 |
|---|---|---|
fraction_train | 1.0 | 参与训练的节点比例 |
fraction_evaluate | 1.0 | 参与评估的节点比例 |
min_train_nodes | 2 | 训练最少节点数(采样不足时兜底) |
min_evaluate_nodes | 2 | 评估最少节点数 |
min_available_nodes | 2 | 系统中最少可用节点数 |
weighted_by_key | "num-examples" | 加权聚合时使用的指标键 |
arrayrecord_key | "arrays" | 消息中存放模型参数的键 |
configrecord_key | "config" | 消息中存放配置的键 |
eta | 1e-1 | 服务端学习率(聚合步长) |
eta_l | 1e-1 | 客户端学习率 |
tau | 1e-3 | 控制算法自适应程度(防止除零) |
此外,fedopt.py中还定义了beta_1/beta_2两个动量参数(FedAdagrad中固定为0.0),而FedAdam、FedYogi等同类策略同样位于 strategy 目录下,你可以在需要时按同样方式替换。
服务端参数评估:Centralized Evaluation 实战
Flower 支持在服务端或客户端两侧评估模型,两者各有优劣。
集中式评估(Centralized Evaluation,即服务端评估)概念上最简单:它和传统集中式机器学习中的评估完全一致。只要服务端有可用的评估数据集,就可以在每轮训练结束后直接评估新聚合出的全局模型,无需把模型再发回客户端。而且整个评估数据集随时完整可用,结果稳定。
联邦式评估(Federated Evaluation,即客户端评估)更复杂,但也更强大:它不要求存在集中式数据集,可以在更大范围的数据上评估模型,往往能得到更接近真实场景的结果。实际上,很多场景必须依赖联邦式评估才能得到有代表性的结果。但它的代价是:如果参与评估的客户端并非每轮都在线,评估数据集会在连续轮次间变化;单个客户端持有的数据也可能逐轮变化。这会导致评估结果不稳定——即使模型完全不变,连续轮次的评估指标也会波动。
你已经在ClientApp中体验过客户端侧评估(通过@app.evaluate装饰器,见 client_app.py)。现在来看如何在服务端评估聚合后的模型参数。
server_app.py中定义的global_evaluate函数正是为此服务的。它作为回调传入策略的start方法,策略会在每一轮联邦学习结束后调用它,并传入两个参数:当前轮数server_round与聚合后的模型参数arrays。从 strategy.py 的源码可以看出,evaluate_fn会在第 0 轮(初始参数)以及每一轮结束后各被调用一次。
global_evaluate的执行步骤:
- 把聚合后的模型参数加载进一个 PyTorch 模型;
- 加载完整的 CIFAR-10 测试集;
- 在测试集上评估模型;
- 将评估指标封装成
MetricRecord返回。
from flwr.app import ArrayRecord, MetricRecord def global_evaluate(server_round: int, arrays: ArrayRecord) -> MetricRecord: """在中心数据上评估模型。""" # 加载模型并用收到的参数初始化 model = Net() model.load_state_dict(arrays.to_torch_state_dict()) device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") model.to(device) # 加载完整测试集 test_dataloader = load_centralized_dataset() # 在测试集上评估全局模型 test_loss, test_acc = test(model, test_dataloader, device) # 返回评估指标 return MetricRecord({"accuracy": test_acc, "loss": test_loss})注意global_evaluate依赖task.py中两个关键函数:load_centralized_dataset()一次性加载uoft-cs/cifar10数据集的完整 test split 并返回批量大小为 128 的DataLoader;test()则计算交叉熵损失与准确率(见 task.py)。
要将该回调接入策略,只需在strategy.start(...)中通过evaluate_fn=global_evaluate传入。quickstart 应用已经这么做了,所以切换到FedAdagrad后务必保留这一行:
@app.main() def main(grid: Grid, context: Context) -> None: """ServerApp 主入口。""" # ... 其余代码不变 # 启动策略,运行 FedAdagrad 共 num_rounds 轮 result = strategy.start( grid=grid, initial_arrays=arrays, train_config=ConfigRecord({"lr": lr}), num_rounds=num_rounds, evaluate_fn=global_evaluate, ) # ... 其余代码不变后续为了快速迭代,我们改用本地模拟运行:
$ flwr run . local --stream运行结束后,服务端会打印每轮回调返回的指标。注意日志末尾的ServerApp-side Evaluate Metrics:
INFO : ServerApp-side Evaluate Metrics: INFO : { 0: {'accuracy': '1.0000e-01', 'loss': '2.3053e+00'}, INFO : 1: {'accuracy': '1.0000e-01', 'loss': '2.3203e+00'}, INFO : 2: {'accuracy': '2.3230e-01', 'loss': '2.0144e+00'}, INFO : 3: {'accuracy': '2.5720e-01', 'loss': '1.9258e+00'}}其中第 0 轮对应初始(未训练)模型在完整测试集上的表现,这正是 strategy.py 中"evaluate starting global parameters"逻辑的输出;这些指标最终被存入Result.evaluate_metrics_serverapp字典,键为轮数。
从策略向客户端发送配置:让学习率按轮次动态衰减
有些场景需要由服务端来配置客户端的训练 / 评估行为。典型例子:服务端希望根据当前轮数让客户端使用不同的学习率。Flower 允许把配置值作为Message的一部分从服务端发往客户端,ClientApp收到后即可读取使用。
理解 ConfigRecord 的传递路径
我们在调用strategy.start(...)时已经传入了一个ConfigRecord({"lr": lr})。这个ConfigRecord会出现在所有发往@app.train()的Message中。从 fedavg.py 的源码可以看到configure_train的实现细节:
- 根据
fraction_train计算并采样参与训练的节点(sample_nodes); - 无条件注入当前轮数:
config["server-round"] = server_round; - 把
ArrayRecord(模型参数)与ConfigRecord(配置)组装进RecordDict,为每个节点构造一条TRAIN类型的Message。
也就是说,train_config中你放入的任何键值对,都会被原样带到客户端。
现在假设我们要实现:每 5 轮把学习率乘以 0.5。这需要覆盖策略的configure_train方法并嵌入这段逻辑。
编写 custom_strategy.py
在pytorchexample目录下新建custom_strategy.py,添加以下代码——它继承自FedAdagrad,只覆盖configure_train一个方法:
from typing import Iterable from flwr.serverapp import Grid from flwr.serverapp.strategy import FedAdagrad from flwr.app import ArrayRecord, ConfigRecord, Message class CustomFedAdagrad(FedAdagrad): def configure_train( self, server_round: int, arrays: ArrayRecord, config: ConfigRecord, grid: Grid ) -> Iterable[Message]: """配置下一轮联邦训练,并实现学习率衰减。""" # 每 5 轮将学习率乘以 0.5 if server_round % 5 == 0 and server_round > 0: config["lr"] *= 0.5 print("LR decreased to:", config["lr"]) # 将更新后的配置与其余参数交给父类处理 return super().configure_train(server_round, arrays, config, grid)关键点在于:configure_train的签名与Strategy抽象基类完全一致(server_round、arrays、config、grid),因此可以无缝覆写。父类FedAvg.configure_train会把传入的config(已被我们修改过)组装进Message发送给客户端。
在 ServerApp 中使用自定义策略
回到server_app.py,导入CustomFedAdagrad并替换原先的FedAdagrad:
# ... 其余代码不变 # 在导入部分添加这一行 from pytorchexample.custom_strategy import CustomFedAdagrad # ... 其余代码不变 @app.main() def main(grid: Grid, context: Context) -> None: """ServerApp 主入口。""" # ... 其余代码不变 # 初始化自定义 FedAdagrad 策略 strategy = CustomFedAdagrad(fraction_evaluate=fraction_evaluate) # ... 其余代码不变再次本地运行,这次把轮数提高到 15 以观察学习率衰减效果:
$ flwr run . local --stream --run-config="num-server-rounds=15"在configure_train阶段(第 5 轮和第 10 轮),学习率会被乘以 0.5,新的学习率会打印到终端。
客户端如何消费新配置
如何确认ClientApp确实在用新的学习率?回顾client_app.py中@app.train()的实现,它从接收到的Message里读取学习率:
@app.train() def train(msg: Message, context: Context): # ... 初始化 # 调用训练函数 train_loss = train_fn( model, trainloader, context.run_config["local-epochs"], msg.content["config"]["lr"], device, ) # ... 构造回复 Message return Message(content=content, reply_to=msg)这里msg.content["config"]正是服务端configure_train组装进Message的ConfigRecord;"config"这个键对应策略参数configrecord_key的默认值。因此,服务端对config["lr"]的任何修改,都会在下一轮被客户端感知到——这就是"服务端动态配置客户端执行"的完整闭环。
至此,你已创建出第一个自定义策略:它让发往客户端的ConfigRecord具备了动态行为。
小结:你学会了什么
本教程展示了如何以极少的代码增量逐步增强系统:
- 更换策略:一行导入、一行实例化,即可把
FedAvg切换为FedAdagrad等自适应优化策略; - 服务端评估:通过
evaluate_fn回调,在每轮聚合后于服务端集中评估全局模型,并获得逐轮的ServerApp-side Evaluate Metrics; - 策略级配置下发:覆盖
configure_train,在策略层实现学习率衰减等动态逻辑,配置经由Message中的ConfigRecord传送到客户端; - 客户端消费配置:
ClientApp从msg.content["config"]读取服务端下发的任何配置项。
这些能力的基础是Strategy抽象(strategy.py)与其内置实现(fedavg.py、fedopt.py、fedadagrad.py),加上start()中"配置 → 收发消息 → 聚合 → 集中评估"的固定循环。理解这条链路后,你可以进一步定制客户端执行逻辑,甚至构建更大规模的联邦学习模拟。
下一步
下一篇教程《从零构建一个策略》将演示如何不依赖内置实现,从头编写一个完全自定义的Strategy,进一步释放 Flower 策略抽象的灵活性。
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考