Flower + PyTorch 联邦学习快速上手:CIFAR-10 图像分类实战指南(Quickstart-Pytorch 深度解析)
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
本篇技术指南以 Flower 仓库中的quickstart-pytorch示例为对象,讲解如何用 PyTorch 构建一个完整的联邦图像分类应用:从安装 Flower、拉取应用模板、安装依赖,到以 Simulation(模拟)与 Deployment(部署)两种模式运行,并深入源码剖析模型、数据分区、ClientApp 与 ServerApp 的协作方式,以及 FedAvg 策略在框架层的实现原理。读完本篇,你将掌握基于 Flower 的 PyTorch 联邦学习应用的完整开发与运行流程,能够独立将任意 PyTorch 模型改造成联邦训练应用。
示例概览:用联邦学习训练一个 CNN 识别 CIFAR-10
quickstart-pytorch是 Flower 的入门级示例,使用 PyTorch 作为深度学习框架。它并不要求读者具备深入的 PyTorch 知识即可运行,但对理解如何将 Flower 适配到自己的使用场景会很有帮助。该示例使用 Flower Datasets 完成 CIFAR-10 数据集的下载、分区与预处理,整个联邦流程由 Flower 的 ServerApp 与 ClientApp 协作完成。
从仓库结构看,该示例位于 examples/quickstart-pytorch,核心代码全部集中在pytorchexample包内,仅包含 4 个 Python 文件(task.py、client_app.py、server_app.py、__init__.py)以及一个pyproject.toml配置文件,结构极其精简,非常适合作为学习 Flower 的起点。
搭建项目:安装 Flower 并获取应用模板
安装 Flower
首先安装 Flower 框架本体:
pip install flwr拉取应用模板
使用 Flower 官方的 CLI 工具拉取quickstart-pytorch应用模板:
flwr new @flwrlabs/quickstart-pytorch该命令会在当前目录下创建一个名为quickstart-pytorch的新目录,其结构如下:
quickstart-pytorch ├── pytorchexample │ ├── __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-pytorch 目录一致:pytorchexample/__init__.py仅包含一行包说明,真正的逻辑分布在另外三个文件与pyproject.toml中。各文件职责如下:
| 文件 | 职责 |
|---|---|
task.py | 定义模型结构、数据加载、训练与评估函数,是纯 PyTorch 逻辑所在 |
client_app.py | 定义 ClientApp,注册联邦训练与评估的处理函数 |
server_app.py | 定义 ServerApp,注册服务端主流程与全局评估函数 |
pyproject.toml | 声明项目元信息、依赖以及 Flower 应用的组件入口与运行配置 |
安装依赖与项目包
进入项目目录后,安装pyproject.toml中声明的依赖以及pytorchexample包本身:
pip install -e .pyproject.toml中声明的核心依赖为(见 pyproject.toml):
dependencies = [ "flwr[simulation]>=1.36.0", "flwr-datasets[vision]>=0.6.1", "torch==2.10.0", "torchvision==0.25.0", ]其中flwr[simulation]额外引入模拟引擎所需的运行环境,flwr-datasets[vision]提供联邦数据集分区能力并附带视觉相关的依赖。
运行项目:Simulation 与 Deployment 两种模式
Flower 项目可以在不修改任何代码的前提下,以Simulation(模拟)与Deployment(部署)两种模式运行。对于 Flower 初学者,官方建议优先使用 Simulation 模式,因为它需要手动启动的组件更少;默认情况下flwr run就会使用 Simulation Engine。
使用 Simulation Engine 运行
在项目根目录执行:
# Run with the default federation (CPU only) flwr run . --stream--stream参数会以流式方式实时输出运行日志,便于观察每个联邦轮次的训练与评估进展。该命令默认执行 CPU 上的联邦训练。需要说明的是,如果 ClientApp 能够访问 GPU,示例运行会更快;关于 Simulation 的原理与优化策略(例如supernode数量、ClientApp并行度等),可以查阅框架文档中的 Simulation Engine 相关章节。
你还可以覆盖pyproject.toml中为 ClientApp 和 ServerApp 定义的运行配置,例如:
flwr run . --run-config "num-server-rounds=5 learning-rate=0.05" --stream这条命令将联邦轮次数从默认的 3 轮提升到 5 轮,并将学习率从默认的0.1调整为0.05,其余配置保持不变。
使用 Deployment Engine 运行
Deployment Engine 是 Flower 面向真实生产环境的运行模式,需要分别启动 SuperLink(服务端协调器)与多个 SuperNode(节点),并将 ClientApp 分发到各节点上执行。如果你想在真实或虚拟设备上运行同一个应用,可以参考框架文档中的 Deployment Engine 使用指南。在跑通部署模式后,通常还会进一步配置:
- TLS 安全通信:为联邦网络启用 TLS 加密连接;
- SuperNode 认证:为 SuperNode 接入联邦网络增加身份验证机制。
如果你已经熟悉 Deployment Engine 的工作方式,还可以通过 Docker 容器化运行整个联邦系统。
深入源码:从模型、数据到 ClientApp 与 ServerApp
这一节将结合仓库源码逐步拆解应用的三个核心文件,帮助你理解联邦训练完整的数据流。
task.py:模型、数据分区与训练/评估逻辑
task.py是纯 PyTorch 逻辑所在(见 task.py)。
模型定义:Net是一个简单的卷积神经网络,改编自 PyTorch 官方教程 "PyTorch: A 60 Minute Blitz":
class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv1 = nn.Conv2d(3, 6, 5) self.pool = nn.MaxPool2d(2, 2) self.conv2 = nn.Conv2d(6, 16, 5) self.fc1 = nn.Linear(16 * 5 * 5, 120) self.fc2 = nn.Linear(120, 84) self.fc3 = nn.Linear(84, 10)网络由两个卷积层(Conv2d+MaxPool2d+ ReLU)和三个全连接层组成,输入为 3 通道的 CIFAR-10 图像,输出 10 类概率 logits。
联邦数据加载:load_data函数是联邦数据流的核心:
def load_data(partition_id: int, num_partitions: int, batch_size: int): global fds if fds is None: partitioner = IidPartitioner(num_partitions=num_partitions) fds = FederatedDataset( dataset="uoft-cs/cifar10", partitioners={"train": partitioner}, ) partition = fds.load_partition(partition_id) partition_train_test = partition.train_test_split(test_size=0.2, seed=42) partition_train_test = partition_train_test.with_transform(apply_transforms) trainloader = DataLoader( partition_train_test["train"], batch_size=batch_size, shuffle=True ) testloader = DataLoader(partition_train_test["test"], batch_size=batch_size) return trainloader, testloader关键点包括:
- 使用
flwr_datasets.partitioner.IidPartitioner将uoft-cs/cifar10数据集划分为num_partitions个独立分区,每个 ClientApp 通过partition_id取用属于自己的那份数据,实现 IID(独立同分布)联邦数据划分; FederatedDataset通过模块级全局变量fds缓存,保证每个进程中只初始化一次,避免重复下载与分区;- 每个节点拿到自己的分区后,再按 80%/20% 切分为本地训练集与本地测试集(
train_test_split(test_size=0.2, seed=42)); apply_transforms将图像转为 Tensor 并做均值 0.5、标准差 0.5 的归一化(Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)))。
训练与评估:train函数使用交叉熵损失(CrossEntropyLoss)与带动量 0.9 的 SGD 优化器进行本地多轮训练,并返回平均训练损失;test函数在测试集上计算损失与准确率。此外,load_centralized_dataset会加载完整的 CIFAR-10 官方测试集,供服务端做全局评估。
client_app.py:ClientApp 如何完成一次联邦参与
ClientApp 定义了客户端侧的联邦行为(见 client_app.py),通过装饰器注册两个处理函数:
训练处理函数@app.train(),完整流程为:
- 用
Message中携带的全局模型参数初始化本地模型:model.load_state_dict(msg.content["arrays"].to_torch_state_dict()); - 自动选择设备:
torch.device("cuda:0" if torch.cuda.is_available() else "cpu"); - 从
context.node_config读取节点身份配置partition-id与num-partitions,从context.run_config读取运行配置batch-size,据此加载本地数据; - 以
context.run_config["local-epochs"]为本地训练轮数、msg.content["config"]["lr"]为学习率执行训练; - 将更新后的模型参数封装为
ArrayRecord,连同train_loss、num-examples指标封装为MetricRecord,组合成RecordDict作为回复Message返回。
评估处理函数@app.evaluate()与训练流程类似:加载收到的模型参数,用本地测试集(valloader)评估,返回eval_loss、eval_acc、num-examples三个指标。
这里体现了 Flower 新一代基于Message/Record的通信模型:ArrayRecord承载张量参数,MetricRecord承载标量指标,二者统一放进RecordDict随Message在网络中传输,客户端与服务端无需关心底层序列化细节。
server_app.py:ServerApp 与 FedAvg 策略编排
ServerApp 定义了服务端侧的联邦编排逻辑(见 server_app.py),核心在@app.main()主函数中:
strategy = FedAvg(fraction_evaluate=fraction_evaluate) result = strategy.start( grid=grid, initial_arrays=arrays, train_config=ConfigRecord({"lr": lr}), num_rounds=num_rounds, evaluate_fn=global_evaluate, )要点如下:
- 服务端首先初始化一个全局模型
Net(),将其参数封装为ArrayRecord作为初始权重; - 实例化
FedAvg策略(联邦平均,基于经典论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》),fraction_evaluate=1.0表示每一轮都让所有节点参与评估; strategy.start()将全局权重、训练配置(学习率)与轮次数交给策略,由策略在每一轮采样节点、下发参数、收集更新并做加权聚合,evaluate_fn=global_evaluate指定每轮结束后在服务端中央测试集上评估全局模型;- 若
context.run_config["save-model"]为真,训练结束后会将最终模型参数保存为final_model.pt,便于后续推理与导出。
其中global_evaluate函数加载全部 CIFAR-10 测试集,返回MetricRecord({"accuracy": test_acc, "loss": test_loss}),这些全局指标会随每轮训练输出,是观察联邦收敛情况的重要依据。
运行配置参数全解
pyproject.toml的[tool.flwr.app.config]段定义了应用的默认运行配置(见 pyproject.toml):
[tool.flwr.app.config] num-server-rounds = 3 fraction-evaluate = 1.0 local-epochs = 1 learning-rate = 0.1 batch-size = 32 save-model = false| 参数 | 默认值 | 作用 |
|---|---|---|
num-server-rounds | 3 | 联邦聚合轮次数,即全局模型被下发-更新-聚合的次数 |
fraction-evaluate | 1.0 | 每轮参与评估的节点比例,1.0表示全部节点参与 |
local-epochs | 1 | 每个 ClientApp 本地训练的 epoch 数 |
learning-rate | 0.1 | 本地训练使用的 SGD 学习率,由服务端通过train_config下发 |
batch-size | 32 | 本地 DataLoader 的批大小 |
save-model | false | 训练结束后是否将最终模型保存为final_model.pt |
这些配置均可通过flwr run . --run-config "key=value ..."在命令行覆盖,例如num-server-rounds=5 learning-rate=0.05。此外[tool.flwr.app.components]段指定了 ServerApp 与 ClientApp 的导入路径(pytorchexample.server_app:app与pytorchexample.client_app:app),[tool.flwr.app]段声明了发布者、FAB 格式版本与目标 Flower 版本(flwr-version-target = "1.37.0"),是flwr run定位应用入口的关键元数据。
FedAvg 策略的框架层实现
示例中使用的FedAvg来自框架的serverapp模块(见 fedavg.py),其构造参数如下:
def __init__( self, fraction_train: float = 1.0, fraction_evaluate: float = 1.0, min_train_nodes: int = 2, min_evaluate_nodes: int = 2, min_available_nodes: int = 2, weighted_by_key: str = "num-examples", arrayrecord_key: str = "arrays", configrecord_key: str = "config", train_metrics_aggr_fn=None, evaluate_metrics_aggr_fn=None, ) -> None:核心机制可以从源码确认:
- 节点采样:
configure_train中按fraction_train计算参与训练节点数num_nodes = int(len(list(grid.get_node_ids())) * self.fraction_train),再与min_train_nodes取较大值,通过sample_nodes在min_available_nodes约束下完成采样; - 参数聚合:训练更新通过
aggregate_arrayrecords聚合,指标通过aggregate_metricrecords聚合,二者都默认以weighted_by_key(默认"num-examples")作为权重进行加权平均——这正是客户端回复中必须携带num-examples指标的原因; - 轮次注入:
configure_train会在下发的配置中注入config["server-round"] = server_round,客户端可据此感知当前轮次; - 消息构造:
_construct_messages为每个被采样节点构造一条携带同一份RecordDict(包含arrays与config)的Message,通过Grid广播。
从源码结构还可以看到,flwr.serverapp.strategy下还提供fedavgm等扩展策略,示例默认使用最基础的FedAvg,后续可平滑替换为更复杂的聚合算法。
总结与延伸
quickstart-pytorch以最小的代码量完整演示了 Flower 联邦学习应用的四个要素:模型与数据(task.py)、客户端逻辑(client_app.py)、服务端编排(server_app.py)、运行配置(pyproject.toml)。你可以沿以下方向继续深入:
- 将该示例替换为自定义模型:只需修改
Net并保证train/test函数的输入输出契约不变; - 将
IidPartitioner替换为flwr_datasets中的其他分区器(如非 IID 的 Dirichlet 分区),即可研究数据异构场景; - 将
FedAvg替换为框架内置的其他策略,或参考仓库 baselines 目录中基于本示例演进的各类基线实现; - 生产化部署时,参考框架文档中关于 Deployment Engine、TLS 连接与 SuperNode 认证的章节(文档源码位于 framework/docs/source)。
本示例对应仓库路径为 examples/quickstart-pytorch,其pyproject.toml声明了flwr>=1.36.0的版本下限,所有命令与配置均以上述仓库实际内容为准。
【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考