news 2026/9/18 0:34:36

Flower 联邦学习策略实战教程:从 FedAvg 到自定义策略(PyTorch)

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Flower 联邦学习策略实战教程:从 FedAvg 到自定义策略(PyTorch)

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:相关的参数摘要(etaeta_lbeta_1beta_2tau),这些摘要由FedOpt.summary()FedAvg.summary()输出。

本地运行

开发和调试阶段同样可以本地运行:

$ flwr run . local --stream

FedAdagrad / FedOpt 家族参数说明

FedAdagrad继承自FedOpt抽象类(见 fedopt.py),其构造函数支持的参数如下(均为关键字参数,默认值取自当前仓库源码):

参数默认值含义
fraction_train1.0参与训练的节点比例
fraction_evaluate1.0参与评估的节点比例
min_train_nodes2训练最少节点数(采样不足时兜底)
min_evaluate_nodes2评估最少节点数
min_available_nodes2系统中最少可用节点数
weighted_by_key"num-examples"加权聚合时使用的指标键
arrayrecord_key"arrays"消息中存放模型参数的键
configrecord_key"config"消息中存放配置的键
eta1e-1服务端学习率(聚合步长)
eta_l1e-1客户端学习率
tau1e-3控制算法自适应程度(防止除零)

此外,fedopt.py中还定义了beta_1/beta_2两个动量参数(FedAdagrad中固定为0.0),而FedAdamFedYogi等同类策略同样位于 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的执行步骤:

  1. 把聚合后的模型参数加载进一个 PyTorch 模型;
  2. 加载完整的 CIFAR-10 测试集;
  3. 在测试集上评估模型;
  4. 将评估指标封装成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 的DataLoadertest()则计算交叉熵损失与准确率(见 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_roundarraysconfiggrid),因此可以无缝覆写。父类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组装进MessageConfigRecord"config"这个键对应策略参数configrecord_key的默认值。因此,服务端对config["lr"]的任何修改,都会在下一轮被客户端感知到——这就是"服务端动态配置客户端执行"的完整闭环。

至此,你已创建出第一个自定义策略:它让发往客户端的ConfigRecord具备了动态行为。

小结:你学会了什么

本教程展示了如何以极少的代码增量逐步增强系统:

  1. 更换策略:一行导入、一行实例化,即可把FedAvg切换为FedAdagrad等自适应优化策略;
  2. 服务端评估:通过evaluate_fn回调,在每轮聚合后于服务端集中评估全局模型,并获得逐轮的ServerApp-side Evaluate Metrics
  3. 策略级配置下发:覆盖configure_train,在策略层实现学习率衰减等动态逻辑,配置经由Message中的ConfigRecord传送到客户端;
  4. 客户端消费配置ClientAppmsg.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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 0:33:21

Unity Shader变体优化:从原理到预加载,解决首帧卡顿与包体膨胀

做Unity客户端三年以上的人,几乎都会遇到一个现象:项目开发阶段一切正常,但第一次进入某个新场景时,画面会愣住几百毫秒甚至一两秒,然后才恢复正常。再严重一点,打了新包上真机,进入战斗首帧直接…

作者头像 李华
网站建设 2026/9/18 0:30:01

数字化康复评估:核心技术、应用与临床实践

1. 康复医疗的数字化变革契机传统康复治疗长期面临效果评估主观性强、数据分散难追溯的痛点。我在三甲医院康复科工作期间,最常听到患者问:"医生,我这个疗程到底进步了多少?"而治疗师往往只能给出"比上周好一点&qu…

作者头像 李华
网站建设 2026/9/18 0:28:21

机械原理PPT教案的拆解与二次开发:从静态课件到互动教学资源

简介:面向机械工程专业本科生、考研学生以及机械设计入门者,这份哈工大机械原理精品课程PPT教案,以76页完整课件系统讲解了机构运动分析与力学计算的核心知识点,尤其适合配合课堂同步复习或考前提纲式回顾。资源包内含1个pptx文件…

作者头像 李华
网站建设 2026/9/18 0:27:42

天线原理与理论基础:从辐射机理到微带天线匹配设计

简介:这份《天线原理天线理论基础》PPT学习教案面向通信工程、电子信息类专业学生及天线设计入门工程师,系统讲解天线理论的核心知识体系。资源为单个pptx课件,压缩包大小809KB,文件内容精炼,适合移动端或课堂教学快速…

作者头像 李华
网站建设 2026/9/18 0:27:27

小米版 Codex 安装配置与 CLI Agent 报错排查实战

小米版 Codex 这东西,我一开始是抱着看热闹的心态装的。结果第一天它就把我手上一个拖了两周的目录重构收尾了,第二天又帮我啃掉了一个老项目里最烦人的接口对齐活。装完之后我最大的感受是:codex 安装、codex 配置这些事本身不复杂&#xff…

作者头像 李华