news 2026/9/17 13:18:22

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

作者头像

张小明

前端开发工程师

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

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.pyclient_app.pyserver_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.IidPartitioneruoft-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(),完整流程为:

  1. Message中携带的全局模型参数初始化本地模型:model.load_state_dict(msg.content["arrays"].to_torch_state_dict())
  2. 自动选择设备:torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
  3. context.node_config读取节点身份配置partition-idnum-partitions,从context.run_config读取运行配置batch-size,据此加载本地数据;
  4. context.run_config["local-epochs"]为本地训练轮数、msg.content["config"]["lr"]为学习率执行训练;
  5. 将更新后的模型参数封装为ArrayRecord,连同train_lossnum-examples指标封装为MetricRecord,组合成RecordDict作为回复Message返回。

评估处理函数@app.evaluate()与训练流程类似:加载收到的模型参数,用本地测试集(valloader)评估,返回eval_losseval_accnum-examples三个指标。

这里体现了 Flower 新一代基于Message/Record的通信模型:ArrayRecord承载张量参数,MetricRecord承载标量指标,二者统一放进RecordDictMessage在网络中传输,客户端与服务端无需关心底层序列化细节。

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-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 "key=value ..."在命令行覆盖,例如num-server-rounds=5 learning-rate=0.05。此外[tool.flwr.app.components]段指定了 ServerApp 与 ClientApp 的导入路径(pytorchexample.server_app:apppytorchexample.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_nodesmin_available_nodes约束下完成采样;
  • 参数聚合:训练更新通过aggregate_arrayrecords聚合,指标通过aggregate_metricrecords聚合,二者都默认以weighted_by_key(默认"num-examples")作为权重进行加权平均——这正是客户端回复中必须携带num-examples指标的原因;
  • 轮次注入configure_train会在下发的配置中注入config["server-round"] = server_round,客户端可据此感知当前轮次;
  • 消息构造_construct_messages为每个被采样节点构造一条携带同一份RecordDict(包含arraysconfig)的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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 13:18:19

MySQL分区表实战:从原理到避坑,解决大表查询与归档难题

MySQL高阶实战:分区表(Partitioning Table)从原理到避坑做MySQL开发或者DBA的,谁没有被大表折磨过。单表几千万行,查询能跑但慢,索引越来越大,维护要半夜搞,删除旧数据直接DELETE能把…

作者头像 李华
网站建设 2026/9/17 13:16:41

Workflow-DSL:告别硬编码,实现可视化编排与零代码流程改造

只要手里管过几条AI工作流,基本都体会过硬编码的痛:需求一变就要翻代码,新模型一发布就要改接口,流程里某一步报错还只能靠日志一点点查。我见过不少团队,明明业务逻辑很简单,代码里却塞满了if...else、循环…

作者头像 李华
网站建设 2026/9/17 13:14:56

Servlet从概念到实战:Maven搭建与大模型HTTP接口调用

很多人第一次接触 Java Web 的时候,教材和视频里张口就是 Servlet,可真让你说清楚 Servlet 到底是什么、它在一次请求里扮演什么角色、为什么现在都用 Spring Boot 了还得回头学它,大部分人是要卡壳的。这篇文章我想把 Servlet 从概念到落地完…

作者头像 李华
网站建设 2026/9/17 13:12:40

电机控制工程师实战成长路径:从参数实测到FOC环路调优

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华