Flower 与 scikit-learn 端到端集成测试:基于逻辑回归与 FedAvg 中心化评估的完整实现解析
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
本文以 Flower 仓库中 e2e-scikit-learn 端到端测试 为核心,深入解析 Flower 框架如何与 scikit-learn 生态协同完成联邦学习任务。你将掌握:如何用NumPyClient封装LogisticRegression模型实现联邦训练、如何通过FederatedDataset与IidPartitioner完成 MNIST 数据的 IID 分区、以及服务端如何基于FedAvg策略配合中心化评估(central evaluation)对训练结果进行断言验证。读完本文,你将能够复现这套最小可运行的 scikit-learn 联邦学习测试闭环,并理解其背后的设计动机。
一、定位:为什么需要 e2e-scikit-learn 测试
在 Flower 仓库中,framework/e2e 目录 集中存放了不同场景的端到端测试,其定位是:在任何改动合并进 Flower 之前,必须通过这些场景的验证。目录下并排陈列着 e2e-pytorch、e2e-tensorflow、e2e-jax、e2e-fastai、e2e-opacus、e2e-pandas、e2e-scikit-learn 等场景,覆盖了 Flower 所支持的主流机器学习生态。
而 e2e-scikit-learn 这一目录承担的任务非常明确:通过一个简单的逻辑回归(logistic regression)任务,验证 Flower 与 scikit-learn 的集成是否正常工作。它采用FedAvg策略并配合中心化评估,是整个 e2e 测试矩阵中验证 sklearn 生态兼容性的关键一环。
需要特别指出的是,这里的 scikit-learn 集成并不依赖任何专用包装器——Flower 通过通用接口NumPyClient与模型进行参数级的交互,这正是该测试能够验证"框架无关性"的原因所在。
二、整体架构与运行入口
整套测试的代码规模非常精简,仅包含 4 个核心文件:
| 文件 | 角色 | 职责 |
|---|---|---|
| client_app.py | 客户端应用 | 定义FlowerClient(NumPyClient),封装逻辑回归的 fit/evaluate |
| server_app.py | 服务端应用 | 构建ServerApp,驱动FedAvg与中心化评估并断言结果 |
| utils.py | 工具层 | 模型参数读写、MNIST 数据加载与 IID 分区 |
| simulation.py | 仿真入口 | 用start_simulation在单机模拟联邦训练 |
工程配置见 pyproject.toml,它同时声明了应用的两种组件入口:
[tool.flwr.app.components] serverapp = "e2e_scikit_learn.server_app:app" clientapp = "e2e_scikit_learn.client_app:app" [tool.flwr.federations] default = "local-simulation" [tool.flwr.federations.local-simulation] options.num-supernodes = 10依赖声明则明确限定版本范围:flwr[simulation](来自仓库父目录的本地源码引用)、flwr-datasets[vision]>=0.5.0,<1.0.0、scikit-learn>=1.1.1,<2.0.0。值得注意的是flwr[simulation]通过{root:parent:parent:uri}直接引用本地框架源码,意味着这套测试始终针对当前仓库的框架实现进行验证,而非已发布的 PyPI 版本。
三、客户端实现:用 NumPyClient 封装逻辑回归
3.1 模型构造的关键技巧
client_app.py 的模型构建是本文最值得咀嚼的部分:
model = LogisticRegression( penalty="l2", max_iter=1, # local epoch warm_start=True, # prevent refreshing weights when fitting )三个参数各有用意:
penalty="l2":使用 L2 正则化,与联邦学习中常见的权重衰减语义对齐;max_iter=1:注释明确说明它等价于"local epoch"——每次fit只做一轮优化,模拟联邦学习中每轮客户端只做一次本地更新的场景;warm_start=True:这是逻辑回归参与联邦学习的前提。sklearn 默认在每次fit时重新初始化并解算模型,而联邦学习要求模型基于服务端下发的全局参数继续迭代;开启warm_start后,fit会以现有参数为起点继续优化,从而"接住"全局聚合结果。
3.2 初始参数的显式设置
逻辑回归在首次fit之前,coef_、intercept_、classes_等属性都是未初始化的。但 Flower 的服务端在启动时就会向客户端索要初始参数(用于第一轮下发),因此必须在训练前显式填充。这一点在 utils.py 的set_initial_params中完成:
def set_initial_params(model: LogisticRegression): n_classes = 10 # MNIST has 10 classes n_features = 784 # Number of features in dataset model.classes_ = np.array([i for i in range(10)]) model.coef_ = np.zeros((n_classes, n_features)) if model.fit_intercept: model.intercept_ = np.zeros((n_classes,))注释中特别给出了依据:sklearn.linear_model.LogisticRegression的文档明确说明这些参数在fit调用前未初始化。这里初始化为全零向量,对应 FedAvg 从零模型开始平均的语义。
3.3 参数序列化与联邦训练循环
utils.py中get_model_parameters/set_model_params完成了 sklearn 原生属性与 Flower 参数列表(List[np.ndarray])之间的双向转换:有fit_intercept时参数列表为[coef_, intercept_],否则仅为[coef_]。
客户端则实现NumPyClient的三个抽象方法:
class FlowerClient(NumPyClient): def get_parameters(self, config): return utils.get_model_parameters(model) def fit(self, parameters, config): utils.set_model_params(model, parameters) with warnings.catch_warnings(): warnings.simplefilter("ignore") model.fit(X_train, y_train) return utils.get_model_parameters(model), len(X_train), {} def evaluate(self, parameters, config): utils.set_model_params(model, parameters) loss = log_loss(y_test, model.predict_proba(X_test)) accuracy = model.score(X_test, y_test) return loss, len(X_test), {"accuracy": accuracy}fit中用warnings.catch_warnings配合simplefilter("ignore")屏蔽收敛告警,这是因为max_iter=1极大概率不满足 sklearn 的收敛判定——这是测试场景下的刻意选择,避免无关告警干扰 e2e 判定。evaluate返回对数损失、样本数与 accuracy 字典,其中 accuracy 会被服务端作为分布式评估指标聚合。
四、数据层:FederatedDataset 与 IID 分区
utils.py 的load_data展示了 Flower Datasets 的标准用法:
partitioner = IidPartitioner(num_partitions=num_partitions) fds = FederatedDataset( dataset="ylecun/mnist", partitioners={"train": partitioner}, ) dataset = fds.load_partition(partition_id, "train").with_format("numpy")要点如下:
IidPartitioner(num_partitions=10):将 MNIST 训练集按 IID(独立同分布)方式切分为 10 份,对应客户端总数;partition_id=np.random.choice(num_partitions):在模块导入时随机抽取一个分区作为"当前客户端"的数据,这是 e2e 测试场景下快速生成多个客户端数据分片的手法;.with_format("numpy"):将 HF Dataset 转换为 NumPy 视图,适配 sklearn 的ndarray输入;- 图像展平:
X = batch["image"].reshape((len(dataset), -1)),将 28×28 图像展平为 784 维向量,正好与set_initial_params中的n_features=784对应; - 边端数据再拆分:每个客户端拿到分区后,再按 80%/20% 切分为本地训练集与本地测试集——这模拟了"数据在设备本地,且设备有自己的本地评估集"的现实场景。
fds被缓存为模块级全局变量(fds = None),确保多次调用只初始化一次FederatedDataset,避免重复下载与分区开销。
五、服务端实现:FedAvg + 中心化评估与结果断言
5.1 基于新 API 的 ServerApp 流程
server_app.py 使用 Flower 新式ServerApp编程模型:
app = fl.serverapp.ServerApp() @app.main() def main(grid, context): context = fl.server.LegacyContext( context=context, config=fl.server.ServerConfig(num_rounds=3), ) workflow = fl.server.workflow.DefaultWorkflow() workflow(grid, context) ...ServerConfig(num_rounds=3)将联邦训练轮数固定为 3 轮,保证 e2e 测试的执行时间可控;LegacyContext位于 framework/py/flwr/server/compat/legacy_context.py,是框架为兼容旧式 API 提供的上下文适配层;DefaultWorkflow定义在 framework/py/flwr/server/workflow/default_workflows.py,内部即执行经典的联邦平均工作流(分发全局参数 → 客户端 fit → 聚合 → 中心化评估)。
5.2 中心化评估与训练成功判据
README 明确本测试采用central evaluation(中心化评估)。在FedAvg默认配置下,ServerApp流程会由服务端在每轮聚合后自行对测试集做一次集中评估(对应FedAvg的evaluate流程)。训练是否成功,由断言判定:
assert ( hist.losses_distributed[-1][1] == 0 or (hist.losses_distributed[0][1] / hist.losses_distributed[-1][1]) >= 0.98 )该判据的语义是:训练必须"有进展"——要么末轮分布式损失恰好为 0(理想收敛),要么首轮损失与末轮损失之比不低于 0.98。换言之,3 轮训练后损失必须下降至少约 2%,否则视为训练异常、测试失败。这个宽松阈值兼顾了 e2e 测试的稳定性(避免随机性导致抖动误报)与有效性(能捕获模型完全不学习的回归缺陷)。
5.3 客户端状态时间戳的单调性检查
服务端还内置了一个有趣的附加校验函数record_state_metrics,用于验证客户端跨轮次状态的正确性:
STATE_VAR = "timestamp" def record_state_metrics(metrics): if STATE_VAR not in metrics[0][1]: return {} states = [] for _, m in metrics: states.append([float(tt) for tt in m[STATE_VAR].split(",")]) for client_state in states: if len(client_state) == 1: continue deltas = np.diff(client_state) assert np.all(deltas > 0), f"Timestamps are not monotonically increasing: {client_state}" return {STATE_VAR: states}它假设客户端会上报以逗号分隔的"时间戳串",随后断言每个客户端的时间戳序列严格单调递增——用于捕获客户端状态在轮次间被意外重置或乱序的 bug。注意在ServerApp主流程中该函数仅在旧式start_server分支被挂接(evaluate_metrics_aggregation_fn=record_state_metrics),且当客户端状态只有单条记录时检查自动跳过。
六、两种运行模式:仿真与真实网络
该测试同时提供了两种运行方式,对应 Flower 的两套执行引擎:
6.1 本地仿真(simulation.py)
hist = fl.simulation.start_simulation( client_fn=client_fn, num_clients=2, config=fl.server.ServerConfig(num_rounds=3), )start_simulation在单进程内以线程/进程方式模拟 2 个客户端,client_fn复用client_app中的工厂函数。这是 CI 中最轻量的验证路径,无需启动任何网络服务。
6.2 真实网络(client/server 直连)
client_app.py 的__main__分支与 server_app.py 的__main__分支共同构成经典的三进程模式:
# client 侧 start_client(server_address="127.0.0.1:8080", client=FlowerClient().to_client()) # server 侧 strategy = fl.server.strategy.FedAvg(evaluate_metrics_aggregation_fn=record_state_metrics) hist = fl.server.start_server( server_address="127.0.0.1:8080", config=fl.server.ServerConfig(num_rounds=3), strategy=strategy, )客户端通过 gRPC 连接到127.0.0.1:8080上的服务端,服务端使用FedAvg策略(其定义位于 framework/py/flwr/server/strategy/fedavg.py)。该模式下record_state_metrics被真正挂载到策略上,并且__main__末尾还有一条与轮次相关的断言:
if STATE_VAR in hist.metrics_distributed: state_metrics_last_round = hist.metrics_distributed[STATE_VAR][-1] assert ( len(state_metrics_last_round[1][0]) == 2 * state_metrics_last_round[0] ), "There should be twice as many entries in the client state as rounds"即:若客户端上报了时间戳状态,则末轮状态条数应为轮数的两倍(对应 3 轮训练中的参数下发与评估两个阶段)。
这种"server + 2 个 client 进程 + 后台等待 + 超时保护"的运行编排方式,与仓库根 e2e 脚本 test_legacy.sh 中的模式一致——后台启动服务端(timeout 3m python server_app.py &),间隔数秒依次拉起多个客户端,最后以服务端退出码判定训练是否成功。
七、实战速览:如何复现与验证
在仓库环境下复现这套测试的完整路径如下:
- 查看测试入口与断言:通读 README.md 了解测试目标(逻辑回归 + FedAvg + 中心化评估);
- 安装依赖:基于 pyproject.toml 安装
flwr[simulation]、flwr-datasets与scikit-learn; - 运行本地仿真:执行
python simulation.py,验证 2 客户端 × 3 轮的仿真流程通过损失断言; - 运行真实网络模式:分别以两个终端启动
server_app.py与一个/多个client_app.py,观察127.0.0.1:8080上的联邦训练与中心化评估; - 判定结果:无论哪种模式,最终都依据"末轮损失趋近 0 或相对首轮下降 ≥2%"的断言是否通过来判断集成是否正常。
八、总结
e2e-scikit-learn 虽然是一个"只有几行说明"的测试目录,但其背后的实现浓缩了 Flower 集成非深度学习框架的全部要点:
- 接口层:
NumPyClient以参数数组为媒介,天然适配任何能暴露coef_/intercept_的 sklearn 模型; - 模型层:
warm_start=True与显式set_initial_params解决了 sklearn 迭代式求解器与联邦"全局参数接力"之间的适配问题; - 数据层:
FederatedDataset+IidPartitioner让 MNIST 的联邦数据切分只需数行代码; - 服务端:
ServerApp+DefaultWorkflow+FedAvg组成标准联邦平均流水线,配合中心化评估与"损失下降 ≥2%"的断言,构成了一套自动化的兼容性回归测试。
对于希望将 scikit-learn 模型(尤其是各类迭代式、可热启动的线性模型与树模型)接入 Flower 的开发者而言,这个测试目录就是最直接的"最小可用参考实现"。
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考