FEATURED · 精选文章

Flower 框架多策略联邦学习端到端测试实战:基于 TensorFlow 的 8 种聚合策略对比验证

发布时间 / 2026/9/17 21:10:53
来源 / 创域科博编辑部
栏目 / 资讯中心
Flower 框架多策略联邦学习端到端测试实战:基于 TensorFlow 的 8 种聚合策略对比验证 Flower 框架多策略联邦学习端到端测试实战基于 TensorFlow 的 8 种聚合策略对比验证【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower导读本文围绕 Flower 框架framework/e2e/strategies端到端测试模块展开讲解如何用 TensorFlowKeras构建联邦学习客户端并在一次模拟中依次验证 FedMedian、FedTrimmedAvg、QFedAvg、FaultTolerantFedAvg、FedAvgM、FedAdam、FedAdagrad、FedYogi 共 8 种服务端聚合策略的正确性与收敛性。读完本文你将掌握 Flower 模拟测试的完整工程结构、每种策略的适用场景与关键参数以及如何通过断言机制对策略做自动化回归验证。一、测试模块定位与整体设计该目录用于在 Flower 框架内对多种服务端策略Strategy做端到端E2E测试验证客户端训练 → 服务端聚合 → 集中式评估整条链路在多策略下均能正常工作。模块文件结构如下framework/e2e/strategies/ ├── README.md # 测试说明 ├── __init__.py # 包声明 ├── client.py # TensorFlow 客户端实现模型 数据 Flower 客户端 ├── test.py # 测试入口8 种策略逐一跑模拟并断言 └── pyproject.toml # 依赖与打包配置需要特别说明的是README 描述本模块使用 CIFAR-10 数据集与 CNN 模型但从当前仓库的实际代码看client.py 加载的是 MNIST 手写数字数据集28×28 灰度图模型也是全连接网络Flatten DenseREADME 与实现存在差异此外 README 提到测试集仅 10 个数据点实际代码中训练集与测试集均截取前SUBSET_SIZE 1000条。下文均以源码为准展开。二、依赖与运行环境pyproject.toml 明确了测试环境的依赖约束Python 版本3.11,4.0Flowerflwr[simulation]并直接引用仓库中 framework 包的本地构建产物 {root:parent:parent:uri}保证测试跑在待发布版本上而非 PyPI 旧版TensorFlowtensorflow-cpu2.18.0使用 CPU 版即可支撑本测试的小模型flwr[simulation]额外安装模拟所需的组件Ray 等使test.py可以在单机上用start_simulation模拟多个客户端无需真实部署分布式节点。三、TensorFlow 客户端实现client.pyclient.py 是联邦学习客户端的完整实现可以拆成三部分理解。3.1 模型与数据模型是两层全连接网络输入为 MNIST 的 28×28 展平向量输出 10 类model tf.keras.models.Sequential([ tf.keras.layers.Flatten(input_shape(28, 28)), tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dense(10), ]) model.compile( optimizertf.keras.optimizers.Adam(0.001), losstf.keras.losses.SparseCategoricalCrossentropy(from_logitsTrue), metrics[tf.keras.metrics.SparseCategoricalAccuracy()], )数据通过tf.keras.datasets.mnist.load_data()加载并用SUBSET_SIZE 1000同时截取训练集与测试集前 1000 条把单轮训练/评估的开销控制到秒级适合 CI 回归。3.2 Flower 客户端抽象FlowerClient继承NumPyClient把 Keras 权重与 Flower 的参数协议互相转换只需实现三个方法get_parameters(config)返回model.get_weights()即当前本地权重列表fit(parameters, config)先用服务端下发的参数model.set_weights(parameters)覆盖本地模型再以batch_size32、epochs1训练一轮返回新权重、样本数len(x_train)与空指标字典evaluate(parameters, config)设置服务端参数后在测试子集上评估返回(loss, len(x_test), {accuracy: accuracy})其中 accuracy 会被服务端聚合。3.3 入口与 ClientAppdef client_fn(context: Context): return FlowerClient().to_client() app ClientApp(client_fnclient_fn)client_fn供模拟环境按需实例化客户端ClientApp是 Flower 新 API 的应用封装。文件末尾的if __name__ __main__分支支持以独立进程方式连接真实服务端start_client(server_address127.0.0.1:8080, ...)意味着同一份客户端代码既能用于单机模拟也能用于真实分布式部署。四、测试驱动入口test.py运行机制test.py 是核心测试脚本用argv[1]指定要测试的策略名称例如python test.py FedAdam。整个脚本的逻辑链如下4.1 策略注册表STRATEGY_LIST [ FedMedian, FedTrimmedAvg, QFedAvg, FaultTolerantFedAvg, FedAvgM, FedAdam, FedAdagrad, FedYogi, ] OPT_IDX 5STRATEGY_LIST的索引顺序与 README 中的列举顺序一致8 种策略按稳健聚合 → 自适应优化排列。get_strat(name)按类名__name__精确匹配从列表中取出(idx, strategy_class)。4.2 策略族差异与 tau 参数OPT_IDX 5是脚本中一个精妙的分类标记列表前 4 项FedMedian、FedTrimmedAvg、QFedAvg、FaultTolerantFedAvg不属于 FedOpt 家族而索引 ≥ 5 的 4 项FedAvgM、FedAdam、FedAdagrad、FedYogi都继承自 FedOpt需要服务端优化器参数。因此脚本这样处理if start_idx OPT_IDX: strat_args[tau] 0.01tau是 FedOpt 家族的适应性控制参数docstring 中描述为Controls the algorithms degree of adaptabilityfedopt.py 中默认值为1e-9测试脚本统一将其放大到0.01避免分母项过小而引发数值不稳定这也是 8 种策略共享同一套公共参数时必要的差异化处理。4.3 公共参数与集中式评估所有策略统一注入两个公共参数strat_args { evaluate_fn: evaluate, initial_parameters: ndarrays_to_parameters(init_model.get_weights()), }initial_parameters用随机初始化的get_model()权重经ndarrays_to_parameters转为 FlowerParameters保证所有策略从同一初始点出发对比公平evaluate_fn服务端集中式评估函数每轮聚合后在完整测试子集1000 条上计算 loss 与 accuracy写入hist.metrics_centralized/hist.losses_centralized供断言使用。4.4 模拟配置与回归断言hist start_simulation( client_fnclient_fn, num_clients2, configServerConfig(num_rounds3), strategystrategy(**strat_args), )模拟配置为 2 个客户端、3 轮联邦学习。运行结束后执行收敛性断言assert ( hist.metrics_centralized[accuracy][0][1] / hist.metrics_centralized[accuracy][-1][1] ) 1.04 or (hist.losses_centralized[0][1] / hist.losses_centralized[-1][1]) 0.96该断言的语义是训练 3 轮后集中式准确率相比首轮不得衰减超过 4%或者损失不得上升超过 4%二者满足其一即可。因为每轮 1 epoch、仅 2 个客户端此阈值足够宽松以容忍波动又能捕获策略参数不兼容导致训练崩溃这类回归——例如某个策略收到非法参数报错、聚合出 NaN 权重首尾指标会剧烈劣化从而触发断言失败。五、8 种策略逐一解析结合源码以下策略全部位于 framework/py/flwr/server/strategy本测试通过统一接口调用源码可作深挖依据。5.1 FedMedian —— 中位数聚合抗拜占庭fedmedian.py 实现 Federated Median [Yin et al., 2018]其aggregate_fit调用aggregate_median逐参数取中位数而非加权平均。中位数对少数离群客户端上传的异常权重不敏感是鲁棒聚合robust aggregation的代表。accept_failures参数控制是否容忍含失败客户端的轮次默认继承 FedAvg 的 True。5.2 FedTrimmedAvg —— 截尾均值聚合fedtrimmedavg.py 实现带截尾均值的联邦平均核心参数beta默认0.2构造时校验必须满足0.0 beta 0.5表示从分布两端各裁掉的比例。聚合时调用aggregate_trimmed_avg(weights_results, self.beta)先按每个参数维度排序去掉头部与尾部各beta比例的极值再对剩余部分取平均兼顾鲁棒性与信息利用效率。5.3 QFedAvg —— 公平资源分配q-FFLqfedavg.py 实现 q-Fair Federated Learning [Li et al., 2020]目标是让损失高的客户端获得更大更新权重。其关键参数q_param默认0.2公平性指数q0 退化为普通 FedAvgq 越大越强调公平qffl_learning_rate默认0.1服务端学习率。实现上它在configure_fit中记录聚合前的pre_weights在aggregate_fit中把客户端更新量换算为梯度(u - v) / learning_rate并用损失估计每个客户端的局部 Lipschitz 常数源码中的hs_ffl最后经aggregate_qffl加权更新把公平性编码进服务端聚合公式。5.4 FaultTolerantFedAvg —— 容错联邦平均fault_tolerant_fedavg.py 在 FedAvg 基础上增加完成率门槛completion_rate len(results) / (len(results) len(failures)) if completion_rate self.completion_rate_fit: return None, {}参数min_completion_rate_fit与min_completion_rate_evaluate默认均0.5分别约束训练/评估轮的最低完成率若本轮成功返回的客户端比例不足则放弃本轮聚合。它在构造时强制accept_failuresTrue是面向真实网络环境客户端掉线、超时的实用策略。5.5 FedAvgM —— 带服务端动量的联邦平均FedAvgM 是经典 FedAvg 加服务端动量momentum的变体用于加速收敛、缓解 Non-IID 数据下的震荡实现于 fedavgm.py。5.6 FedAdam / FedAdagrad / FedYogi —— 服务端自适应优化这三个策略同属 FedOptAdaptive Federated Optimization [Reddi et al., 2020]家族把 Adam、Adagrad、Yogi 优化器搬到服务端客户端仍做普通 SGD 训练服务端则维护一阶/二阶矩估计来更新全局模型。公共参数在 fedopt.py 中定义eta默认1e-1服务端学习率eta_l默认1e-1客户端学习率用于梯度换算beta_1/beta_2默认0.0矩估计衰减系数三者的差异主要体现在 beta_2 的用法上tau默认1e-9适应性控制参数本测试中统一设为0.01见 4.2 节。三者实现分别在 fedadam.py、fedadagrad.py、fedyogi.py在 Non-IID 与异质数据场景下通常比普通 FedAvg 收敛更快、更稳。六、运行方式在framework/e2e下安装依赖后pip install -e ./strategies或按仓库 e2e 流程构建即可逐个策略运行python test.py FedMedian python test.py FedTrimmedAvg python test.py QFedAvg python test.py FaultTolerantFedAvg python test.py FedAvgM python test.py FedAdam python test.py FedAdagrad python test.py FedYogi每次运行都会打印 3 轮模拟的集中式评估曲线若对应策略的训练出现异常劣化脚本会以非零退出码结束从而接入 CI 流水线。得益于client.py底部的start_client分支同一套FlowerClient也可以脱离模拟器连接到127.0.0.1:8080的真实服务端参与联邦训练。七、小结framework/e2e/strategies是一个教科书级的策略回归测试样例用 1000 条 MNIST 数据 两层全连接网络在 3 轮 × 2 客户端的最小规模下覆盖 8 种策略并用首尾指标比值做自动断言。它同时展示了 Flower 的几个关键工程模式NumPyClient与 Keras 的无缝对接、start_simulation的单机多客户端模拟、initial_parameters保证对比公平以及通过OPT_IDX按策略族差异化注入参数tau的写法。对于想要在真实项目中选择聚合策略的开发者既可以对照 strategy 目录 的源码理解每种算法的数值细节也可以把本测试作为模板替换为自己的模型与数据集快速验证策略在目标任务上的表现。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