ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Flower 实战指南:使用 PyTorch 编写并运行你的第一个 Flower App(CIFAR-10 图像分类)

Flower 实战指南:使用 PyTorch 编写并运行你的第一个 Flower App(CIFAR-10 图像分类) Flower 实战指南使用 PyTorch 编写并运行你的第一个 Flower AppCIFAR-10 图像分类【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower在完成 Flower 入门教程运行flwrlabs/demo、理解ServerApp/ClientApp/策略/pyproject.toml的协作方式之后本篇指南将带你使用同样的工作流编写一个真实的 PyTorch Flower App在 CIFAR-10 图像分类数据集上训练一个小型卷积神经网络。读完本文你将掌握如何用flwr new拉取 PyTorch 快速开始模板、理解ClientApp与ServerApp的完整实现、在 SuperGrid 云端与本地两种模式下运行 App并通过--run-config在运行时覆盖训练超参数。本文对应的官方教程位于 framework/docs/source/tutorial-series-write-your-first-flower-app-pytorch.rst完整可运行的 App 源码位于 examples/quickstart-pytorch。本教程在 Flower 协作 AI 教程系列中的位置本教程是 Flower 协作 AI 教程系列Collaborative AI Tutorial的第三部分。系列脉络如下部分主题对应文档Part 1在 SuperGrid 上运行现成 App、认识模拟联邦tutorial-series-get-started-with-flower.rstPart 2下载 App、运行并理解各组件如何配合tutorial-series-write-your-first-flower-app.rstPart 3本文用 PyTorch 编写第一个真实的 Flower Apptutorial-series-write-your-first-flower-app-pytorch.rstPart 4自定义联邦学习策略tutorial-series-use-a-federated-learning-strategy-pytorch.rst与上一个使用 NumPy 数组的 demo 相比本教程的 App 拥有真实的模型、真实的数据集、真实的本地训练循环但 Flower 的整体骨架完全一致ServerApp启动运行FedAvg策略协调每一轮联邦学习每个ClientApp使用所在 SuperNode 上的数据训练或评估模型。创建 App从 Flower Hub 拉取 PyTorch 快速开始模板安装好 Flower 之后在终端执行$ 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这个 App 的工作负载是在CIFAR-10数据集上训练一个小型卷积神经网络。CIFAR-10 是包含 10 个类别的图像分类数据集类别包括飞机airplane、汽车automobile、鸟bird、猫cat、狗dog、船ship、卡车truck等。在仓库中该模板的源码位于 examples/quickstart-pytorch其 pyproject.toml 声明了如下依赖flwr[simulation]1.36.0、flwr-datasets[vision]0.6.1、torch、torchvision。安装依赖并注册本地包$ pip install -e .快速认识各文件职责在运行 App 之前先弄清每个文件负责什么pytorchexample/task.py包含 PyTorch 专属代码——神经网络定义、CIFAR-10 数据加载与分区、本地训练循环、评估循环以及服务端评估辅助函数。pytorchexample/client_app.py定义ClientApp。其中app.train()处理器接收当前全局模型加载一份 CIFAR-10 分区在本地训练模型并回复更新后的模型参数与指标app.evaluate()处理器在本地验证数据上评估收到的模型并回复指标。pytorchexample/server_app.py定义ServerApp。它创建初始 PyTorch 模型将模型参数包装为ArrayRecord创建FedAvg策略并启动联邦学习运行。pyproject.toml声明 App 元数据与依赖将 Flower 指向ServerApp和ClientApp对象并定义运行配置值例如服务端轮数、批次大小、本地训练轮数、学习率与评估设置。本 App 使用Flower Datasets下载 CIFAR-10 并划分为多个分区每个分区对应一个模拟客户端。这种从单一集中式数据集切分分区的做法非常适合**模拟Simulation**场景让你即使只有一个中心化数据集也能实验联邦学习。而在模拟之外的真实部署场景中通常不会人为创建分区每个ClientApp直接加载其运行所在 SuperNode 上已有的数据。在 SuperGrid 上运行 AppSuperGrid 是 Flower 托管的云端协作 AI 平台也是运行 Flower 协作工作流的推荐方式。若你还没有 SuperGrid 账号和模拟联邦请先完成入门教程。打开终端、激活 Python 环境先登录 SuperGrid# This will open a browser window where you can enter your SuperGrid credentials. $ flwr login supergrid登录成功后进入 App 目录并运行# Navigate to the directory of the app you want to run $ cd /path/to/quickstart-pytorch # Run the app $ flwr run . supergridSuperGrid 会为该 App 启动一次新的 run。打开 SuperGrid 仪表盘选择你的 federation点击这次新的 run 即可跟踪进度并查看日志。在日志中你会看到 Flower 启动FedAvg策略并运行多轮联邦学习。每一轮包括在选中的ClientApp实例上进行本地训练、在ServerApp中进行聚合以及eval_loss、eval_acc等评估指标。运行时覆盖配置你可以通过--run-config在运行时覆盖pyproject.toml中的配置值# Run the app for five rounds instead of the default three rounds $ flwr run . supergrid \ --run-config num-server-rounds5 # Run the app for five rounds and a smaller batch size $ flwr run . supergrid \ --run-config num-server-rounds5 \ --run-config batch-size16在 SuperGrid 上可以使用--federation标志配合 federation ID 来指定运行目标若省略Flower 默认使用your-account/workspace。关于创建和管理 federation 的更多细节参见在 SuperGrid 上创建与管理联邦。在本地运行 App在开发或调试阶段在本地运行同一个 App 也很有用。进入 App 下载目录执行下面的命令$ cd /path/to/quickstart-pytorch $ flwr run . local --streamFlower 会启动一个托管的本地 SuperLinkSuperGrid 的精简版本并在你的机器上以模拟 SuperNodes 执行 App。第一次运行耗时较长因为 App 需要下载 CIFAR-10 并安装依赖。加上--stream标志后你可以在终端实时看到本地运行的日志。流式输出大致如下INFO : Starting FedAvg strategy: INFO : ├── Number of rounds: 3 INFO : ... INFO : [ROUND 1/3] INFO : configure_train: Sampled 2 SuperNodes (out of 2) INFO : aggregate_train: Received 2 results and 0 failures INFO : └── Aggregated MetricRecord: {train_loss: 2.149280} INFO : configure_evaluate: Sampled 2 SuperNodes (out of 2) INFO : aggregate_evaluate: Received 2 results and 0 failures INFO : └── Aggregated MetricRecord: {eval_loss: 2.31319, eval_acc: 0.13004} INFO : [ROUND 2/3] INFO : ... INFO : [ROUND 3/3] INFO : ... INFO : Strategy execution finished两个注意事项上面flwr run命令没有指定 federation因为本地原型开发只有一个 federation 可用因此不需要--federation标志。在 Windows 上若看到异常终端输出例如□[32m□[1m请查阅 FAQ 中关于 Windows 意外输出的条目。关于使用 Flower CLI 与本地运行的 SuperLink 交互的更多细节包括如何列出 run、查看日志参见使用托管的本地 SuperLink 在本地运行 Flower。深入解析 Appflwrlabs/quickstart-pytorch展示了一个简单的联邦学习工作流服务端将全局模型参数发送给客户端客户端用收到的参数初始化本地模型在本地数据上训练这会改变本地模型参数再把更新后的模型参数发回服务端或者只发送梯度而非完整模型参数。定义 Flower ClientApp联邦学习系统由服务端和多个客户端SuperNodes组成。在 Flower 中我们分别创建ServerApp和ClientApp来运行服务端与客户端代码。ClientApp的核心职责是利用其所在 SuperNode例如边缘设备、数据中心服务器或笔记本电脑能够访问的本地数据执行某些操作。在本教程中这个操作就是使用本地训练集和验证集训练并评估前面定义的小型 CNN 模型。加载数据本 App 在 CIFAR-10 上训练一个小型卷积神经网络。由于本教程使用Simulation Runtime参见如何运行模拟所有数据都源自一个中心化数据集并被切分为多个分区每个分区对应一个模拟 SuperNode。task.py中的load_data()函数使用 Flower Datasets 加载一个分区、将其拆分为训练集与验证集、应用 PyTorch 变换并返回两个DataLoaderdef load_data(partition_id: int, num_partitions: int, batch_size: int): Load partition CIFAR10 data. # Only initialize FederatedDataset once 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) # Divide data on each SuperNode: 80% train, 20% test partition_train_test partition.train_test_split(test_size0.2, seed42) pytorch_transforms Compose( [ToTensor(), Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))] ) def apply_transforms(batch): Apply transforms to the partition from FederatedDataset. batch[img] [pytorch_transforms(img) for img in batch[img]] return batch 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关键点解读对应源码 examples/quickstart-pytorch/pytorchexample/task.py使用IidPartitioner(num_partitionsnum_partitions)按 IID 方式切分数据FederatedDataset通过模块级全局变量fds缓存保证每个进程中只初始化一次每个 SuperNode 上的分区再按80% 训练 / 20% 测试划分seed42保证可复现变换采用ToTensor()加Normalize((0.5,0.5,0.5),(0.5,0.5,0.5))将像素值归一化到[-1, 1]训练集DataLoader开启shuffleTrue测试集不洗牌。这种分区只在模拟时需要。在部署场景中每个 SuperNode 通常会直接加载自己的本地数据例如通过--node-config传入数据路径。训练通过用app.train()装饰器包装函数来定义ClientApp如何执行训练。该函数此处命名为train始终接收两个参数Message从服务端收到的消息包含模型参数以及服务端发送的其他配置信息Context包含执行ClientApp的 SuperNode 信息与当前 run 信息的上下文对象。通过Context可以取回 App 在pyproject.toml中定义的配置。Context还可用于在多次train或evaluate调用之间持久化客户端状态。在 Flower 中ClientApp是临时对象ephemeral只为执行一条Message而实例化当回复发回服务端后即被销毁。下面是使用前述 PyTorch CNN 模型的ClientApp实现通过消息应用来自ServerApp的参数、加载本地数据、用train_fn训练模型并生成一条包含更新后模型参数及若干指标的回复Messagefrom pytorchexample.task import train as train_fn # Flower ClientApp app ClientApp() app.train() def train(msg: Message, context: Context): Train the model on local data. # Load the model and initialize it with the received weights model Net() model.load_state_dict(msg.content[arrays].to_torch_state_dict()) device torch.device(cuda:0 if torch.cuda.is_available() else cpu) model.to(device) # Load the data partition_id context.node_config[partition-id] num_partitions context.node_config[num-partitions] batch_size context.run_config[batch-size] trainloader, _ load_data(partition_id, num_partitions, batch_size) # Call the training function train_loss train_fn( model, trainloader, context.run_config[local-epochs], msg.content[config][lr], device, ) # Construct and return reply Message model_record ArrayRecord(model.state_dict()) metrics { train_loss: train_loss, num-examples: len(trainloader.dataset), } metric_record MetricRecord(metrics) content RecordDict({arrays: model_record, metrics: metric_record}) return Message(contentcontent, reply_tomsg)这里有几个值得注意的细节partition-id和num-partitions由Simulation Runtime提供写入context.node_config。在部署场景中ClientApp通常加载 SuperNode 上已有的数据例如启动 SuperNode 时通过--node-config data-path/path/to/data传入路径然后在代码中读取context.node_config[data-path]。train_fn只是对task.py中训练函数的别名。调用时传入本地待训练的模型、数据加载器、本地训练轮数local-epochs与学习率lr。注意local-epochs通过Context从run config读取而lr则通过Message从服务端发送的ConfigRecord读取——这样服务端可以在每一轮动态调整学习率当不需要这种动态性时从 run config 读取lr同样完全有效。训练完成后ClientApp构造一条回复Message其content是一个RecordDict通常包含两条记录ArrayRecord更新后的模型参数MetricRecord相关指标此处为训练损失和训练样本数。必须在指标中返回num-examples键因为FedAvg等策略默认依赖该键来按权重聚合模型与指标除非你覆盖weighted_by_key参数例如FedAvg(weighted_by_keymy-different-key)。这一点在FedAvg的源码 framework/py/flwr/serverapp/strategy/fedavg.py 的weighted_by_key参数注释中有明确说明默认值为num-examples。构造完回复Message后ClientApp将其返回Flower 会自动把回复发回服务端。评估典型的联邦学习设置中ClientApp还会实现app.evaluate()函数在本地验证数据上评估从ServerApp收到的模型。这特别有助于在训练过程中监控全局模型在每个客户端上的表现。evaluate的实现与train几乎相同区别在于它调用task.py中定义的test_fn函数实现 PyTorch 评估循环并且返回的Message只包含一个MetricRecord评估期间不更新模型参数因此没有ArrayRecordfrom pytorchexample.task import test as test_fn app.evaluate() def evaluate(msg: Message, context: Context): Evaluate the model on local data. # Load the model and initialize it with the received weights model Net() model.load_state_dict(msg.content[arrays].to_torch_state_dict()) device torch.device(cuda:0 if torch.cuda.is_available() else cpu) model.to(device) # Load the data partition_id context.node_config[partition-id] num_partitions context.node_config[num-partitions] batch_size context.run_config[batch-size] _, valloader load_data(partition_id, num_partitions, batch_size) # Call the evaluation function eval_loss, eval_acc test_fn( model, valloader, device, ) # Construct and return reply Message metrics { eval_loss: eval_loss, eval_acc: eval_acc, num-examples: len(valloader.dataset), } metric_record MetricRecord(metrics) content RecordDict({metrics: metric_record}) return Message(contentcontent, reply_tomsg)如你所见evaluate与train几乎一致只是调用test_fn而非train_fn且返回的Message只含评估相关指标eval_loss、eval_acc——均为标量。同样指标中必须包含num-examples键服务端才能正确聚合评估指标。定义 Flower ServerApp在服务端需要配置一个**策略strategy**来封装联邦学习方法/算法例如联邦平均Federated Averaging, FedAvg。Flower 内置了多种策略也支持使用自定义策略实现来定制联邦学习方法的几乎所有方面。本教程使用内置FedAvg并微调了参与每轮训练的 SuperNode 比例。构造ServerApp需要定义其app.main()方法该方法接收两个输入参数Grid用于与运行ClientApp的 SuperNodes 交互的对象可让它们参与训练/评估/查询等轮次Context提供对运行配置访问的上下文对象。在通过策略的start()方法启动前需要初始化全局模型——它会在第一轮联邦学习中被发送给各客户端上的ClientApp。做法是创建模型实例Net、取出其state_dict中的参数、构造ArrayRecord并通过start()的initial_arrays参数将其提供给策略。还可以可选地向start()传入一个ConfigRecord其中包含希望传给客户端的设置这些设置会随携带模型参数的Message一起发送app ServerApp() app.main() def main(grid: Grid, context: Context) - None: Main entry point for the ServerApp. # Read run config fraction_evaluate: float context.run_config[fraction-evaluate] num_rounds: int context.run_config[num-server-rounds] lr: float context.run_config[learning-rate] # Load global model global_model Net() arrays ArrayRecord(global_model.state_dict()) # Initialize FedAvg strategy strategy FedAvg(fraction_evaluatefraction_evaluate) # Start strategy, run FedAvg for num_rounds result strategy.start( gridgrid, initial_arraysarrays, train_configConfigRecord({lr: lr}), num_roundsnum_rounds, evaluate_fnglobal_evaluate, ) # Save final model to disk print(\nSaving final model to disk...) state_dict result.arrays.to_torch_state_dict() torch.save(state_dict, final_model.pt)ServerApp的大部分执行发生在strategy.start()方法内部。在运行完指定轮数num_rounds之后start()返回一个Result对象其中包含最终的模型参数以及从客户端收到或由策略自身生成的指标。随后即可将最终模型保存到磁盘供后续使用。仓库中的实现还额外提供了global_evaluate函数对应源码 examples/quickstart-pytorch/pytorchexample/server_app.py它在完整测试集上评估全局模型并返回MetricRecord({accuracy: ..., loss: ...})并且只有context.run_config[save-model]为真时才保存模型——这两个细节都对应pyproject.toml中的配置项。配置项一览模板默认配置定义在 examples/quickstart-pytorch/pyproject.toml 的[tool.flwr.app.config]段配置键默认值作用num-server-rounds3联邦学习总轮数由ServerApp读取并传给strategy.start(num_rounds...)fraction-evaluate1.0参与评估的 SuperNode 比例传给FedAvg(fraction_evaluate...)local-epochs1每个客户端本地训练轮数由ClientApp从 run config 读取learning-rate0.1学习率由ServerApp写入ConfigRecord({lr: lr})随消息发给客户端batch-size32训练/评估批次大小由ClientApp从 run config 读取save-modelfalse是否在训练结束后将最终模型保存为final_model.pt同时[tool.flwr.app.components]段将 Flower 指向具体对象serverapp pytorchexample.server_app:app、clientapp pytorchexample.client_app:app。背后原理一轮联邦学习是如何执行的当执行flwr run使用默认本地连接配置时Flower 会把 run 提交给托管的本地 SuperLink。默认情况下本地 SuperLink 将 Simulation Runtime 配置为使用两个 SuperNode每个 SuperNode 都会运行一个前面定义的ClientApp实例。本地 SuperLink 随后启动ServerApp并要求它通过FedAvg策略向这些 SuperNodes 下发指令。在本示例中FedAvg配置了两个关键参数fraction-train1.0→ 选择100%可用客户端参与训练fraction-evaluate1.0→ 选择100%可用客户端参与评估。因此在示例中所有客户端SuperNodes都会同时被采样参与训练轮与评估轮。从FedAvg的源码 framework/py/flwr/serverapp/strategy/fedavg.py 可以看到其默认参数还包括min_train_nodes2、min_evaluate_nodes2、min_available_nodes2、weighted_by_keynum-examples、arrayrecord_keyarrays、configrecord_keyconfig——当fraction_*计算出的采样数低于最小值时仍会采样至少min_*_nodes个节点。典型的一轮训练与评估流程如下训练轮FedAvg选择所有客户端2/2Flower 向每个被选中的ClientApp发送TRAIN消息每个ClientApp调用app.train()装饰的函数返回包含ArrayRecord更新后的模型参数和MetricRecord训练损失与样本数的MessageServerApp收到所有回复FedAvg将所有ArrayRecord聚合成代表新全局模型的ArrayRecord并合并所有MetricRecord。评估轮FedAvg选择所有客户端2/2Flower 向每个ClientApp发送EVALUATE消息每个ClientApp调用app.evaluate()装饰的函数返回只含MetricRecord评估损失、准确率与样本数的MessageServerApp收到所有回复FedAvg聚合所有MetricRecord。训练与评估都完成后下一轮开始又一次训练、又一次评估……直到达到配置的轮数。收尾与下一步至此你已经成功在 SuperGrid 与本地运行了一个 PyTorch Flower App。与 NumPy demo 相比这个 App 使用了真实的模型、真实的数据集、真实的本地训练但 Flower 的结构完全一致ServerApp、ClientApp、策略与pyproject.toml。接下来可以继续学习系列第四部分使用联邦学习策略PyTorch了解如何自定义联邦学习策略改变服务端协调训练与评估的方式。如果你想直接动手实验也可以在 examples/quickstart-pytorch 中查看本教程 App 的完整源码用flwr run . local --stream在本地复现上述全部流程。【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表