拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Flower 与 scikit-learn 联邦学习快速入门:在 Iris 数据集上训练 Logistic Regression

Flower 与 scikit-learn 联邦学习快速入门在 Iris 数据集上训练 Logistic Regression【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower本教程基于 Flower 官方 Quickstart 文档framework/docs/source/tutorial-quickstart-scikitlearn.rst及其配套示例examples/quickstart-sklearn演示如何用 Flower 的现代 Message API 与 scikit-learn 构建一个联邦逻辑回归系统在 Iris 数据集上通过flwr new一键生成工程、用 Flower Datasets 做数据分区、以ClientApp/ServerApp描述训练与聚合逻辑最后用flwr run在本地模拟联邦环境端到端跑通。读完本文你将掌握scikit-learn 模型如何接入 Flower 联邦训练的完整套路numpy 参数与ArrayRecord的双向转换、IidPartitioner数据分区、FedAvg策略的启动方式以及如何通过--run-config覆盖超参数。准备工作与项目生成教程建议先在独立的 Python 虚拟环境中操作避免污染全局环境。环境就绪后先安装 Flower# In a new Python environment $ pip install flwr然后使用flwr new从 Flower Labs 官方模板拉取一个完整的 Flower scikit-learn 项目$ flwr new flwrlabs/quickstart-sklearn该命令会生成一个名为quickstart-sklearn的新目录其中包含运行两节点联邦所需的全部文件。默认情况下生成的应用带有一个本地模拟 profileflwr run会把运行提交给一个受管理的本地 SuperLink再由 SuperLink 通过 Flower Simulation Runtime 执行这次联邦运行数据集的划分由 Flower Datasets 的IidPartitioner完成。quickstart-sklearn ├── sklearnexample │ ├── __init__.py │ ├── client_app.py # Defines your ClientApp │ ├── server_app.py # Defines your ServerApp │ └── task.py # Defines your model, training and data loading ├── pyproject.toml # Project metadata like dependencies and configs └── README.md如果你想在仓库中直接查看这份示例的完整代码它的结构与上面完全一致入口分别位于 examples/quickstart-sklearn/sklearnexample/client_app.py、examples/quickstart-sklearn/sklearnexample/server_app.py 和 examples/quickstart-sklearn/sklearnexample/task.py。项目依赖与入口声明生成的pyproject.toml中声明了三个关键依赖见 examples/quickstart-sklearn/pyproject.tomldependencies [ flwr[simulation]1.36.0, flwr-datasets[vision]0.6.1, scikit-learn1.6.1, ]flwr[simulation]Flower 框架本体simulationextra 提供本地模拟所需的运行时组件flwr-datasets负责数据集下载、分区与预处理本示例使用其IidPartitioner与FederatedDatasetscikit-learn机器学习库提供LogisticRegression模型。同时[tool.flwr.app.components]段声明了联邦应用的入口flwr run正是据此定位你的ClientApp与ServerApp[tool.flwr.app.components] serverapp sklearnexample.server_app:app clientapp sklearnexample.client_app:app依赖安装可直接执行pip install -e .运行联邦训练进入项目目录并启动运行$ cd quickstart-sklearn # Run with default arguments and stream logs $ flwr run . --stream--stream表示持续流式输出日志如果使用普通的flwr run .命令会提交运行、打印运行 ID 后直接返回不再跟随日志输出。使用默认参数运行时你会看到类似下面的流式日志Starting local SuperLink on 127.0.0.1:39091... Successfully started run 1859953118041441032 INFO : Starting FedAvg strategy: INFO : ├── Number of rounds: 3 INFO : [ROUND 1/3] INFO : configure_train: Sampled 2 nodes (out of 2) INFO : aggregate_train: Received 2 results and 0 failures INFO : └── Aggregated MetricRecord: {train_logloss: 1.3937176081476854} INFO : configure_evaluate: Sampled 2 nodes (out of 2) INFO : aggregate_evaluate: Received 2 results and 0 failures INFO : └── Aggregated MetricRecord: {test_logloss: 1.23306, accuracy: 0.69154, precision: 0.68659, recall: 0.68046, f1: 0.65752} INFO : [ROUND 2/3] INFO : ... INFO : [ROUND 3/3] INFO : ... INFO : Strategy execution finished in 17.87s INFO : Final results: INFO : ServerApp-side Evaluate Metrics: INFO : {}从日志可以清晰看到联邦训练的完整生命周期本地 SuperLink 启动、运行提交成功、FedAvg策略初始化、每轮采样节点 → 训练聚合 → 采样评估 → 评估聚合最后策略执行结束并输出最终指标。需要说明的是日志中展示的是 3 轮运行的示例输出项目默认的num-server-rounds配置为 25 轮定义在pyproject.toml的[tool.flwr.app.config]段你可以按下一节的方式覆盖它。用 --run-config 覆盖超参数[tool.flwr.app.config]段中的参数可以在命令行按需覆盖无需改动代码# Override some arguments $ flwr run . --run-config num-server-rounds5 local-epochs2--run-config后面以空格分隔的keyvalue会被注入运行配置ClientApp与ServerApp通过context.run_config[...]读取。README 中还给出了覆盖字符串型参数的写法示例examples/quickstart-sklearn/README.mdflwr run . --run-config penaltyl1 --stream该项目pyproject.toml中预置的配置项及其默认值如下配置键默认值说明penaltyl2LogisticRegression的正则化类型可覆盖为l1等num-server-rounds25FedAvg 联邦训练的轮数min-available-clients2系统中最少可用的客户端数量save-modelfalse是否在训练结束后把最终模型保存到本地磁盘min-available-clients与save-model等键在代码中的读取位置可参见 examples/quickstart-sklearn/sklearnexample/server_app.py。数据Flower Datasets 与 IidPartitioner 分区本示例使用 Flower Datasets 下载并划分 Iris 数据集hitorilabs/iris并采用IidPartitioner生成num_partitions个独立同分布分区。Iris 是经典的表格分类数据集本示例只使用其中 4 个数值特征列标签为花的品种3 类。数据加载逻辑实现在 examples/quickstart-sklearn/sklearnexample/task.py 的load_data()中FEATURES [petal_length, petal_width, sepal_length, sepal_width] partitioner IidPartitioner(num_partitionsnum_partitions) fds FederatedDataset(datasethitorilabs/iris, partitioners{train: partitioner}) dataset fds.load_partition(partition_id, train).with_format(pandas)[:] X dataset[FEATURES] y dataset[species] # Split the on-edge data: 80% train, 20% test X_train, X_test X[: int(0.8 * len(X))], X[int(0.8 * len(X)) :] y_train, y_test y[: int(0.8 * len(y))], y[int(0.8 * len(y)) :] return X_train.values, y_train.values, X_test.values, y_test.values关键点IidPartitioner(num_partitions...)把完整训练集随机分成num_partitions个分区每个分区都近似代表整体分布。其源码位于 datasets/flwr_datasets/partitioner/iid_partitioner.py每个分区从数据集中随机采样是它的核心语义。FederatedDataset(datasethitorilabs/iris, partitioners{train: partitioner})声明训练集按给定分区器切分。fds.load_partition(partition_id, train)取出指定客户端所属的分区with_format(pandas)转换为 pandas DataFrame 以便按列索引。每个ClientApp都会调用这个函数用context.node_config提供的partition-id和num-partitions构造自己的本地数据加载器。在分区内部再做 80/20 切分前者用于本地训练后者用于本地评估。如果IidPartitioner不满足需求Flower Datasets 还提供其他分区器如按标签分布的 Non-IID 分区器可按需替换。模型scikit-learn LogisticRegression 的联邦化改造模型定义同样在task.py中核心是create_log_reg_and_instantiate_parameters()函数def create_log_reg_and_instantiate_parameters(penalty): model LogisticRegression( penaltypenalty, max_iter1, # local epoch warm_startTrue, # prevent refreshing weights when fitting, solversaga, ) # Setting initial parameters, akin to model.compile for keras models set_initial_params(model, n_featureslen(FEATURES), n_classeslen(UNIQUE_LABELS)) return model几个参数的含义与联邦场景的适配max_iter1把 scikit-learn 的优化迭代次数当作本地 epoch来用每次fit只做一轮优化warm_startTrue再次fit时沿用上一次的权重而不是重新初始化这是联邦学习中在服务端下发的参数基础上继续训练的前提solversaga适合中小规模数据且支持l1/l2正则的求解器penalty正则化类型来自context.run_config[penalty]默认l2。与 Keras 等框架不同scikit-learn 的LogisticRegression在fit之前参数是未初始化的而联邦流程要求服务端启动时就能拿到一组全局初始参数。因此task.py提供了set_initial_params()显式设置classes_、把coef_置零、把intercept_置零作为联邦的初始全局模型。对应地get_model_params()和set_model_params()负责在 numpy ndarray 列表与模型对象之间搬运参数def get_model_params(model: LogisticRegression) - NDArrays: if model.fit_intercept: params [model.coef_, model.intercept_] else: params [model.coef_] return params def set_model_params(model: LogisticRegression, params: NDArrays) - LogisticRegression: model.coef_ params[0] if model.fit_intercept: model.intercept_ params[1] return modelClientApp把 Message 中的 ArrayRecord 接入模型Flower 与 scikit-learn 对接时最主要的改动集中在参数序列化格式的转换上ClientApp从Message中收到的模型参数是ArrayRecord需要先转成 numpy ndarray 再写回模型训练结束后再把更新后的 ndarray 打包回ArrayRecord随Message返回。这些转换可以直接使用ArrayRecord内置的方法完成app.train() def train(msg: Message, context: Context): # Create LogisticRegression Model penalty context.run_config[penalty] # Create LogisticRegression Model model create_log_reg_and_instantiate_parameters(penalty) # Apply received parameters ndarrays msg.content[arrays].to_numpy_ndarrays() set_model_params(model, ndarrays) # Train the model ... # Extract the updated model parameters with auxhiliary function ndarrays get_model_params(model) # Pack the updated parameters into an ArrayRecord model_record ArrayRecord(ndarrays)从框架源码看framework/py/flwr/app/message/arrayrecord.pyArrayRecord是字符串键 → Array的带类型字典用于存放命名数组模型参数、梯度、嵌入向量等内部行为类似dict[str, Array]可以理解为 PyTorchstate_dict的等价物但保存的是序列化形式的数组它属于RecordDict支持的记录类型之一因此可以放进Message的content或Context的state中。ClientApp提供三个可实现的装饰器方法train用本地数据训练收到的模型、evaluate在验证集上评估收到的模型、query查询执行该ClientApp的节点信息。本教程只用到train和evaluate。train 方法本地训练并回传参数train接收来自ServerApp的Message默认携带两部分内容一个ArrayRecord存放待联邦训练的模型参数默认可通过消息内容中的键arrays获取一个ConfigRecord存放ServerApp下发的配置默认可通过键config获取。train还接收Context用于访问运行配置与节点配置运行配置run config的超参数定义在 Flower App 的pyproject.toml中节点配置node config只能在以 Deployment Runtime 运行 Flower 时设置模拟Simulation模式下不可直接配置。完整的train实现如下与仓库 examples/quickstart-sklearn/sklearnexample/client_app.py 一致app ClientApp() app.train() def train(msg: Message, context: Context): Train the model on local data. # Create LogisticRegression Model penalty context.run_config[penalty] # Create LogisticRegression Model model create_log_reg_and_instantiate_parameters(penalty) # Apply received parameters ndarrays msg.content[arrays].to_numpy_ndarrays() set_model_params(model, ndarrays) # Load the data partition_id context.node_config[partition-id] num_partitions context.node_config[num-partitions] X_train, y_train, _, _ load_data(partition_id, num_partitions) # Ignore convergence failure due to low local epochs with warnings.catch_warnings(): warnings.simplefilter(ignore) # Train the model on local data model.fit(X_train, y_train) # Lets compute train loss y_train_pred_proba model.predict_proba(X_train) train_logloss log_loss(y_train, y_train_pred_proba, labelsUNIQUE_LABELS) accuracy model.score(X_train, y_train) # Construct and return reply Message ndarrays get_model_params(model) model_record ArrayRecord(ndarrays) metrics { num-examples: len(X_train), train_logloss: train_logloss, train_accuracy: accuracy, } metric_record MetricRecord(metrics) content RecordDict({arrays: model_record, metrics: metric_record}) return Message(contentcontent, reply_tomsg)值得注意的细节因为max_iter1会导致收敛警告训练时用warnings.catch_warnings()静默掉相关告警训练完成后除回传模型参数外还计算train_logloss对数损失与train_accuracy连同样本数num-examples一起放进MetricRecord供服务端聚合返回的Message使用reply_tomsg指向请求消息形成完整的请求-响应闭环UNIQUE_LABELS [0, 1, 2]用于log_loss的labels参数确保概率矩阵与标签对齐。evaluate 方法本地评估app.evaluate与train结构镜像但只加载测试集X_test, y_test评估收到的模型返回MetricRecord中的评估损失与准确率不包含模型权重——因为评估过程不会修改模型app.evaluate() def evaluate(msg: Message, context: Context): Evaluate the model on local data. penalty context.run_config[penalty] model create_log_reg_and_instantiate_parameters(penalty) ndarrays msg.content[arrays].to_numpy_ndarrays() set_model_params(model, ndarrays) _, _, X_test, y_test load_data(partition_id, num_partitions) y_test_pred_proba model.predict_proba(X_test) accuracy model.score(X_test, y_test) loss log_loss(y_test, y_test_pred_proba, labelsUNIQUE_LABELS) metrics { num-examples: len(X_test), test_logloss: loss, accuracy: accuracy, } metric_record MetricRecord(metrics) content RecordDict({metrics: metric_record}) return Message(contentcontent, reply_tomsg)ServerApp用 FedAvg 编排联邦训练服务端通过ServerApp的app.main()方法构建。main接收两个参数Grid对象用于与运行ClientApp的节点交互把节点拉入一轮 train/evaluate/query 等联邦流程Context对象提供对运行配置的访问。本示例使用FedAvgFederated Averaging策略。从框架源码看framework/py/flwr/serverapp/strategy/fedavg.pyFedAvg基于论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》arXiv:1602.05629实现fraction_train与fraction_evaluate的默认值均为 1.0min_train_nodes、min_evaluate_nodes、min_available_nodes的默认值均为 2——即采样多少比例的节点参与训练/评估由这两个比例参数控制。随后调用策略的start()方法启动执行需要传入Grid对象一个携带随机初始化模型的ArrayRecord作为待联邦训练的全局模型包含训练超参数、需要下发给客户端的ConfigRecord策略在发送前还会把当前轮数写入该配置num_rounds参数指定执行多少轮FedAvg。完整实现如下与仓库 examples/quickstart-sklearn/sklearnexample/server_app.py 一致app ServerApp() app.main() def main(grid: Grid, context: Context) - None: Main entry point for the ServerApp. # Read run config num_rounds: int context.run_config[num-server-rounds] # Create LogisticRegression Model penalty context.run_config[penalty] model create_log_reg_and_instantiate_parameters(penalty) # Construct ArrayRecord representation arrays ArrayRecord(get_model_params(model)) # Initialize FedAvg strategy strategy FedAvg(fraction_train1.0, fraction_evaluate1.0) # Start strategy, run FedAvg for num_rounds result strategy.start( gridgrid, initial_arraysarrays, num_roundsnum_rounds, ) if context.run_config[save-model]: # Save final model parameters print(\nSaving final model to disk...) ndarrays result.arrays.to_numpy_ndarrays() set_model_params(model, ndarrays) joblib.dump(model, logreg_model.pkl)这段代码把整个联邦流程讲得很清楚从运行配置读取轮数num-server-rounds与正则化类型penalty在服务端创建一个逻辑回归模型用get_model_params()取出其参数并包装成ArrayRecord作为全局初始模型构造FedAvg策略本例中fraction_train1.0, fraction_evaluate1.0即全部节点参与每轮训练与评估strategy.start()驱动多轮联邦每轮向节点下发全局参数、收集本地更新、按 FedAvg 聚合出新全局模型循环至num_rounds轮结束若save-model为true把最终聚合参数写回模型并用joblib.dump保存为logreg_model.pkl。总结与延伸至此你已经在 Iris 数据集上用 Flower scikit-learn 跑通了一个完整的联邦学习系统。整个流程的核心要点可以归纳为脚手架flwr new flwrlabs/quickstart-sklearn一键生成可运行的 Flower Appflwr run . --stream在本地模拟联邦环境数据Flower Datasets 的FederatedDatasetIidPartitioner完成数据下载与分区每个客户端只看到自己的分区模型scikit-learn 的LogisticRegression配合warm_startTrue、max_iter1把fit变成联邦中的本地 epoch参数转换ArrayRecord与 numpy ndarray 之间的双向转换是 scikit-learn 模型接入 Flower Message API 的关键策略编排ServerApp用FedAvg聚合客户端更新--run-config可在不改代码的情况下调整轮数、正则化等超参数。如果要进一步深入官方文档建议学习如何配置与运行更大规模的 Flower 模拟Simulation可查阅 how-to-run-simulations 指南同样的 App 代码无需修改即可切换到 Deployment Runtime 运行真实的多机联邦并可进一步配置 TLS 安全通信与 SuperNode 认证如果想看到另一个基于 scikit-learn 的 Flower App 示例可以参考quickstart-sklearn-tabular的源码对于 tabular 类任务的更丰富分区与预处理方式Flower Datasets 提供的其他 partitioner 都值得一试。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门