资讯动态

Flower + PyTorch 联邦学习快速上手:CIFAR-10 图像分类实战指南(Quickstart-Pytorch 深度解析)

发布时间:2026/9/17 13:18:33 来源:尧图企业网站定制
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-10quickstart-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.tomldependencies [ flwr[simulation]1.36.0, flwr-datasets[vision]0.6.1, torch2.10.0, torchvision0.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-rounds5 learning-rate0.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 Blitzclass 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)网络由两个卷积层Conv2dMaxPool2d 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_partitionsnum_partitions) fds FederatedDataset( datasetuoft-cs/cifar10, partitioners{train: partitioner}, ) partition fds.load_partition(partition_id) partition_train_test partition.train_test_split(test_size0.2, seed42) partition_train_test partition_train_test.with_transform(apply_transforms) trainloader DataLoader( partition_train_test[train], batch_sizebatch_size, shuffleTrue ) testloader DataLoader(partition_train_test[test], batch_sizebatch_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_size0.2, seed42)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.pyClientApp 如何完成一次联邦参与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.pyServerApp 与 FedAvg 策略编排ServerApp 定义了服务端侧的联邦编排逻辑见 server_app.py核心在app.main()主函数中strategy FedAvg(fraction_evaluatefraction_evaluate) result strategy.start( gridgrid, initial_arraysarrays, train_configConfigRecord({lr: lr}), num_roundsnum_rounds, evaluate_fnglobal_evaluate, )要点如下服务端首先初始化一个全局模型Net()将其参数封装为ArrayRecord作为初始权重实例化FedAvg策略联邦平均基于经典论文《Communication-Efficient Learning of Deep Networks from Decentralized Data》fraction_evaluate1.0表示每一轮都让所有节点参与评估strategy.start()将全局权重、训练配置学习率与轮次数交给策略由策略在每一轮采样节点、下发参数、收集更新并做加权聚合evaluate_fnglobal_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-rounds3联邦聚合轮次数即全局模型被下发-更新-聚合的次数fraction-evaluate1.0每轮参与评估的节点比例1.0表示全部节点参与local-epochs1每个 ClientApp 本地训练的 epoch 数learning-rate0.1本地训练使用的 SGD 学习率由服务端通过train_config下发batch-size32本地 DataLoader 的批大小save-modelfalse训练结束后是否将最终模型保存为final_model.pt这些配置均可通过flwr run . --run-config keyvalue ...在命令行覆盖例如num-server-rounds5 learning-rate0.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_fnNone, evaluate_metrics_aggr_fnNone, ) - 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声明了flwr1.36.0的版本下限所有命令与配置均以上述仓库实际内容为准。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价