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 App:
- README.md:模块说明(即本指南对应的原文档)
- pyproject.toml:Flower App 的工程元数据与联邦配置
- client_app.py:客户端侧完整实现(数据、模型、训练、指标)
- server_app.py:服务端聚合逻辑与最终断言
- simulation.py:经典模拟运行入口(
start_simulation) - simulation_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-10,Cifar10Dataset 封装了load_dataset("uoft-cs/cifar10", split=split),并在__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_size=32, shuffle=True); - 测试集只取前 10 个样本,用于
evaluate阶段的快速验证; - 整个数据集首次访问时从 Hugging Face 下载,后续由
datasets库本地缓存。
简化版 CNN:PyTorch 60 分钟入门教程风格
模型 Net 是一个参数规模很小的 CNN,注释表明其改编自 "PyTorch: A 60 Minute Blitz" 教程:
conv1:3 通道输入 → 4 通道,5×5 卷积核,接 2×2 MaxPoolconv2:4 通道 → 8 通道,5×5 卷积核,接 2×2 MaxPoolfc1:8×5×5 展平 → 32;fc2:32 → 16;fc3:16 → 10(对应 CIFAR-10 的 10 个类别)
训练与评估函数同样保持精简:train使用CrossEntropyLoss与SGD(lr=0.001, momentum=0.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,返回(新参数, 训练样本数, 指标字典);evaluate:set_parameters后计算损失与准确率,返回(loss, 测试样本数, {"accuracy": ..., ...});set_parameters(源码):用OrderedDict按state_dict的键顺序重建张量字典,并以strict=True加载,保证参数结构与模型严格对齐。
用 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_metrics(server_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( context=context, config=fl.server.ServerConfig(num_rounds=2), ) 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_rounds=2)语义,并执行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:8080(start_client),并以空RecordDict()初始化状态; - 服务端
start_server使用FedAvg(evaluate_metrics_aggregation_fn=record_state_metrics)与ServerConfig(num_rounds=2)。
这种模式与仓库根级 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_fn=record_state_metrics) hist = fl.simulation.start_simulation( client_fn=client_fn, num_clients=2, config=fl.server.ServerConfig(num_rounds=2), strategy=strategy, )其中client_fn直接复用client_app.py中从Context构造客户端的工厂函数(client_fn(context)返回FlowerClient(context.state).to_client()),因此模拟模式下客户端状态同样会被真实创建与持久化。
方式三:run_simulation + ServerApp(新一代推荐方式)
simulation_next.py 展示的是当前推荐的写法:以ServerApp(config=ServerConfig(num_rounds=2))与服务端、以ClientApp为客户端,调用fl.simulation.run_simulation(server_app=..., client_app=..., num_supernodes=2)。这种方式与 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 除了声明依赖(datasets>=4.0.0,<5.0.0、torch>=2.10.0,<3.0.0、torchvision>=0.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),仅供参考