ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Flower + TensorFlow 联邦学习实战:用 Quickstart TensorFlow 在 CIFAR-10 上训练 CNN

Flower + TensorFlow 联邦学习实战:用 Quickstart TensorFlow 在 CIFAR-10 上训练 CNN Flower TensorFlow 联邦学习实战用 Quickstart TensorFlow 在 CIFAR-10 上训练 CNN【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower本教程基于 Flower 框架官方文档《Quickstart TensorFlow》编写讲解如何用 Flower 与 TensorFlow/Keras 在 CIFAR-10 数据集上构建并运行一个两节点联邦学习系统从flwr new脚手架生成项目到编写数据加载、模型定义、ClientApp与ServerApp再到以flwr run启动 FedAvg 联邦训练。读完本文你将掌握 Flower 的核心消息与记录类型Message、ArrayRecord、MetricRecord在 TensorFlow 场景下的完整使用方式并能独立把现有 Keras 模型改造成可联邦训练的 Flower App。教程概览与运行环境准备本教程推荐在独立虚拟环境中完成。首先安装 Flower# 在一个全新的 Python 环境中 $ pip install flwr然后使用flwr new命令从 Flower Labs 拉取现成的快速开始模板它会生成一个完整可运行的 Flower TensorFlow 项目该项目使用 FedAvg 策略组织两个节点的联邦训练$ flwr new flwrlabs/quickstart-tensorflow命令执行后当前目录下会出现一个名为quickstart-tensorflow的新目录其结构如下quickstart-tensorflow ├── tfexample │ ├── __init__.py │ ├── client_app.py # 定义你的 ClientApp │ ├── server_app.py # 定义你的 ServerApp │ └── task.py # 定义模型、训练与数据加载 ├── pyproject.toml # 项目元数据依赖与配置 └── README.md在本仓库中与模板等价的可运行完整实现位于 examples/quickstart-tensorflow其pyproject.toml声明了核心依赖pyproject.tomlflwr[simulation]1.36.0Flower 主框架并附带 Simulation Engine 所需的依赖flwr-datasets[vision]0.6.1Flower Datasets用于下载与切分 CIFAR-10tensorflow2.20.0TensorFlow/Keras 深度学习框架。说明本教程默认以本地 Simulation 模式运行。flwr run会向本机托管的 SuperLink 提交一次运行由 Flower Simulation Runtime 负责调度无需手工启动多个进程同一份代码也可切换到 Deployment 模式在真实节点上运行。运行联邦训练并理解输出日志进入项目目录后用下面的命令启动联邦训练$ cd quickstart-tensorflow # 使用默认参数运行并流式输出日志 $ flwr run . --stream这里的--stream表示实时流式查看日志不带该参数时flwr run .只会提交运行、打印 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_loss: 2.0013, train_acc: 0.2624} INFO : configure_evaluate: Sampled 2 nodes (out of 2) INFO : aggregate_evaluate: Received 2 results and 0 failures INFO : └── Aggregated MetricRecord: {eval_acc: 0.1216, eval_loss: 2.2686} INFO : [ROUND 2/3] INFO : ... INFO : [ROUND 3/3] INFO : ... INFO : Strategy execution finished in 16.60s INFO : Final results: INFO : ServerApp-side Evaluate Metrics: INFO : {}从日志可以读出几个关键信息SuperLink 监听在127.0.0.1:39091FedAvg 共执行 3 轮每轮训练会从 2 个节点中采样 2 个fraction_train1.0聚合训练/评估指标以MetricRecord形式返回最终服务端汇总的评估指标为空字典{}因为本示例的评估指标由各节点在evaluate中返回服务端并未额外评估。通过 run-config 覆盖超参数flwr run支持覆盖pyproject.toml中[tool.flwr.app.config]段定义的参数# 覆盖部分参数 $ flwr run . --run-config num-server-rounds5 batch-size16该示例的默认配置如下pyproject.toml配置键默认值说明num-server-rounds3FedAvg 联邦训练的轮数local-epochs1每个客户端本地训练的 epoch 数batch-size32本地训练的批大小learning-rate0.005Adam 优化器的学习率fraction-train1.0每轮参与训练的节点比例verbosefalse是否打印训练过程详细日志save-modelfalse是否在训练结束后保存最终模型pyproject.toml中还有两段与 App 生命周期直接相关的配置[tool.flwr.app.components]指定了服务端与客户端的入口对象tfexample.server_app:app与tfexample.client_app:app[tool.flwr.app]记录了发布者flwrlabs、FAB 格式版本与应用目标 Flower 版本flwr-version-target 1.37.0。The Data用 Flower Datasets 加载并切分 CIFAR-10本教程使用 Flower Datasets 下载并切分 CIFAR-10 数据集。它借助IidPartitioner把训练集切分为num_partitions份IID即独立同分布切分每个ClientApp在运行时按自己的partition-id取出对应分片partitioner IidPartitioner(num_partitionsnum_partitions) fds FederatedDataset( datasetuoft-cs/cifar10, partitioners{train: partitioner}, ) partition fds.load_partition(partition_id, train) partition.set_format(numpy) # 在每个节点上把数据再划分80% 训练20% 测试 partition partition.train_test_split(test_size0.2) x_train, y_train partition[train][img] / 255.0, partition[train][label] x_test, y_test partition[test][img] / 255.0, partition[test][label]仓库中的完整实现位于 examples/quickstart-tensorflow/tfexample/task.py它额外做了三件事使用模块级全局变量fds缓存FederatedDataset避免每个客户端重复下载数据集通过partition.set_format(typenumpy, columns[img, label])显式把图片与标签转为 NumPy 格式将像素值除以 255.0 归一化到[0, 1]并把图片转为float32以匹配 Keras 的输入要求。如果你需要非 IID 的数据分布更贴近真实联邦场景可以在 Flower Datasets 的 partitioner 集合中选择其他实现如 Dirichlet 分布切分器只需替换IidPartitioner即可其余数据加载流程保持不变。The Model定义 CIFAR-10 卷积神经网络接下来是模型部分。教程定义了一个简单的 CNN你可以自由替换为更复杂的网络结构def load_model(learning_rate: float 0.001): # 为 CIFAR-10 定义一个简单 CNN并使用 Adam 优化器 model keras.Sequential( [ keras.Input(shape(32, 32, 3)), layers.Conv2D(32, kernel_size(3, 3), activationrelu), layers.MaxPooling2D(pool_size(2, 2)), layers.Conv2D(64, kernel_size(3, 3), activationrelu), layers.MaxPooling2D(pool_size(2, 2)), layers.Flatten(), layers.Dropout(0.5), layers.Dense(10, activationsoftmax), ] ) optimizer keras.optimizers.Adam(learning_rate) model.compile( optimizeroptimizer, losssparse_categorical_crossentropy, metrics[accuracy], ) return model这个模型接受(32, 32, 3)的 CIFAR-10 彩色图片输入经过两层「卷积 最大池化」提取特征再接Flatten与Dropout(0.5)防止过拟合最后由 10 个神经元的softmax层输出分类概率。注意两点其一load_model(learning_rate...)的学习率来自运行配置而非硬编码其二sparse_categorical_crossentropy配合整数标签使用因此数据加载时无需做 one-hot 编码。The ClientApp把 Keras 权重装进 Message 往返传输在 Flower 中客户端与服务器之间的一切交互都通过Message完成。要在 TensorFlow 场景下使用 Flower核心改动是把Message中收到的ArrayRecord转成 NumPy 数组供 Keras 的set_weights()使用训练结束后再把get_weights()得到的数组重新打包进ArrayRecord并随Message返回app.train() def train(msg: Message, context: Context): # 加载模型 model load_model(context.run_config[learning-rate]) # 从 Message 中取出 ArrayRecord 并转换为 numpy ndarrays model.set_weights(msg.content[arrays].to_numpy_ndarrays()) # 训练模型 ... # 将模型权重打包进 ArrayRecord model_record ArrayRecord(model.get_weights())ClientApp提供三个核心方法train、evaluate、query分别用于不同目的train用本地数据训练收到的模型evaluate在验证集上评估收到的模型性能query查询执行ClientApp的节点信息。本教程只使用train和evaluate。train方法收到的Message默认携带两类内容一个ArrayRecord存放待联邦训练的模型参数数组默认通过键arrays从消息内容中取出一个ConfigRecord存放ServerApp下发的配置默认通过键config取出。此外train还接收Context参数它提供 run 级配置与 node 级配置run 配置超参数定义在pyproject.toml中node 配置如partition-id、num-partitions只能在 Deployment Runtime 下设置Simulation 模式下由仿真运行时自动注入、不可直接配置。train 方法的完整实现# Flower ClientApp app ClientApp() app.train() def train(msg: Message, context: Context): 使用本地数据训练模型。 # 重置本地 Tensorflow 状态 keras.backend.clear_session() # 加载数据 partition_id context.node_config[partition-id] num_partitions context.node_config[num-partitions] x_train, y_train, _, _ load_data(partition_id, num_partitions) # 加载模型 model load_model(context.run_config[learning-rate]) model.set_weights(msg.content[arrays].to_numpy_ndarrays()) epochs context.run_config[local-epochs] batch_size context.run_config[batch-size] verbose context.run_config.get(verbose) # 训练模型 history model.fit( x_train, y_train, epochsepochs, batch_sizebatch_size, verboseverbose, ) # 提取训练指标 train_loss history.history[loss][-1] if loss in history.history else None train_acc ( history.history[accuracy][-1] if accuracy in history.history else None ) # 打包模型权重与指标并作为消息返回 model_record ArrayRecord(model.get_weights()) metrics {num-examples: len(x_train)} if train_loss is not None: metrics[train_loss] train_loss if train_acc is not None: metrics[train_acc] train_acc content RecordDict({arrays: model_record, metrics: MetricRecord(metrics)}) return Message(contentcontent, reply_tomsg)完整的可运行版本见 examples/quickstart-tensorflow/tfexample/client_app.py。这里有几点值得深挖keras.backend.clear_session()用于重置本地 TensorFlow 图状态避免多轮训练之间产生残留权重经model.get_weights()得到list[np.ndarray]直接传入ArrayRecord(...)构造器即可完成打包。从源码看framework/py/flwr/app/message/arrayrecord.pyArrayRecord是一个str - Array的 TypedDict官方将其类比为 PyTorch 的state_dict但内部以序列化形式持有数组支持从空容器、dict[str, Array]、NumPy 数组列表或 PyTorchstate_dict四种方式初始化——本示例使用第三种方式反向转换to_numpy_ndarrays()见 framework/py/flwr/app/message/arrayrecord.py把ArrayRecord还原成 NumPy 数组列表供model.set_weights()使用返回消息中的metrics是一个MetricRecord其中num-examples是 FedAvg 按样本数加权聚合的关键权重键详见下文服务端实现。evaluate 方法的实现app.evaluate()与train几乎相同只有两点差异(1) 模型不做本地训练而是直接在本地留出的验证集上评估其性能(2) 由于模型未被本地修改回复的Message中不再需要携带模型权重app.evaluate() def evaluate(msg: Message, context: Context): 在本地数据上评估模型。 # 重置本地 Tensorflow 状态 keras.backend.clear_session() # 加载数据 partition_id context.node_config[partition-id] num_partitions context.node_config[num-partitions] _, _, x_test, y_test load_data(partition_id, num_partitions) # 加载模型 model load_model(context.run_config[learning-rate]) model.set_weights(msg.content[arrays].to_numpy_ndarrays()) # 评估模型 eval_loss, eval_acc model.evaluate(x_test, y_test, verbose0) # 打包评估指标并作为消息返回 metrics { eval_acc: eval_acc, eval_loss: eval_loss, num-examples: len(x_test), } content RecordDict({metrics: MetricRecord(metrics)}) return Message(contentcontent, reply_tomsg)注意这里content只包含metrics键、不包含arrays这正是文档所述「不再需要包含模型」的具体体现。The ServerApp用 FedAvg 编排联邦学习轮次服务端通过定义app.main()方法构造ServerApp。该方法接收两个参数Grid对象用于与运行ClientApp的节点交互把它们组织进一轮 train/evaluate/query 等操作Context对象提供运行配置的访问入口。本示例使用 FedAvg 策略其fraction_train从运行配置读取默认值见pyproject.toml。随后调用策略的start方法启动联邦训练向其传入Grid对象、一个携带随机初始化全局模型的ArrayRecord以及联邦轮数num_rounds# 创建 ServerApp app ServerApp() app.main() def main(grid: Grid, context: Context) - None: ServerApp 的主入口。 # 加载配置 num_rounds context.run_config[num-server-rounds] fraction_train context.run_config[fraction-train] # 加载初始模型 model load_model() arrays ArrayRecord(model.get_weights()) # 定义并启动 FedAvg 策略 strategy FedAvg( fraction_trainfraction_train, ) result strategy.start( gridgrid, initial_arraysarrays, num_roundsnum_rounds, ) if context.run_config[save-model]: # 保存最终模型 ndarrays result.arrays.to_numpy_ndarrays() final_model_name final_model.keras print(fSaving final model to disk as {final_model_name}...) model.set_weights(ndarrays) model.save(final_model_name)完整实现见 examples/quickstart-tensorflow/tfexample/server_app.py。start方法返回一个结果对象其中包含联邦学习过程的全部关键信息以ArrayRecord形式存在的最终模型权重以及以MetricRecord形式存在的联邦训练与评估指标。你可以用 Python 的pprint打印这些指标也可以像上面那样用 TensorFlow 的save()方法把最终权重落盘为final_model.keras。从FedAvg源码看framework/py/flwr/serverapp/strategy/fedavg.py该策略基于经典论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》arXiv:1602.05629实现除fraction_train外还支持以下关键参数参数默认值说明fraction_train1.0训练时采样的节点比例若min_train_nodes大于fraction_train × 总连接节点数仍会采样到min_train_nodes个节点fraction_evaluate1.0验证时采样的节点比例min_train_nodes2训练阶段最少参与的节点数min_evaluate_nodes2验证阶段最少参与的节点数min_available_nodes2系统中最少的可用节点总数weighted_by_keynum-examples计算加权平均时使用的指标键对应客户端返回的样本数arrayrecord_keyarrays构造 Message 时存放ArrayRecord的键configrecord_keyconfig构造 Message 时存放ConfigRecord的键train_metrics_aggr_fn/evaluate_metrics_aggr_fnNone自定义训练/评估指标聚合函数默认使用按weighted_by_key加权的aggregate_metricrecords这正是客户端metrics中必须包含num-examples的原因服务端FedAvg默认按它对各节点的权重数组与指标做加权平均。小结与进阶方向至此你已经成功构建并运行了第一个联邦学习系统flwr new生成项目骨架task.py完成数据加载与模型定义client_app.py实现train/evaluate两个方法完成本地训练与评估并打包Messageserver_app.py用FedAvg聚合多节点权重。整个过程只用了flwr run . --stream一条命令Simulation Engine 自动完成了本地 SuperLink 启动、两节点采样与多轮聚合。如果希望进一步深入可以从两个方向继续探索仿真调优查看 Simulation 配置与运行指南若需原文路径可参考仓库framework/docs/source/下的对应文档学习如何配置并优化 Flower 仿真部署落地将同一份代码切换到 Deployment Engine 在真实设备上运行并可进一步配置 TLS 加密通信与 SuperNode 认证若要替换策略可在server_app.py中把FedAvg换成FedAvgM等其他内置策略实现位于 framework/py/flwr/serverapp/strategy。结合本仓库的 examples/quickstart-tensorflow 完整源码与 教程文档你可以按需修改网络结构、替换数据切分器或调整超参数把这一套「数据加载 → 本地训练 → 权重打包 → 服务端聚合」的范式快速迁移到你自己的 TensorFlow 联邦学习任务中。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进