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的新目录,其中包含运行"两节点联邦"所需的全部文件。默认情况下,生成的应用带有一个本地模拟 profile:flwr 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.toml):
dependencies = [ "flwr[simulation]>=1.36.0", "flwr-datasets[vision]>=0.6.1", "scikit-learn>=1.6.1", ]flwr[simulation]:Flower 框架本体,simulationextra 提供本地模拟所需的运行时组件;flwr-datasets:负责数据集下载、分区与预处理(本示例使用其IidPartitioner与FederatedDataset);scikit-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-rounds=5 local-epochs=2"--run-config后面以空格分隔的key=value会被注入运行配置,ClientApp与ServerApp通过context.run_config[...]读取。README 中还给出了覆盖字符串型参数的写法示例(examples/quickstart-sklearn/README.md):
flwr run . --run-config penalty="'l1'" --stream该项目pyproject.toml中预置的配置项及其默认值如下:
| 配置键 | 默认值 | 说明 |
|---|---|---|
penalty | "l2" | LogisticRegression的正则化类型,可覆盖为"l1"等 |
num-server-rounds | 25 | FedAvg 联邦训练的轮数 |
min-available-clients | 2 | 系统中最少可用的客户端数量 |
save-model | false | 是否在训练结束后把最终模型保存到本地磁盘 |
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_partitions=num_partitions) fds = FederatedDataset(dataset="hitorilabs/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(dataset="hitorilabs/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( penalty=penalty, max_iter=1, # local epoch warm_start=True, # prevent refreshing weights when fitting, solver="saga", ) # Setting initial parameters, akin to model.compile for keras models set_initial_params(model, n_features=len(FEATURES), n_classes=len(UNIQUE_LABELS)) return model几个参数的含义与联邦场景的适配:
max_iter=1:把 scikit-learn 的优化迭代次数当作"本地 epoch"来用,每次fit只做一轮优化;warm_start=True:再次fit时沿用上一次的权重而不是重新初始化,这是联邦学习中"在服务端下发的参数基础上继续训练"的前提;solver="saga":适合中小规模数据且支持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.py),ArrayRecord是"字符串键 → 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) # Let's compute train loss y_train_pred_proba = model.predict_proba(X_train) train_logloss = log_loss(y_train, y_train_pred_proba, labels=UNIQUE_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(content=content, reply_to=msg)值得注意的细节:
- 因为
max_iter=1会导致收敛警告,训练时用warnings.catch_warnings()静默掉相关告警; - 训练完成后除回传模型参数外,还计算
train_logloss(对数损失)与train_accuracy,连同样本数num-examples一起放进MetricRecord,供服务端聚合; - 返回的
Message使用reply_to=msg指向请求消息,形成完整的请求-响应闭环; 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, labels=UNIQUE_LABELS) metrics = { "num-examples": len(X_test), "test_logloss": loss, "accuracy": accuracy, } metric_record = MetricRecord(metrics) content = RecordDict({"metrics": metric_record}) return Message(content=content, reply_to=msg)ServerApp:用 FedAvg 编排联邦训练
服务端通过ServerApp的@app.main()方法构建。main接收两个参数:
Grid对象:用于与运行ClientApp的节点交互,把节点拉入一轮 train/evaluate/query 等联邦流程;Context对象:提供对运行配置的访问。
本示例使用FedAvg(Federated Averaging)策略。从框架源码看(framework/py/flwr/serverapp/strategy/fedavg.py),FedAvg基于论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》(arXiv:1602.05629)实现,fraction_train与fraction_evaluate的默认值均为 1.0,min_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_train=1.0, fraction_evaluate=1.0) # Start strategy, run FedAvg for `num_rounds` result = strategy.start( grid=grid, initial_arrays=arrays, num_rounds=num_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_train=1.0, fraction_evaluate=1.0,即全部节点参与每轮训练与评估); strategy.start()驱动多轮联邦:每轮向节点下发全局参数、收集本地更新、按 FedAvg 聚合出新全局模型,循环至num_rounds轮结束;- 若
save-model为true,把最终聚合参数写回模型并用joblib.dump保存为logreg_model.pkl。
总结与延伸
至此,你已经在 Iris 数据集上用 Flower + scikit-learn 跑通了一个完整的联邦学习系统。整个流程的核心要点可以归纳为:
- 脚手架:
flwr new @flwrlabs/quickstart-sklearn一键生成可运行的 Flower App,flwr run . --stream在本地模拟联邦环境; - 数据:Flower Datasets 的
FederatedDataset+IidPartitioner完成数据下载与分区,每个客户端只看到自己的分区; - 模型:scikit-learn 的
LogisticRegression配合warm_start=True、max_iter=1,把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),仅供参考