news 2026/9/17 6:20:15

Flower Datasets 与 PyTorch 集成指南:从 FederatedDataset 到 DataLoader 的完整实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Flower Datasets 与 PyTorch 集成指南:从 FederatedDataset 到 DataLoader 的完整实战

Flower Datasets 与 PyTorch 集成指南:从 FederatedDataset 到 DataLoader 的完整实战

【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower

本文讲解如何在 Flower 联邦学习框架中,将flwr-datasets(Flower Datasets)下载与划分好的联邦数据集无缝接入 PyTorch 的DataLoader,并保留你惯用的 PyTorch Transform 流水线。读完本文,你将掌握FederatedDataset的标准用法、特征名检查、with_transformmap两种数据转换方式、多种数据二次划分方案,以及联邦训练循环中字典型 batch 的正确取数方式,可直接用于联邦学习项目的客户端数据准备。

前置准备:安装与背景

flwr-datasets是 Flower 生态中负责数据集下载与划分的库,它基于 Hugging Face 的datasets实现,因此天然支持 Hugging Face、PyTorch、TensorFlow、NumPy、Pandas、JAX、Arrow 等多种格式。在 PyTorch 场景下,建议安装带视觉扩展的版本(以便处理图像数据集):

python -m pip install "flwr-datasets[vision]"

详细安装方式可参考 如何安装 flwr-datasets。本文的所有示例均以 CIFAR-10 为例,但流程适用于 Hugging Face Hub 上的任意数据集。

标准流程:创建 FederatedDataset 并加载分区

使用FederatedDataset完成"下载数据集 + 按客户端划分 + 保留集中式评估集"三件事,是使用flwr-datasets的标准起点:

from flwr_datasets import FederatedDataset fds = FederatedDataset(dataset="cifar10", partitioners={"train": 10}) partition = fds.load_partition(0, "train") centralized_dataset = fds.load_split("test")

逐行解读:

  • FederatedDataset(dataset="cifar10", partitioners={"train": 10}):指定数据集名称为cifar10,并将trainsplit 划分为 10 份(IID 划分),对应 10 个联邦客户端;
  • fds.load_partition(0, "train"):取出编号为 0 的客户端分区(partition_id取值范围为0num_partitions - 1),用于本地训练;
  • fds.load_split("test"):加载完整的testsplit,不参与划分,通常用于服务端的集中式评估。

从源码(federated_dataset.py)可以看到,FederatedDataset的完整构造参数还包括subset(数据子集/版本)、preprocessor(重划分等预处理)、shuffle(默认True,划分前随机打乱样本顺序)、seed(默认42,控制打乱的随机性)以及load_dataset_kwargs(透传给datasets.load_dataset的额外参数,如num_proc=4trust_remote_code=True)。此外,partitioners的值除了整数(代表 IID 划分为多少份)外,也可以是Partitioner对象——例如用DirichletPartitioner模拟非独立同分布(non-IID)场景,此时可传入num_partitionspartition_byalpha等参数。

值得注意的底层行为:数据集的下载是**惰性(lazy)**的,即只有在第一次调用load_partitionload_split时才真正触发datasets.load_dataset,随后依次执行打乱(shuffle)、预处理(preprocessor)与分区分配(见源码中的_prepare_dataset方法)。这意味着你可以在创建FederatedDataset后立即返回给各个客户端,由各客户端按需触发下载。

确认特征名:partition.features

load_partition返回的是 Hugging Face 的Dataset对象,数据以列(feature)组织。在写转换代码前,务必先确认特征名——不同数据集的命名习惯不同,可能是"img""image""label""labels"

partition.features

CIFAR-10 的输出如下:

