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.toml):
flwr[simulation]>=1.36.0:Flower 主框架,并附带 Simulation Engine 所需的依赖;flwr-datasets[vision]>=0.6.1:Flower Datasets,用于下载与切分 CIFAR-10;tensorflow==2.20.0:TensorFlow/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:39091;FedAvg 共执行 3 轮;每轮训练会从 2 个节点中采样 2 个(fraction_train=1.0);聚合训练/评估指标以MetricRecord形式返回;最终服务端汇总的评估指标为空字典{}(因为本示例的评估指标由各节点在evaluate中返回,服务端并未额外评估)。
通过 run-config 覆盖超参数
flwr run支持覆盖pyproject.toml中[tool.flwr.app.config]段定义的参数:
# 覆盖部分参数 $ flwr run . --run-config "num-server-rounds=5 batch-size=16"该示例的默认配置如下(pyproject.toml):
| 配置键 | 默认值 | 说明 |
|---|---|---|
num-server-rounds | 3 | FedAvg 联邦训练的轮数 |
local-epochs | 1 | 每个客户端本地训练的 epoch 数 |
batch-size | 32 | 本地训练的批大小 |
learning-rate | 0.005 | Adam 优化器的学习率 |
fraction-train | 1.0 | 每轮参与训练的节点比例 |
verbose | false | 是否打印训练过程详细日志 |
save-model | false | 是否在训练结束后保存最终模型 |
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_partitions=num_partitions) fds = FederatedDataset( dataset="uoft-cs/cifar10", partitioners={"train": partitioner}, ) partition = fds.load_partition(partition_id, "train") partition.set_format("numpy") # 在每个节点上把数据再划分:80% 训练,20% 测试 partition = partition.train_test_split(test_size=0.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(type="numpy", 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), activation="relu"), layers.MaxPooling2D(pool_size=(2, 2)), layers.Conv2D(64, kernel_size=(3, 3), activation="relu"), layers.MaxPooling2D(pool_size=(2, 2)), layers.Flatten(), layers.Dropout(0.5), layers.Dense(10, activation="softmax"), ] ) optimizer = keras.optimizers.Adam(learning_rate) model.compile( optimizer=optimizer, loss="sparse_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, epochs=epochs, batch_size=batch_size, verbose=verbose, ) # 提取训练指标 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(content=content, reply_to=msg)完整的可运行版本见 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.py),ArrayRecord是一个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, verbose=0) # 打包评估指标并作为消息返回 metrics = { "eval_acc": eval_acc, "eval_loss": eval_loss, "num-examples": len(x_test), } content = RecordDict({"metrics": MetricRecord(metrics)}) return Message(content=content, reply_to=msg)注意这里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_train=fraction_train, ) result = strategy.start( grid=grid, initial_arrays=arrays, num_rounds=num_rounds, ) if context.run_config["save-model"]: # 保存最终模型 ndarrays = result.arrays.to_numpy_ndarrays() final_model_name = "final_model.keras" print(f"Saving 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_train | 1.0 | 训练时采样的节点比例;若min_train_nodes大于fraction_train × 总连接节点数,仍会采样到min_train_nodes个节点 |
fraction_evaluate | 1.0 | 验证时采样的节点比例 |
min_train_nodes | 2 | 训练阶段最少参与的节点数 |
min_evaluate_nodes | 2 | 验证阶段最少参与的节点数 |
min_available_nodes | 2 | 系统中最少的可用节点总数 |
weighted_by_key | "num-examples" | 计算加权平均时使用的指标键,对应客户端返回的样本数 |
arrayrecord_key | "arrays" | 构造 Message 时存放ArrayRecord的键 |
configrecord_key | "config" | 构造 Message 时存放ConfigRecord的键 |
train_metrics_aggr_fn/evaluate_metrics_aggr_fn | None | 自定义训练/评估指标聚合函数,默认使用按weighted_by_key加权的aggregate_metricrecords |
这正是客户端metrics中必须包含"num-examples"的原因:服务端FedAvg默认按它对各节点的权重数组与指标做加权平均。
小结与进阶方向
至此,你已经成功构建并运行了第一个联邦学习系统:flwr new生成项目骨架,task.py完成数据加载与模型定义,client_app.py实现train/evaluate两个方法完成本地训练与评估并打包Message,server_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),仅供参考