ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从零自定义 Flower Strategy:重写 `start` 方法实现模型存档与 WB 指标追踪(PyTorch 实战)

从零自定义 Flower Strategy:重写 `start` 方法实现模型存档与 WB 指标追踪(PyTorch 实战) 从零自定义 Flower Strategy重写start方法实现模型存档与 WB 指标追踪PyTorch 实战【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower导读本篇教程基于 Flower 框架的官方教程系列tutorial-series-build-a-strategy-from-scratch-pytorch.rst围绕「自定义Strategy的start方法」展开。你将继承上一篇教程中已经学会的configure_train定制思路在FedAdagrad基础上实现一个CustomFedAdagrad每当集中式评估发现新的全局最优准确率时自动保存模型 checkpoint并把每轮训练/评估指标实时写入 Weights BiasesWB。读完本文你将掌握 Flower 消息制策略的执行骨架、start的完整生命周期以及如何在server_app.py中注入自定义路径配置。1. 教程定位与前置准备本教程是 Flower Collaborative AI 教程系列中的一环。此前你已经完成了在 SuperGrid 上创建模拟联邦federation运行并自定义 Flower App从 NumPy demo 迁移到 PyTorch quickstart app在上一篇教程中通过重写configure_train完成了学习率衰减并把更新后的学习率放进ConfigRecord随Message发送给客户端。本教程在此基础上更进一步直接重写策略的主入口start方法。最终的CustomFedAdagrad需要实现三个目标保存全局模型副本当发现新的全局最优准确率时把当前全局模型存到磁盘记录运行指标把每轮产生的训练、评估指标实时写入 WB沿用上一篇的configure_train定制每 5 轮将学习率衰减为原来的一半。1.1 安装依赖与创建 App如果从本教程直接开始跳过上一篇需要先创建 PyTorch quickstart App# 安装 Flower $ pip install -U flwr # 用 PyTorch quickstart 模板创建新 Flower App $ flwr new flwrlabs/quickstart-pytorch进入 App 目录后为pyproject.toml的依赖列表追加wandb$ cd quickstart-pytorchwandb0.17.8Flower 会在 App 运行时把该依赖安装到隔离的运行时环境中具体机制参见仓库文档 how-to-install-app-dependencies-at-runtime。注意如果你是第一次安装wandb运行前可能需要在当前终端环境中先执行wandb login完成账号注册与登录$ wandb login2. 理解start方法策略的主循环Flower 的Strategy抽象基类定义在 framework/py/flwr/serverapp/strategy/strategy.py它把「可被定制」的部分抽象成一组方法configure_train/aggregate_train配置并聚合训练轮次configure_evaluate/aggregate_evaluate配置并聚合客户端侧评估summary打印策略配置摘要start策略的总入口包含整个联邦学习流程的驱动循环strategy.py。从 strategy.py 的源码签名可以看到start的完整参数def start( self, grid: Grid, initial_arrays: ArrayRecord, num_rounds: int 3, timeout: float 3600, train_config: ConfigRecord | None None, evaluate_config: ConfigRecord | None None, evaluate_fn: Callable[[int, ArrayRecord], MetricRecord | None] | None None, ) - Result:各参数含义如下参数类型默认值说明gridGrid—用于向执行ClientApp的节点发送/接收Message的网格实例initial_arraysArrayRecord—初始模型参数数组联邦学习的起点num_roundsint3要执行的联邦学习轮数timeoutfloat3600等待节点响应的超时时间秒train_configConfigRecordNone训练轮次下发给节点的配置未设置时使用空ConfigRecordevaluate_configConfigRecordNone评估轮次下发给节点的配置未设置时使用空ConfigRecordevaluate_fnCallableNone服务端集中式评估函数接收轮号与ArrayRecord返回MetricRecord提供时会在第 0 轮初始参数与每一轮结束后被调用start返回一个 Result 数据类其中包含arrays最终的全局模型参数train_metrics_clientapp按轮号索引的、来自 ClientApp 的聚合训练指标evaluate_metrics_clientapp按轮号索引的、来自 ClientApp 的聚合评估指标evaluate_metrics_serverapp按轮号索引的、ServerApp 侧集中式评估产生的指标。2.1 每一轮包含的三个阶段从start的源码strategy.py可以看出其主体是一个for current_round in range(1, num_rounds 1)循环每一轮由三个明确划分的阶段组成训练阶段ClientApp 侧调用configure_train构造消息通过grid.send_and_receive发送给采样到的客户端并等待回复strategy.py随后调用aggregate_train聚合得到新的全局ArrayRecord与聚合训练指标strategy.py。send_and_receive的实现位于 framework/py/flwr/serverapp/grid/grid.py负责把消息推送到指定节点并拉取回复。评估阶段ClientApp 侧调用configure_evaluategrid.send_and_receiveaggregate_evaluate让一部分客户端在本地验证集上评估更新后的全局模型strategy.py。集中式评估阶段ServerApp 侧可选如果传入了evaluate_fn则对最新全局模型做集中式评估——这正是上一篇教程通过global_evaluate回调启用的机制strategy.py。3. 完整实现CustomFedAdagrad下面给出完整的策略实现。它继承自FedAdagrad源码位于 framework/py/flwr/serverapp/strategy/fedadagrad.py基于 [Reddi et al., 2020] 的自适应联邦优化算法并新增三个关键组件configure_train沿用上一篇的学习率衰减逻辑_update_best_acc每当集中式评估出现新的最优准确率时保存模型 checkpointset_save_path设置 WB 日志与模型 checkpoint 的保存目录重写的start在原生主循环的各个阶段嵌入 WB 日志与存档逻辑。import io import time from logging import INFO from pathlib import Path from typing import Callable, Iterable, Optional import torch import wandb from flwr.app import ArrayRecord, ConfigRecord, Message, MetricRecord from flwr.common import log, logger from flwr.serverapp import Grid from flwr.serverapp.strategy import FedAdagrad, Result from flwr.serverapp.strategy.strategy_utils import log_strategy_start_info PROJECT_NAME FLOWER-advanced-pytorch class CustomFedAdagrad(FedAdagrad): def configure_train( self, server_round: int, arrays: ArrayRecord, config: ConfigRecord, grid: Grid ) - Iterable[Message]: Configure the next round of federated training and maybe do LR decay. # Decrease learning rate by a factor of 0.5 every 5 rounds if server_round % 5 0 and server_round 0: config[lr] * 0.5 print(LR decreased to:, config[lr]) # Pass the updated config and the rest of arguments to the parent class return super().configure_train(server_round, arrays, config, grid) def set_save_path(self, path: Path): Set the path where wandb logs and model checkpoints will be saved. self.save_path path def _update_best_acc( self, current_round: int, accuracy: float, arrays: ArrayRecord ) - None: Update best accuracy and save model checkpoint if current accuracy is higher. if accuracy self.best_acc_so_far: self.best_acc_so_far accuracy logger.log(INFO, New best global model found: %f, accuracy) # Save the PyTorch model file_name fmodel_state_acc_{accuracy}_round_{current_round}.pth torch.save(arrays.to_torch_state_dict(), self.save_path / file_name) logger.log(INFO, New best model saved to disk: %s, file_name) def start( self, grid: Grid, initial_arrays: ArrayRecord, num_rounds: int 3, timeout: float 3600, train_config: Optional[ConfigRecord] None, evaluate_config: Optional[ConfigRecord] None, evaluate_fn: Optional[ Callable[[int, ArrayRecord], Optional[MetricRecord]] ] None, ) - Result: Execute the federated learning strategy logging results to WB and saving them to disk. # Init WB name f{str(self.save_path.parent.name)}/{str(self.save_path.name)}-ServerApp wandb.init(projectPROJECT_NAME, namename) # Keep track of best acc self.best_acc_so_far 0.0 log(INFO, Starting %s strategy:, self.__class__.__name__) log_strategy_start_info( num_rounds, initial_arrays, train_config, evaluate_config ) self.summary() log(INFO, ) # Initialize if None train_config ConfigRecord() if train_config is None else train_config evaluate_config ConfigRecord() if evaluate_config is None else evaluate_config result Result() t_start time.time() # Evaluate starting global parameters if evaluate_fn: res evaluate_fn(0, initial_arrays) log(INFO, Initial global evaluation results: %s, res) if res is not None: result.evaluate_metrics_serverapp[0] res arrays initial_arrays for current_round in range(1, num_rounds 1): log(INFO, ) log(INFO, [ROUND %s/%s], current_round, num_rounds) # ----------------------------------------------------------------- # --- TRAINING (CLIENTAPP-SIDE) ----------------------------------- # ----------------------------------------------------------------- # Call strategy to configure training round # Send messages and wait for replies train_replies grid.send_and_receive( messagesself.configure_train( current_round, arrays, train_config, grid, ), timeouttimeout, ) # Aggregate train agg_arrays, agg_train_metrics self.aggregate_train( current_round, train_replies, ) # Log training metrics and append to history if agg_arrays is not None: result.arrays agg_arrays arrays agg_arrays if agg_train_metrics is not None: log(INFO, \t└── Aggregated MetricRecord: %s, agg_train_metrics) result.train_metrics_clientapp[current_round] agg_train_metrics # Log to WB wandb.log(dict(agg_train_metrics), stepcurrent_round) # ----------------------------------------------------------------- # --- EVALUATION (CLIENTAPP-SIDE) --------------------------------- # ----------------------------------------------------------------- # Call strategy to configure evaluation round # Send messages and wait for replies evaluate_replies grid.send_and_receive( messagesself.configure_evaluate( current_round, arrays, evaluate_config, grid, ), timeouttimeout, ) # Aggregate evaluate agg_evaluate_metrics self.aggregate_evaluate( current_round, evaluate_replies, ) # Log training metrics and append to history if agg_evaluate_metrics is not None: log(INFO, \t└── Aggregated MetricRecord: %s, agg_evaluate_metrics) result.evaluate_metrics_clientapp[current_round] agg_evaluate_metrics # Log to WB wandb.log(dict(agg_evaluate_metrics), stepcurrent_round) # ----------------------------------------------------------------- # --- EVALUATION (SERVERAPP-SIDE) --------------------------------- # ----------------------------------------------------------------- # Centralized evaluation if evaluate_fn: log(INFO, Global evaluation) res evaluate_fn(current_round, arrays) log(INFO, \t└── MetricRecord: %s, res) if res is not None: result.evaluate_metrics_serverapp[current_round] res # Maybe save to disk if new best is found self._update_best_acc(current_round, res[accuracy], arrays) # Log to WB wandb.log(dict(res), stepcurrent_round) log(INFO, ) log(INFO, Strategy execution finished in %.2fs, time.time() - t_start) log(INFO, ) log(INFO, Final results:) log(INFO, ) for line in io.StringIO(str(result)): log(INFO, \t%s, line.strip(\n)) log(INFO, ) return result3.1 关键实现点拆解WB 初始化与 run 命名。start开头调用wandb.init(projectPROJECT_NAME, namename)其中name由保存路径的父目录名与目录名拼接而成形如YYYY-MM-DD/HH-MM-SS-ServerApp这样每次运行都能在 WB 中对应一个可辨识的 run。辅助方法_update_best_acc。它维护self.best_acc_so_far初始为0.0当accuracy self.best_acc_so_far时更新最优值并通过torch.save(arrays.to_torch_state_dict(), ...)把全局模型写盘。这里的ArrayRecord.to_torch_state_dict由框架提供实现于 framework/py/flwr/app/message/arrayrecord.py会把ArrayRecord转换为 PyTorch 的state_dictOrderedDict[str, torch.Tensor]因此可以无缝对接torch.save。每轮三类指标的 WB 写入。与原生start相比定制版本在三个位置追加了wandb.log(dict(...), stepcurrent_round)训练阶段wandb.log(dict(agg_train_metrics), stepcurrent_round)客户端评估阶段wandb.log(dict(agg_evaluate_metrics), stepcurrent_round)集中式评估阶段wandb.log(dict(res), stepcurrent_round)。MetricRecord本质上是受类型约束的字典定义见 framework/py/flwr/app/message/metricrecord.py所以dict(...)转换后可直接交给wandb.log。日志信息的复用。定制版调用了框架的log_strategy_start_info定义于 framework/py/flwr/serverapp/strategy/strategy_utils.py来打印轮数、ArrayRecord大小MB以及训练/评估ConfigRecord的摘要并调用self.summary()输出策略配置保证与原版start的日志体验一致。3.2 为什么不直接改父类源码start是策略执行的主循环它的默认实现位于 strategy.py。本教程采用「继承 覆写」的方式而不是修改框架源码FedAdagrad的构造函数参数如fraction_train1.0、min_train_nodes2、weighted_by_keynum-examples、eta1e-1、eta_l1e-1、tau1e-3等见 fedadagrad.py被完整保留自定义逻辑与内置策略解耦既便于后续升级 Flower 版本也更容易在多个 App 间复用。4. 在server_app.py中接入自定义策略策略定义好之后还需要在 ServerApp 入口中完成两件事实例化策略、调用set_save_path设置保存目录。关键是必须在实例化之后、start被调用之前完成set_save_path的调用。在server_app.py中追加导入并构造带时间戳的保存目录# ... unchanged # add this to the imports from datetime import datetime from pathlib import Path # ... unchanged app.main() def main(grid: Grid, context: Context) - None: Main entry point for the ServerApp. # ... unchanged # Initialize FedAdagrad strategy # strategy CustomFedAdagrad( ... ) # Get the current date and time current_time datetime.now() run_dir current_time.strftime(%Y-%m-%d/%H-%M-%S) # Save path is based on the current directory save_path Path.cwd() / foutputs/{run_dir} save_path.mkdir(parentsTrue, exist_okFalse) # Set the path where results and model checkpoints will be saved strategy.set_save_path(save_path) # ... rest unchanged这里的设计意图很明确目录名基于当前日期时间生成YYYY-MM-DD/HH-MM-SS因此每次执行flwr run都会产生一个全新的目录不同运行的结果互不覆盖。mkdir(parentsTrue, exist_okFalse)也保证了如果同一秒内重复运行会因目录已存在而报错从而避免意外覆盖历史实验结果。5. 本地运行与结果验证本教程会向当前工作目录写入模型 checkpoint并向 WB 上报指标因此直接在本地运行便于检查产出$ flwr run . local --stream其中--stream表示以流式方式输出运行日志若执行不带--stream的flwr run . local则只提交运行、打印 run ID 并立即返回不会持续输出日志。完整的本地运行工作流可参考仓库文档 how-to-run-flower-locally对应文件位于 framework/docs/source/。启动运行后可以观察到两件事本地目录产出outputs/YYYY-MM-DD/HH-MM-SS目录下会陆续出现模型 checkpoint 文件命名形如model_state_acc_{accuracy}_round_{round}.pth。回忆一下存档触发条件只有集中式评估阶段发现新的全局最优准确率时才会写入一个 checkpoint因此目录下的文件数量反映了「打破最优」的次数。WB 项目产出你的 WB 项目中会新建一个 run训练指标、客户端评估指标与集中式评估指标随轮次实时可视化。6. 回顾与下一步本教程的核心收获start是任何 Flower 策略的主入口掌握它的三段式循环训练 → 客户端评估 → 可选的集中式评估是定制策略的前提通过覆写start可以在不触碰框架源码的前提下把 WB 指标日志、模型 checkpoint 存档等横切逻辑平滑嵌入联邦学习主循环配合上一篇学到的configure_train覆写技巧同一策略类可以同时定制「每轮如何配置客户端」与「整体如何驱动训练」两者互不干扰set_save_path这类辅助 setter 方法让外部server_app.py能够在策略实例化后注入运行时配置是一种简单且灵活的扩展模式。下一步你可以继续学习教程系列的下一部分tutorial-series-customize-the-client-pytorch文档位于 framework/docs/source/tutorial-series-customize-the-client-pytorch.rst通过序列化自定义数据并将其封装进Message在ClientApp与ServerApp之间传递更多附加信息。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表