{'img': Image(decode=True, id=None), 'label': ClassLabel(names=['airplane', 'automobile', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck'], id=None)}

即该数据集的特征名为img(图像)和label(类别标签,共 10 类)。后续所有转换和取数都要用这两个 key。

方式一:with_transform 实时应用 PyTorch Transform

datasets.Dataset.with_transform()是推荐的第一种转换方式,它最大的特点是**按需(on-the-fly)**执行:你指定的转换只在你真正访问数据时才生效,这与 PyTorch 生态中 Transform 的工作方式一致,不会预先物化整个数据集。

需要特别留意的是:with_transform中的函数作用在批量(batch)数据上——即使你只取一个元素,它也会被表示为一个大小为 1 的 batch。因此需要在函数内遍历该批次的每个样本并逐一应用转换:

from torch.utils.data import DataLoader from torchvision.transforms import ToTensor transforms = ToTensor() def apply_transforms(batch): batch["img"] = [transforms(img) for img in batch["img"]] return batch partition_torch = partition.with_transform(apply_transforms) # 可选:先通过 partition_torch[0] 检查转换是否有误 dataloader = DataLoader(partition_torch, batch_size=64)

完成with_transform之后,partition_torch直接就是一个符合 PyTorchDataset协议的datasets.Dataset,可以直接交给torch.utils.data.DataLoader。建议在正式训练前先执行一次partition_torch[0]做冒烟测试,确认转换逻辑没有书写错误。

仓库中的 PyTorch 端到端测试(pytorch_test.py)完整验证了这条链路:它使用FederatedDataset(dataset="cifar10", partitioners={"train": 100})加载分区后,用with_transform应用ToTensor()Compose([ToTensor(), Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]),并断言:DataLoader 产出的 batch 是dict类型、batch["img"]Tensor且形状为(batch_size, 3, 32, 32),同时用一轮训练验证 loss 不为 NaN/Inf。这说明"分区 → 实时转换 → DataLoader → 训练"的完整流程是被官方测试覆盖的可靠实践。

方式二:map 即时转换并配合 with_format("torch")

如果你希望转换立即执行(而不是访问时触发),可以使用map()函数。它与with_transform/set_transform不同:操作是即时完成的;同时要注意,map返回的字典中如果 key 已存在,会就地修改该特征,若 key 不存在则会新增一个特征。下面把数据集的"img"特征直接转换为 PyTorch Tensor:

from torch.utils.data import DataLoader from torchvision.transforms import ToTensor transforms = ToTensor() partition_torch = partition.map( lambda img: {"img": transforms(img)}, input_columns="img" ).with_format("torch") dataloader = DataLoader(partition_torch, batch_size=64)

这里map对每条样本执行{"img": transforms(img)}(使用input_columns="img"指定输入列),随后调用.with_format("torch")让数据在访问时以 PyTorch Tensor 格式返回,再交给DataLoader

两种方式如何取舍?简单说:with_transform是"转换随取随用",适合希望在训练循环里保留完整 PyTorch 语义、避免重复转换开销的场景;map是"转换立即落盘",适合希望一次性完成预处理、后续只做格式切换的场景。对于小数据集,两者性能差异很小,可按个人习惯选择。

为什么建议保留 ToTensor()

官方文档特别建议保留ToTensor()(尤其是你的 PyTorch 代码原本就使用它时),原因在于它完成了通道维度的交换:将形状从(H x W x C)变为(C x H x W)这种通道在前的顺序正是带卷积层(Conv2D)的模型所期望的输入布局。跳过这一步直接喂入原始图像,会因维度顺序不符而导致模型无法正确训练。

数据二次划分:训练/验证/测试子集

联邦场景中经常需要把某个客户端的分区再拆成训练集、验证集、测试集。flwr-datasets提供了三种方案,可在把数据集交给DataLoader之前的任意时刻使用。

方案一:Hugging Face 原生的 train_test_split

partition_train_test = partition.train_test_split(test_size=0.2, seed=42) partition_train = partition_train_test["train"] partition_test = partition_train_test["test"]

这是最简方案:按 80:20 拆分为训练与测试,seed=42保证可复现。缺点是一次只能拆成两份。

方案二:divide_dataset 按比例拆成多份

如果你需要保持样本顺序不变,并且要拆成 2 份或更多份,可以使用flwr_datasets.utils中的divide_dataset

from flwr_datasets.utils import divide_dataset train, valid, test = divide_dataset(partition, [0.6, 0.2, 0.2])

从实现(utils.py)来看,divide_dataset按给定比例从数据集开头依次切分division可以是一个list/tuple(如[0.6, 0.2, 0.2],返回list[Dataset]),也可以是一个dict(如{"train": 0.6, "valid": 0.2, "test": 0.2},返回带名字的DatasetDict)。源码中的校验逻辑要求:每个比例必须大于 0 且小于等于 1,各比例之和不能超过 1;若总和小于 1,会给出警告提示部分数据未被使用。

方案三:手动计算索引

最简单的就是自己计算索引范围,然后借助select切片:

partition_len = len(partition) # 将 partition 按 80:20 拆分 num_train_examples = int(0.8 * partition_len) # 使用前 80% partition_train = partition.select(range(num_train_examples)) # 使用后 20% partition_test = partition.select(range(num_train_examples, partition_len))

这种方法自由度最高,适合需要完全自定义划分边界的场景。

训练循环:字典型 batch 的取数方式

最后,训练循环里需要做一处关键调整。普通的 PyTorch DataLoader 每次迭代返回一个列表:

for batch in all_from_pytorch_dataloader: images, labels = batch # 或者: # images, labels = batch[0], batch[1]

flwr-datasets产出的数据集返回的是字典,需要通过 key 而不是下标取数:

for batch in dataloader: images, labels = batch["img"], batch["label"]

在端到端测试(pytorch_test.py)中,训练循环正是按此模式编写:inputs, labels = data['img'].to(device), data['label'].to(device),随后执行optimizer.zero_grad()、前向、loss.backward()optimizer.step(),与标准 PyTorch 训练完全兼容。

完整示例:联邦客户端数据准备 + PyTorch 训练

将上述要点串起来,一个典型的客户端侧数据准备与训练代码结构如下:

from flwr_datasets import FederatedDataset from torch.utils.data import DataLoader from torchvision.transforms import ToTensor # 1. 下载并划分数据集(10 个客户端),test 用于集中式评估 fds = FederatedDataset(dataset="cifar10", partitioners={"train": 10}) partition = fds.load_partition(0, "train") centralized_dataset = fds.load_split("test") # 2. 确认特征名(partition.features 输出 img / label) # 3. 实时应用 Transform transforms = ToTensor() def apply_transforms(batch): batch["img"] = [transforms(img) for img in batch["img"]] return batch # 4.(可选)本地二次划分为 train / valid / test train, valid, test = divide_dataset(partition, [0.8, 0.1, 0.1]) # 5. 构建 DataLoader trainloader = DataLoader(train.with_transform(apply_transforms), batch_size=64) # 6. 训练循环:按 key 取数 for batch in trainloader: images, labels = batch["img"], batch["label"] # 前向、反向、更新……

这套流程可以与 Flower 联邦学习框架的客户端实现直接衔接:每个客户端用自己的partition_id调用load_partition获取专属分区,再经由上述转换与 DataLoader 封装后进入本地训练;服务端则通过load_split("test")获取完整测试集进行集中式评估。相关的 NumPy 与 TensorFlow 集成方式可分别参考 如何与 NumPy 一起使用 与 如何与 TensorFlow 一起使用,两篇文档共享同一套FederatedDataset加载逻辑,仅在数据格式转换环节有所不同。

【免费下载链接】flowerFlower: A Friendly Federated AI Framework项目地址: https://gitcode.com/GitHub_Trending/flo/flower

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

低功耗IC设计实战:从DVFS到AVS与NTC的工程落地

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

作者头像 李华
网站建设 2026/9/17 6:19:35

FluidVoice全局快捷键设置教程:如何一键唤起全系统语音输入

FluidVoice全局快捷键设置教程:如何一键唤起全系统语音输入 【免费下载链接】FluidVoice Fastest and only macOS Dictation app with on-device STT and custom trained AI enhancement model. Windows pre-build available! A local Wispr Flow alternative. DM u…

作者头像 李华
网站建设 2026/9/17 6:19:09

RAPTOR可视化流程图引擎:零代码学算法与逻辑思维训练

简介:本资源是一份面向编程初学者与高校计算机基础课程教学的RAPTOR可视化程序设计入门教程,聚焦算法思维培养与流程图式编程实践。PPT课件系统讲解RAPTOR环境搭建、四大基本符号(输入/输出/赋值/过程调用)、变量定义与动态赋值机…

作者头像 李华
网站建设 2026/9/17 6:18:34

从零到入门:6个月成为机器人工程师的实战路线图

先说结论:6个月成为一名机器人工程师,这个目标可以实现,但和你想象的“看一遍视频就能调通一台机械臂”不是一回事。我自己走过这条路,也带过几个从电气、软件、机械转岗过来的同事,说实话,能不能成&#x…

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

一人成团做漫剧:豆包、即梦、剪映、扣子AI流水线实战指南

普通人一人成团做漫剧:把豆包、即梦、剪映、扣子串成一条AI流水线后台一直有人问我,漫剧到底是不是智商税?一个人到底能不能做?我的回答一直是:能,但前提是你别把四个工具当成四个孤岛,而是当成…

作者头像 李华