OneFlow 数据加载完全指南:从 DataLoader 架构到单/多进程实战
发布时间:2026/9/25 3:32:00来源:尧图网络
深度学习分布式训练模型优化【免费下载链接】oneflowOneFlow is a deep learning framework designed to be user-friendly, scalable and efficient.项目地址https://gitcode.com/gh_mirrors/one/oneflow点击查看免费下载oneflow.utils.data是 OneFlow 深度学习框架中负责数据加载的核心模块其灵魂是oneflow.utils.data.DataLoader类——一个建立在数据集Dataset之上的 Python 可迭代对象。本文以该模块的官方文档docs/source/utils.data.rst为主体骨架结合仓库中的真实实现源码python/oneflow/utils/data/展开讲解帮助你掌握 map/iterable 两种数据集类型、采样器定制加载顺序、自动/手动批处理、单进程与多进程加载以及 CUDA 内存固定memory pinning等完整技术链路最终能在自己的训练脚本中正确、高效地配置 DataLoader。DataLoader数据加载的枢纽类DataLoader组合了一个数据集与一个采样器对外提供对给定数据集的可迭代访问。它支持的完整能力包括map-style 与 iterable-style 两种数据集自定义数据加载顺序通过 Sampler自动批处理automatic batching单进程与多进程数据加载自动内存固定automatic memory pinning。以上选项全部由DataLoader的构造函数参数配置其完整签名为DataLoader(dataset, batch_size1, shuffleFalse, samplerNone, batch_samplerNone, num_workers0, collate_fnNone, pin_memoryFalse, drop_lastFalse, timeout0, worker_init_fnNone, *, prefetch_factor2, persistent_workersFalse)构造参数详解参数默认值作用dataset必填数据来源必须是 Dataset 对象batch_size1每个批次加载多少样本设为None时关闭自动批处理shuffleFalse每个 epoch 是否重新打乱数据samplerNone自定义采样策略与shuffle互斥batch_samplerNone一次产出「一批索引」的采样器与batch_size、shuffle、sampler、drop_last互斥num_workers0用于数据加载的子进程数0表示在主进程内加载collate_fnNone将样本列表合并为 mini-batch 的函数pin_memoryFalse返回数据前将 Tensor 拷贝到 CUDA 固定内存drop_lastFalse数据集大小不能被 batch_size 整除时是否丢弃最后一个不完整的 batchtimeout0从 worker 收集一个 batch 的超时秒数必须非负worker_init_fnNone每个 worker 子进程在完成随机种子设置后、开始加载数据前被调用入参为 worker id[0, num_workers-1]区间整数prefetch_factor2关键字参数每个 worker 预先加载的样本批数2表示全部 worker 合计预取2 * num_workers批样本persistent_workersFalse为True时数据集被消费一轮后 worker 进程不关闭保持 Dataset 实例存活源码层面这些参数的校验逻辑集中在 python/oneflow/utils/data/dataloader.py#L178-L336。值得注意的实现细节num_workers与timeout均不允许为负违反会直接抛出ValueErrornum_workers0时若显式指定非默认的prefetch_factor会报错因为预取只存在于多进程模式persistent_workersTrue必须搭配num_workers 0DataLoader 初始化完成后batch_size、batch_sampler、sampler、drop_last、dataset、persistent_workers等属性不允许再被修改通过__setattr__拦截见 dataloader.py#L345-L359。另外OneFlow 的 DataLoader 在设计上保持了与 PyTorch v1.7 的接口兼容性源码注释明确说明见 dataloader.py#L105因此熟悉 PyTorch 数据管线的用户可以零成本迁移。Dataset 的两种风格map-style 与 iterable-styledataset参数是 DataLoader 构造函数中最重要的参数。OneFlow 支持两种类型的数据集二者以不同的协议实现区分。Map-style datasets映射式数据集map-style 数据集实现了__getitem__与__len__两个协议表示从可能非整数的索引/键到数据样本的映射。例如dataset[idx]可以从磁盘读取第idx张图片及其标签。基类定义见 python/oneflow/utils/data/dataset.py#L57-L77class Dataset(Generic[T_co]): def __getitem__(self, index) - T_co: raise NotImplementedError def __add__(self, other: Dataset[T_co]) - ConcatDataset[T_co]: return ConcatDataset([self, other])子类只需要覆写__getitem__即可__len__可选择性实现但众多 Sampler 实现与 DataLoader 的默认行为都依赖它返回数据集大小。Iterable-style datasets可迭代式数据集iterable-style 数据集是IterableDataset的子类实现__iter__协议表示对数据样本流的迭代。它特别适合随机读取代价高昂甚至不可能的场景以及 batch 大小取决于所取数据本身的场景。例如iter(dataset)可以返回一条从数据库、远程服务器乃至实时日志中读取的数据流。从源码看IterableDataset还内置了一套函数注册机制register_function/register_datapipe_as_function以及用于进程间传输的自定义__reduce_ex__钩子见 dataset.py#L80-L187。例如以下最小实现class MyIterableDataset(flow.utils.data.IterableDataset): def __init__(self, start, end): super(MyIterableDataset).__init__() self.start start self.end end def __iter__(self): return iter(range(self.start, self.end)) ds MyIterableDataset(start3, end7) # 单进程加载得到 [3, 4, 5, 6] print(list(flow.utils.data.DataLoader(ds, num_workers0)))重要注意当IterableDataset与多进程数据加载num_workers 0搭配使用时同一数据集对象会在每个 worker 进程中被复制一份因此各副本必须配置不同例如在__iter__中按 worker 划分数据区间否则会产生重复数据。官方文档建议通过worker_init_fn或__iter__内的划分逻辑实现详见 dataset.py#L96-L135 中的两个示例。仓库内置的 Dataset 派生类模块还提供了开箱即用的数据集工具类dataset.pyTensorDataset包装一组第一维大小相同的 Tensordataset[i]返回各 Tensor 在第i行的切片元组构造时校验各 Tensor 第一维是否一致ConcatDataset将多个 map-style 数据集首尾拼接为一个数据集内部维护cumulative_sizes前缀和并用bisect定位样本所属子数据集不支持 IterableDatasetChainDataset将多个IterableDataset按顺序链式拼接拼接过程按需on-the-fly进行适合大规模数据流Subset按给定索引序列取原数据集的子集random_split(dataset, lengths, generator)将数据集随机切分为不重叠的若干份长度之和必须等于数据集长度使用flow._C.randperm生成随机置换可传入flow.Generator固定随机种子实现结果可复现。数据加载顺序与 Sampler对于 iterable-style 数据集加载顺序完全由用户自定义的迭代器控制这方便实现分块读取与动态 batch 大小例如每次 yield 一个已打包的 batch。本节其余内容针对 map-style 数据集。Sampler类用于指定数据加载过程中使用的索引/键序列它是对数据集索引的可迭代对象。例如在随机梯度下降SGD的常见场景中Sampler 可以随机打乱索引列表并逐个产出或每次产出少量索引用于 mini-batch SGD。Sampler 基类定义在 python/oneflow/utils/data/sampler.py#L25-L67。默认采样器的自动构造DataLoader 会根据shuffle参数自动构造顺序或随机采样器shuffleFalse→SequentialSampler(dataset)按range(len(dataset))顺序产出索引shuffleTrue→RandomSampler(dataset, generatorgenerator)每次迭代产出打乱后的索引。也可以显式传入自定义sampler对象它每次 yield 下一个要取数据的索引/键。若想一次产出「一批索引的列表」可将其作为batch_sampler传入。源码见 dataloader.py#L300-L319。约束与注意事项sampler与shuffle互斥sampler is not None and shuffle同时出现会抛ValueErrorbatch_sampler与batch_size、shuffle、sampler、drop_last互斥sampler与batch_sampler均不兼容 iterable-style 数据集——因为此类数据集没有键/索引的概念。源码中若dataset是IterableDataset而同时指定了shuffle/sampler/batch_sampler会直接抛错dataloader.py#L257-L275并为 iterable 数据集自动挂上一个无限的_InfiniteConstantSampler等价于itertools.repeat(None, None)见 dataloader.py#L78-L91。分布式训练DistributedSampler模块的 python/oneflow/utils/data/distributed.py 中还提供了DistributedSampler用于在多卡/多机分布式训练时对数据集做分片确保每个 rank 加载互不重叠的数据切片。批处理与 collate_fn自动 vs 手动DataLoader 通过batch_size、drop_last、batch_sampler和collate_fn带默认实现支持将逐条取出的样本自动整理collate成 batch。自动批处理默认开启这是最常见的情况取出一个 mini-batch 的数据并整理为批量样本——即包含一个 batch 维通常为第一维的 Tensor。当batch_size默认1不为None时数据加载器产出批量样本而非单个样本。batch_size与drop_last用于说明数据加载器如何获取「一批数据集键」。对 map-style 数据集也可以改用batch_sampler它每次产出键的列表。关键机制batch_size与drop_last本质上是从sampler构造batch_sampler的参数源码见 dataloader.py#L312-L314使用BatchSampler(sampler, batch_size, drop_last)。对 map-style 数据集sampler要么由用户提供要么由shuffle参数构造对 iterable-style 数据集sampler是那个无限的 dummy sampler。此外当以多进程方式加载 iterable-style 数据集时drop_last会丢弃每个 worker 数据集副本的最后一个不完整 batch。自动批处理下map-style 数据集的加载大致等价于for indices in batch_sampler: yield collate_fn([dataset[i] for i in indices])iterable-style 数据集的加载大致等价于dataset_iter iter(dataset) for indices in batch_sampler: yield collate_fn([next(dataset_iter) for _ in indices])禁用自动批处理当batch_size与batch_sampler均为Nonebatch_sampler默认值即None时自动批处理被禁用。此时数据集返回的每个样本都会经collate_fn处理后直接从数据加载器产出。这类场景包括希望在数据集代码中手动处理批处理、直接加载单个样本、从数据库批量读取或连续内存块更划算、batch 大小依赖数据本身、程序按单样本设计等。禁用自动批处理时map-style 数据集的加载大致等价于for index in sampler: yield collate_fn(dataset[index])iterable-style 数据集的加载大致等价于for data in iter(dataset): yield collate_fn(data)对应源码中的取数器fetcher逻辑实现于 python/oneflow/utils/data/_utils/fetch.py_MapDatasetFetcher与_IterableDatasetFetcher分别根据auto_collation标志决定是取「索引列表对应的样本列表」还是「单个索引对应的单个样本」。深入理解 collate_fncollate_fn在自动批处理开启与否时行为略有不同禁用自动批处理时collate_fn被逐个样本调用其输出即数据加载器迭代器产出的内容。此时默认的collate_fn就是default_convert——简单地把 NumPy 数组转换为 OneFlow Tensor其余类型原样保留。开启自动批处理时collate_fn每次被调用时收到一个样本列表负责把它们整理成一个 batch。默认实现为default_collate在 python/oneflow/utils/data/_utils/collate.py#L74-L114。例如若每个样本是「一张 3 通道图片 一个整数类别标签」组成的元组(image, class_index)默认collate_fn会把样本列表整理成「批量图片 Tensor 批量标签 Tensor」的元组。default_collate具备如下性质总是前置一个新维度作为 batch 维Tensor 分支调用flow._C.stack(batch, dim0)自动将 NumPy 数组和 Python 数值转换为 OneFlow Tensornumpy ndarray 会先flow.tensor(b)再递归 collate标量 numpy 数据、float转flow.float64、int均被转为 Tensor保留数据结构若每个样本是 dict输出同键 dict 但值为批量 Tensor无法转换时退化为 listlist、tuple、namedtuple同理字符串str/bytes保持原样返回Sequence元素若长度不一致会抛RuntimeError。用户可用自定义collate_fn实现定制批处理例如沿非第一维 collate、对不同长度序列做 padding 补齐到 batch 最大长度、或为自定义数据类型增加支持。当 DataLoader 输出的维度或类型与预期不符时优先检查collate_fn。单进程与多进程数据加载DataLoader默认使用单进程数据加载。Python 的全局解释器锁GIL阻止了线程间真正并行的 Python 代码执行。为避免数据加载阻塞计算代码OneFlow 提供了一个简单开关把num_workers设为正整数即可启用多进程数据加载。单进程模式默认此模式下数据抓取发生在 DataLoader 被初始化的同一个进程内因此数据加载可能阻塞计算。但在以下场景它反而更优进程间共享数据的资源共享内存、文件描述符有限整个数据集很小、可以全部载入内存单进程加载的报错回溯更易读便于调试。源码中_SingleProcessDataLoaderIter的实现非常直接取索引 → fetcher 抓数据 →若pin_memoryTrue执行固定内存见 dataloader.py#L561-L580。多进程模式将num_workers设为正整数即开启指定 worker 进程数的多进程数据加载。内存消耗警告迭代若干轮后worker 进程对「父进程中所有被访问到的 Python 对象」会消耗与父进程同等规模的 CPU 内存。若 Dataset 在构造时保存了大量数据如一个非常大的文件名列表且 worker 数较多总内存占用约为「worker 数 × 父进程大小」。最简单的规避方式是把 Python 对象替换为无引用计数的表示例如 Pandas、NumPy 或 PyArrow 对象。此模式下每次创建 DataLoader 迭代器例如调用enumerate(dataloader)时都会创建num_workers个 worker 进程并把dataset、collate_fn、worker_init_fn传给每个 worker用于初始化与取数。这意味着数据集访问及其内部 IO、变换包括collate_fn都在 worker 进程中执行。对 map-style 数据集主进程用sampler生成索引并分发给 worker因此打乱随机化在主进程完成由它指导「取哪些索引的数据」对 iterable-style 数据集每个 worker 持有数据集对象的一个副本朴素的多进程加载常导致数据重复。可用worker_init_fn独立配置每个副本参见 dataset.py#L96-L135 的示例同理多进程下drop_last会丢弃每个 worker 的 iterable 副本的最后一个不完整 batch。worker 在迭代结束或迭代器被垃圾回收时关闭。源码中多进程迭代器_MultiProcessingDataLoaderIter的核心数据流模型是见 dataloader.py#L583-L604 的注释主进程 ──{index_queue}── worker 进程 ──{worker_result_queue}── 主进程的 pin_memory 线程 ──{data_queue}── 数据输出即主进程把待取索引放入每个 worker 的index_queueworker 从index_queue取任务、经 fetcher 与collate_fn处理后把结果放入worker_result_queue若pin_memoryTrue主进程会启动一个pin_memory_thread从worker_result_queue读取结果、执行固定内存后写入data_queuequeue.Queue。worker 与 pin_memory 线程均设置为 daemon配合atexit钩子、workers_done_event、SIGCHLD 处理器等一整套复杂的关闭协议确保迭代器耗尽或进程异常退出时各方都能优雅退出、不挂死。CUDA Tensor 警告多进程加载中一般不建议直接返回 CUDA Tensor因为 CUDA 的使用与跨进程共享存在诸多微妙问题。推荐改用自动内存固定pin_memoryTrue它可加速数据向 CUDA 显卡的传输。平台相关行为worker 依赖 Python 的multiprocessing因此启动方式在 Windows 与 Unix 上不同Unix默认fork()启动方式子 worker 可直接通过克隆的地址空间访问dataset与 Python 参数函数Windows / macOS默认spawn()启动方式会启动另一个解释器运行主脚本内部 worker 函数通过pickle序列化接收dataset、collate_fn等参数。spawn 的独立序列化要求你做两件事以兼容 Windows 的多进程加载将主脚本大部分代码放进if __name__ __main__:块避免每个 worker 进程启动时重新执行主脚本很可能报错。Dataset 与 DataLoader 实例的创建逻辑可以放这里因为无需在 worker 中重新执行确保所有自定义的collate_fn、worker_init_fn或dataset代码声明为顶层定义、位于__main__检查之外以保证在 worker 进程中可见函数只按引用 pickle不携带字节码。此外由于 spawn 下worker_init_fn需要可 picklelambda 等不可 pickle 的对象不能用作worker_init_fn。多进程数据加载的随机性默认情况下每个 worker 的 OneFlow 随机种子被设为base_seed worker_id其中base_seed由主进程用其 RNG 生成强制消耗一次 RNG 状态或由指定的generator产生。但其他库的种子在 worker 初始化时可能相同导致各 worker 返回相同的随机数。可在worker_init_fn中用oneflow.initial_seed()读取每个 worker 的 OneFlow 种子并用它去播种其他库见 dataloader.py#L505-L508 中_base_seed的生成逻辑。worker 数量合理性检查多进程模式下OneFlow 还会做 worker 数量合理性检查check_worker_number_rationality见 dataloader.py#L426-L488若创建的 worker 数超过系统可用 CPU 数优先读取os.sched_getaffinity否则回退到os.cpu_count()会发出警告提示过多 worker 可能导致 DataLoader 变慢甚至卡死。内存固定pin_memory当数据从主机内存拷贝到 GPU 时若源数据来自固定内存page-locked memory拷贝速度会显著更快。对数据加载而言给 DataLoader 传入pin_memoryTrue会自动把取到的数据 Tensor 放入固定内存从而加速向 CUDA 设备的传输。默认的内存固定逻辑只识别 Tensor 以及包含 Tensor 的映射和可迭代对象。若collate_fn返回的是自定义 batch 类型或 batch 中每个元素是自定义类型固定逻辑无法识别它们会原样返回而不固定内存。要对自定义 batch/数据类型启用内存固定需要在这些自定义类型上定义pin_memory方法。源码中固定逻辑实现于 python/oneflow/utils/data/_utils/pin_memory.pypin_memory(data)递归处理 Tensor调用data.pin_memory()、字符串、Mapping、namedtuple、Sequence并且对任何定义了pin_memory方法的对象直接调用其方法这正是自定义类型固定内存的入口_pin_memory_loop是主进程中独立守护线程的循环体其中还会调用flow.set_num_threads(1)避免固定内存拷贝占满全部 CPU 核。官方文档给出的完整示例自定义 batch 类型 自定义 pin_memory 方法 collate 包装器class SimpleCustomBatch: def __init__(self, data): transposed_data list(zip(*data)) self.inp oneflow.stack(transposed_data[0], 0) self.tgt oneflow.stack(transposed_data[1], 0) # custom memory pinning method on custom type def pin_memory(self): self.inp self.inp.pin_memory() self.tgt self.tgt.pin_memory() return self def collate_wrapper(batch): return SimpleCustomBatch(batch) inps oneflow.arange(10 * 5, dtypeoneflow.float32).view(10, 5) tgts oneflow.arange(10 * 5, dtypeoneflow.float32).view(10, 5) dataset TensorDataset(inps, tgts) loader DataLoader(dataset, batch_size2, collate_fncollate_wrapper, pin_memoryTrue) for batch_ndx, sample in enumerate(loader): print(sample.inp.is_pinned()) print(sample.tgt.is_pinned())完整 API 一览oneflow.utils.data模块对外暴露的核心 API 包括核心类DataLoader、Dataset、IterableDataset、TensorDataset、ConcatDataset、Subset工具函数random_split采样器Sampler、SequentialSampler、RandomSampler、SubsetRandomSampler、BatchSampler、distributed.DistributedSampler内部辅助python/oneflow/utils/data/_utils/下default_collate/default_convertcollate.py、map/iterable 取数器fetch.py、worker 循环与get_worker_infoworker.py、内存固定pin_memory.py等均可在python/oneflow/utils/data/目录中深入查阅。实战建议小结默认配置最省心DataLoader(dataset, batch_size32, shuffleTrue)即可获得「每 epoch 打乱 自动成批」的标准训练管线内部自动完成 SequentialSampler/RandomSampler 与 BatchSampler 的组装需要自定义采样如按权重采样、子集采样时使用sampler/batch_sampler但注意与shuffle、batch_size、drop_last的互斥关系数据在本地小文件时保持num_workers0避免多进程的内存复制开销与调试困难数据量大、IO 密集时再开多进程并按 CPU 核数合理设置num_workers必要时配合prefetch_factor提升吞吐GPU 训练务必开启pin_memoryTrue加速主机到设备拷贝返回自定义 batch 类型时记得实现pin_memory方法分布式训练使用distributed.DistributedSampler为每个 rank 分片数据若在 RDMA 分布式训练中启用 OneFlowpersistent_workers必须为True否则会触发段错误见 dataloader.py#L144-L146 的参数说明。赞分享深度学习分布式训练模型优化【免费下载链接】oneflowOneFlow is a deep learning framework designed to be user-friendly, scalable and efficient.项目地址https://gitcode.com/gh_mirrors/one/oneflow点击查看免费下载相关推荐PyTorch torch.utils.data 数据加载全指南从 DataLoader 构造到多进程与内存固定PyTorch torch.utils.data 数据加载全指南从 DataLoader 构造到多进程与内存固定 本文以 PyTorch 源码仓库中的官方数据人工智能机器学习深度学习分布式训练模型编译MXNet Gluon Dataset 与 DataLoader 实战指南从内存数据、图像目录到自定义数据集与多进程加载MXNet Gluon Dataset 与 DataLoader 实战指南从内存数据、图像目录到自定义数据集与多进程加载 Gluon 的 Dataset 与深度学习人工智能机器学习分布式训练graph-gophers/dataloader 迁移指南从 v1 到 v5 的 API 演进与 Go 数据加载器实战graph gophers/dataloader 迁移指南从 v1 到 v5 的 API 演进与 Go 数据加载器实战 导读 github.com/graph后端任务调度工作流自动化微服务上一篇React Query 的 useIsMutating 全面解析精确统计应用中正在执行的 mutation 数量下一篇Falco开源项目品牌资产库资源管理创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网