尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Flower 端到端测试实战:基于 PyTorch 与 CIFAR-10 验证 FedAvg 联邦训练全链路

Flower 端到端测试实战:基于 PyTorch 与 CIFAR-10 验证 FedAvg 联邦训练全链路 Flower 端到端测试实战基于 PyTorch 与 CIFAR-10 验证 FedAvg 联邦训练全链路【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower本篇技术指南围绕 Flower 框架仓库中的framework/e2e/e2e-pytorch端到端测试模块展开讲解它如何在发布前用 PyTorch、CIFAR-10 数据集与一个轻量 CNN 模型对FedAvg策略的完整联邦训练链路客户端训练、服务端聚合、指标回传、客户端状态记录进行自动化验证。读完本文你将掌握该测试模块的架构设计、源码级实现细节以及它的三种运行方式与通过判定标准并可直接将其作为模板编写你自己的框架级端到端测试。测试模块定位发布前的全链路体检Flower 仓库的framework/e2e目录集中存放了用于验证框架不同能力组合的端到端测试场景其根目录 README 明确说明该目录下的每个子目录对应一个在改动合入 Flower 之前必须被测试并验证的场景。e2e-pytorch正是其中之一它负责回答一个核心问题当用户以 PyTorch 编写客户端、以FedAvg作为服务端策略时从数据加载、模型训练到指标聚合与状态传递的整条链路是否工作正常。从目录结构看该模块是麻雀虽小、五脏俱全的完整 Flower AppREADME.md模块说明即本指南对应的原文档pyproject.tomlFlower App 的工程元数据与联邦配置client_app.py客户端侧完整实现数据、模型、训练、指标server_app.py服务端聚合逻辑与最终断言simulation.py经典模拟运行入口start_simulationsimulation_next.py新一代模拟运行入口run_simulation ServerApp根据原文档的描述该测试的核心设定为使用 CIFAR-10 数据集与一个 CNN 模型测试 Flower 与 PyTorch 的集成采用FedAvg策略并提供自定义的evaluate_metrics_aggregation_fn训练数据使用一个子集、测试仅使用 10 个数据点以控制运行时长。需要说明的是原文档记载训练子集规模为 1000而当前仓库代码中 client_app.py 定义的SUBSET_SIZE 100实际生效值以代码为准写作时请留意 README 与代码之间的这一细微差异。数据与模型面向测试速度的最小化设计测试要频繁运行因此数据与模型都做了刻意的最小化设计保证几轮训练可以在秒级完成。从 Hugging Face 加载 CIFAR-10客户端通过 Hugging Face 的datasets库加载 CIFAR-10Cifar10Dataset 封装了load_dataset(uoft-cs/cifar10, splitsplit)并在__getitem__中应用变换、返回(img, label)元组。预处理使用标准的ToTensor()加Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))load_data。关键的最小化设定体现在Subset截取训练集只取前SUBSET_SIZE当前代码为 100个样本DataLoader(batch_size32, shuffleTrue)测试集只取前 10 个样本用于evaluate阶段的快速验证整个数据集首次访问时从 Hugging Face 下载后续由datasets库本地缓存。简化版 CNNPyTorch 60 分钟入门教程风格模型 Net 是一个参数规模很小的 CNN注释表明其改编自 PyTorch: A 60 Minute Blitz 教程conv13 通道输入 → 4 通道5×5 卷积核接 2×2 MaxPoolconv24 通道 → 8 通道5×5 卷积核接 2×2 MaxPoolfc18×5×5 展平 → 32fc232 → 16fc316 → 10对应 CIFAR-10 的 10 个类别训练与评估函数同样保持精简train使用CrossEntropyLoss与SGD(lr0.001, momentum0.9)轮数由调用方传入test在no_grad下计算损失与准确率并返回。计算设备通过DEVICE torch.device(cuda:0 if torch.cuda.is_available() else cpu)自动选择。客户端实现NumPyClient 与客户端状态记录FlowerClient 继承NumPyClient是测试的核心验证对象之一。除了标准的参数收发它还刻意演示了 Flower 的**客户端状态state**机制用于端到端验证客户端在一次运行中跨多轮保持状态这一能力。参数交换与训练评估get_parameters将net.state_dict()的各张量转为 NumPy 数组返回这是NumPyClient约定的序列化边界fit先set_parameters用服务端下发的参数覆写模型再训练 1 个 epoch返回(新参数, 训练样本数, 指标字典)evaluateset_parameters后计算损失与准确率返回(loss, 测试样本数, {accuracy: ..., ...})set_parameters源码用OrderedDict按state_dict的键顺序重建张量字典并以strictTrue加载保证参数结构与模型严格对齐。用 ConfigRecord 记录时间戳状态客户端通过state.config_records维护一个名为timestamp的累积状态变量STATE_VAR timestamp_record_timestamp_to_state把当前时间戳datetime.now().timestamp()追加到该变量的逗号分隔字符串中若已有值则追加,新时间戳_retrieve_timestamp_from_state读取当前累计的时间戳串fit与evaluate每次执行都会调用记录函数并把读回的时间戳字符串放进返回给服务端的指标字典。这意味着每个客户端每完成一轮fit或evaluate其状态中就会多一个时间戳条目——服务端可以据此验证客户端状态确实跨轮持续存在且按时间单调递增。这正是该测试超越能跑通层面的深层验证点。服务端与指标聚合FedAvg 单调时间戳断言服务端逻辑在 server_app.py 中完整实现了原文档承诺的FedAvg策略 evaluate_metrics_aggregation_fn组合。自定义指标聚合函数record_state_metricsserver_app.py接收各客户端回传的指标元组列表逐客户端把逗号分隔的时间戳串解析为浮点数列表然后用np.diff计算相邻时间戳差值并断言差值必须全部大于 0即同一客户端的状态时间戳严格单调递增若断言失败会抛出明确的错误信息Timestamps are not monotonically increasing。此外该函数对缺少timestamp键的指标做了防御性处理直接返回空字典避免破坏其他不含该指标的客户端。这个函数随后被传入FedAvg的evaluate_metrics_aggregation_fn参数即 README.md 中所指的自定义评估指标聚合逻辑。ServerApp 主流程与损失收敛断言以 Flower 新一代 API 编写的ServerApp主流程如下app fl.serverapp.ServerApp() app.main() def main(grid, context): context fl.server.LegacyContext( contextcontext, configfl.server.ServerConfig(num_rounds2), ) workflow fl.server.workflow.DefaultWorkflow() workflow(grid, context) hist context.history assert ( hist.losses_distributed[-1][1] 0 or (hist.losses_distributed[0][1] / hist.losses_distributed[-1][1]) 0.98 )要点拆解通过LegacyContext把新框架的上下文包装成传统ServerConfig(num_rounds2)语义并执行DefaultWorkflow内部完成客户端选择、fit、evaluate等默认流程运行结束后从context.history取出分布式损失序列断言最后一轮损失为 0或首轮损失与末轮损失之比 ≥ 0.98。由于测试集仅 10 个样本且训练轮次很少这一宽泛的收敛条件保证了测试不会被训练效果的不确定性干扰只验证链路正确性。三种运行方式从传统 CLI 到新一代模拟引擎该模块提供了多套入口覆盖了 Flower 不同历史阶段的运行范式是理解框架 API 演进的极佳素材。方式一进程级 start_server / start_client传统驱动模式client_app.py与server_app.py底部的__main__分支支持以独立进程方式运行客户端监听127.0.0.1:8080start_client并以空RecordDict()初始化状态服务端start_server使用FedAvg(evaluate_metrics_aggregation_fnrecord_state_metrics)与ServerConfig(num_rounds2)。这种模式与仓库根级 test_legacy.sh 的编排方式同源后台先启动python server_app.py随后并行启动两个python client_app.py等待服务端进程退出后按其退出码判定训练是否成功。可见 e2e-pytorch 同样可以被这类脚本化方式拉起多个客户端进程进行联调。方式二start_simulation进程内模拟simulation.py 直接在单进程内模拟整个联邦strategy fl.server.strategy.FedAvg(evaluate_metrics_aggregation_fnrecord_state_metrics) hist fl.simulation.start_simulation( client_fnclient_fn, num_clients2, configfl.server.ServerConfig(num_rounds2), strategystrategy, )其中client_fn直接复用client_app.py中从Context构造客户端的工厂函数client_fn(context)返回FlowerClient(context.state).to_client()因此模拟模式下客户端状态同样会被真实创建与持久化。方式三run_simulation ServerApp新一代推荐方式simulation_next.py 展示的是当前推荐的写法以ServerApp(configServerConfig(num_rounds2))与服务端、以ClientApp为客户端调用fl.simulation.run_simulation(server_app..., client_app..., num_supernodes2)。这种方式与 pyproject.toml 中声明的 App 组件serverapp e2e_pytorch.server_app:app、clientapp e2e_pytorch.client_app:app完全对应也是flwr run命令行执行时实际加载的入口。通过标准双重断言把关无论走哪条运行路径测试都以两组断言作为通过的唯一标准损失收敛断言所有入口共有hist.losses_distributed[-1][1] 0或首末轮损失比 ≥ 0.98状态规模断言simulation.py与server_app.py的__main__分支取hist.metrics_distributed[timestamp][-1]断言len(客户端时间戳列表) 2 * 轮数。这是因为每轮每个客户端会执行一次fit与一次evaluate各追加一个时间戳因此在 2 轮模拟下应恰好积累 4 个时间戳条目。第二组断言从指标回传的维度反向验证了客户端状态机制的正确性若fit/evaluate未执行、状态未持久化或指标未回传长度校验必然失败。这一设计让测试不仅验证能训练还验证了状态真的跨轮存在。工程配置pyproject.toml 中的联邦元数据pyproject.toml 除了声明依赖datasets4.0.0,5.0.0、torch2.10.0,3.0.0、torchvision0.25.0,0.26.0、tqdm及flwr[simulation]还携带 Flower App 的标准配置段[tool.flwr.app.components]声明serverapp与clientapp的模块级入口供flwr run发现[tool.flwr.federations]定义名为local-simulation的联邦options.num-supernodes 10指定模拟引擎下的默认 SuperNode 数量default local-simulation将该联邦设为flwr run的默认目标。对比 e2e 根目录的 pyproject.toml同样以local-simulation联邦、10 个 SuperNode 作为默认配置可以看出这是 e2e 测试家族统一的工程约定。小结一份可复用的框架级测试模板e2e-pytorch的价值不在于训练效果而在于它把框架发布前的链路验证做成了标准动作最小化数据与模型保证测试速度FedAvg evaluate_metrics_aggregation_fn覆盖策略扩展点客户端ConfigRecord状态机制提供跨轮状态验证多套运行入口兼容传统 CLI 与新一代模拟引擎。当你需要为新的框架能力编写端到端测试时以 client_app.py 为客户端骨架、server_app.py 为服务端断言骨架、simulation_next.py 为运行入口即可快速搭建起同样严谨的测试场景。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表