完全指南:Dataset、DataLoader、Sampler 与 Transform)
PyTorch C 数据加载 APItorch::data完全指南Dataset、DataLoader、Sampler 与 Transform【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch本文以当前仓库中torch::data的官方文档docs/cpp/source/api/data/index.md 及其子页面为骨架系统讲解 PyTorch 在 C 前端中加载与预处理训练数据的完整生态。你将掌握torch::data四大核心组件——Dataset单样本/批次数据访问、DataLoader批量化与多 worker 并行加载、Sampler数据访问顺序控制、Transform归一化/增广等预处理——的接口设计、组合方式与在 MNIST 训练任务中的端到端用法。文中所引接口均可在仓库源码中交叉验证可直接用于你自己的 C 训练代码。一、torch::data 是什么什么时候用torch::data命名空间为训练过程中的数据加载与处理提供了一整套工具其核心定位在 docs/cpp/source/api/data/index.md 中被概括为四件事数据集抽象、用于批量与打乱的数据加载器、控制数据访问模式的采样器、用于数据增广的变换。从使用场景看以下情况应优先考虑torch::data需要按 batch 加载训练数据需要多 worker 并行数据加载需要实现自定义数据集或自定义 transform。官方索引页进一步给出了组件速览Dataset定义如何访问单个样本需实现get()与size()DataLoader对样本做批量化并可选择打乱顺序、并行加载Sampler控制样本被访问的顺序Transform对样本做预处理归一化、增广等。在 C 前端中所有数据相关头文件最终聚合在单个总入口之下。torch/data.h对应磁盘文件 torch/csrc/api/include/torch/data.h实际只是对四个子头文件的统一转出#include torch/data/dataloader.h #include torch/data/datasets.h #include torch/data/samplers.h #include torch/data/transforms.h并在torch::data命名空间中用using声明将datasets::BatchDataset与datasets::Dataset直接导出方便书写。文档列出的四个官方头文件与仓库实际位置对应如下头文件文档给出仓库实际路径torch/data.h主数据头torch/csrc/api/include/torch/data.htorch/data/dataloader.htorch/csrc/api/include/torch/data/dataloader.htorch/data/datasets.htorch/csrc/api/include/torch/data/datasets.htorch/data/samplers.htorch/csrc/api/include/torch/data/samplers.htorch/data/datasets.h、samplers.h、transforms.h又分别聚合各自子目录下的具体实现详见后文各节。注意实际使用时通常只需要#include torch/torch.h即包含上述全部 API。二、一个从零到一的快速示例文档在索引页给出了最具代表性的最小用法加载内置 MNIST 数据集 → 链式施加归一化与 Stack 变换 → 构造带批大小与 worker 数的 DataLoader → 遍历 batch。#include torch/torch.h // 加载内置数据集 auto dataset torch::data::datasets::MNIST(./data) .map(torch::data::transforms::Normalize(0.1307, 0.3081)) .map(torch::data::transforms::Stack()); // 创建带 batching 与 shuffling 的数据加载器 auto data_loader torch::data::make_data_loader( std::move(dataset), torch::data::DataLoaderOptions().batch_size(64).workers(4)); // 遍历批次 for (auto batch : *data_loader) { auto images batch.data; // Shape: [64, 1, 28, 28] auto labels batch.target; // Shape: [64] }其中有三点值得展开MNIST(./data)从./data目录读取 MNIST。仓库中该数据集的实现在 torch/csrc/api/include/torch/data/datasets/mnist.h若目录下无数据首次运行时数据集构造会负责从标准镜像源下载并解压到该目录训练、测试读取逻辑以及图像字节序处理都在此头文件与对应源文件中完成。两次.map(...)map是 datasets/base.h 中BatchDataset提供的成员模板它会返回一个MapDatasetSelf, TransformType把后面的变换“包”进数据集本身连续调用即形成预处理流水线Normalize在单样本级把像素归一化Stack把整批单样本堆叠成四维张量。其声明为template typename TransformType MapDatasetSelf, TransformType map(TransformType transform) ;make_data_loader(std::move(dataset), options)构造函数模板它先从dataset.size()拿到样本总数并默认构造一个RandomSampler用于训练时打乱顺序源码 dataloader.h 中可见这一逻辑先dataset.size()随后调用Sampler(*size)若数据集无 size则通过TORCH_CHECK报错提示。因此上方示例天然是“每个 epoch 数据都被随机打乱”的训练加载方式。三、Dataset定义单样本的访问方式Dataset是数据抽象的核心。官方 datasets 文档对应正文页 docs/cpp/source/api/data/datasets.md明确所有数据集继承自Dataset且必须实现get()与size()。3.1 类层次BatchDataset → Dataset → StreamDataset从源码 datasets/base.h 可以看清其继承结构BatchDatasetSelf, Batch, BatchRequest是最基础抽象提供纯虚get_batch(BatchRequest request)与size()并声明静态常量is_stateful通过判断BatchType是否为std::optional推导见同文件detail::is_optional。它同时是.map()方法的宿主。DatasetSelf, SingleExample继承自BatchDatasetSelf, std::vectorSingleExample。文档说“Dataset也是BatchDataset因为它支持随机访问”源码正是这么实现的子类只需实现逐样本的get(size_t index)基类默认的get_batch会循环调用get()填充批次virtual std::vectorExampleType get_batch(ArrayRefsize_t indices) override { std::vectorExampleType batch; batch.reserve(indices.size()); for (const auto i : indices) { batch.push_back(get(i)); } return batch; }需要的话子类也可以重写get_batch做自定义批量读取。StreamDataset别名表示“可能是无限的数据流”它的BatchRequest不再是索引数组而只是一个代表批次大小的size_ttemplate typename Self, typename Batch std::vectorExample using StreamDataset BatchDatasetSelf, Batch, /*BatchRequest*/size_t;3.2 单样本的载体Exampleget()返回的类型默认是Example定义在 torch/csrc/api/include/torch/data/example.htemplate typename Data at::Tensor, typename Target at::Tensor struct Example { Data data; Target target; };默认data与target都是at::Tensor即一个样本 输入张量 标签张量。对无标签的场景ExampleData, example::NoTarget提供了特化并在torch::data命名空间预置别名TensorExample Exampleat::Tensor, example::NoTarget——Example中目标模板参数设为void时还会隐式转换回底层数据类型方便无监督/仅生成任务直接取出数据。3.3 自定义 Dataset写一个训练示例datasets 文档给出了自定义数据集的模板class CustomDataset : public torch::data::datasets::DatasetCustomDataset { public: explicit CustomDataset(const std::string root) { // Load data from root directory } torch::data::Example get(size_t index) override { return {images_[index], labels_[index]}; } torch::optionalsize_t size() const override { return images_.size(0); } private: torch::Tensor images_, labels_; };CRTP 用法DatasetCustomDataset让基类的map()能返回携带具体子类类型Self的MapDataset从而保证流水线类型信息不丢失、无需类型擦除。注意size()的返回类型是torch::optionalsize_t即std::optionalsize_t允许数据集声明自己“尺寸未知”。3.4 其余数据集类datasets 文档还罗列了若干进阶类型可通过 Doxygen 块查看各成员StatefulDataset跨批次自行管理内部状态例如数据流中当前位置的数据集它直接产出 batch 而非依赖外部 sampler见 datasets/stateful.h。ChunkDataReader/ChunkDataset面向大规模数据的“分块读取”接口配套使用。MapDataset把某个 Transform 应用到源数据集后得到的组合数据集是.map()链的结果类型见 datasets/map.h。SharedBatchDataset在多线程多 worker场景下共享 batch 结果的包装见 datasets/shared.h。内置数据集 MNIST直接构造即用示例见上文Normalize参数取自 MNIST 官方均值/标准差 0.1307、0.3081。四、DataLoader批量化、打乱与并行加载DataLoader 是训练循环对数据的主入口负责从数据集取样本、按采样器组织索引、做批量化并在workers 1时并行加载。文档页 docs/cpp/source/api/data/dataloader.md 将其定位为“迭代训练数据的主要接口”。4.1 make_data_loader 的三种重载由 dataloader.h 源码可知make_data_loader依据“数据集是否有状态”与“是否显式传 sampler”提供三套重载Stateless 数据集 显式 sampler返回unique_ptrStatelessDataLoaderDataset, Sampler把数据集、采样器、选项三者组装Stateless 数据集 仅 options最常用从dataset.size()推断样本数并默认构造RandomSampler模板默认参数Sampler samplers::RandomSampler若数据集未实现size()optional为空会触发TORCH_CHECK报错Stateful 数据集 options返回unique_ptrStatefulDataLoaderDataset此时由数据集自己管理批量逻辑不需要 sampler。这些重载通过Dataset::is_stateful与std::is_constructible_vSampler, size_t在编译期做 SFINAE 匹配。4.2 DataLoaderOptions 的常用配置DataLoaderOptions见 dataloader_options.h以流式 setter 方式配置 DataLoader。文档示例中出现的两个核心参数batch_size(N)每个 mini-batch 的样本数。示例用 64。它同时作为next(batch_size)传给 sampler见下节 Sampler 接口Sampler 据此决定每次返回多少个索引。workers(N)并行 worker 线程数。示例分别出现workers(4)与workers(2)。该值决定内部DataShuttle启动多少个后台线程进行数据预取与装配相关实现在 dataloader/stateful.h、dataloader/stateless.h 及 detail/data_shuttle.h。调大 worker 数可让加载与计算重叠但应不超过数据源/IO 的实际瓶颈。其余文档以 Doxygen 暴露的成员还包括超时、队列容量、跨设备CUDA加载开关等以头文件中实际成员为准。4.3 无状态 / 有状态两类 DataLoader文档区分了两个类StatelessDataLoader面向“必须由外部 sampler 组织 batch”的Dataset对应前文重载 1、2StatefulDataLoader面向自行管理批量逻辑的StatefulDataset对应重载 3。它们都派生自DataLoaderBase迭代返回BatchType。无论哪种用户代码都以“for (auto batch : *data_loader)”的方式消费数据。4.4 完整训练示例文档原样继承dataloader 文档给出了含模型与优化器的完整训练骨架需自行定义模型Net例如用torch::nn组装两层卷积加全连接、最后接LogSoftmax#include torch/torch.h int main() { // Load dataset auto dataset torch::data::datasets::MNIST(./data) .map(torch::data::transforms::Normalize(0.1307, 0.3081)) .map(torch::data::transforms::Stack()); // Create data loader auto data_loader torch::data::make_data_loader( std::move(dataset), torch::data::DataLoaderOptions().batch_size(64).workers(2)); // Create model and optimizer auto model std::make_sharedNet(); auto optimizer torch::optim::Adam(model-parameters(), 0.001); // Training loop for (size_t epoch 1; epoch 10; epoch) { for (auto batch : *data_loader) { optimizer.zero_grad(); auto output model-forward(batch.data); auto loss torch::nll_loss(output, batch.target); loss.backward(); optimizer.step(); } } }与 Python 侧DataLoader的体验一致每轮for都从迭代器拿下一个 batchbatch.data/batch.target形状由batch_size与样本张量形状共同决定文档索引页标注 MNIST 场景下为[64, 1, 28, 28]与[64]。一个 DataLoader 在遍历完一个 epoch 后会自动复位内部 sampler调用其reset()因此外层再加一层 epoch 循环即可得到多轮训练。五、Sampler控制样本访问顺序Sampler 决定 DataLoader 取数据的索引序列是“训练要打乱、评估要顺序”的关键旋钮。samplers 文档页 docs/cpp/source/api/data/samplers.md 给出了完整类族。5.1 基类接口从源码 samplers/base.h 看Sampler接口非常精简template typename BatchRequest std::vectorsize_t class Sampler { public: // 复位内部状态通常在新 epoch 前调用可传入新的数据集大小 virtual void reset(std::optionalsize_t new_size) 0; // 返回下一批索引若本 epoch 已耗尽则返回空 optional virtual std::optionalBatchRequest next(size_t batch_size) 0; // 序列化 / 反序列化用于 checkpoint virtual void save(serialize::OutputArchive archive) const 0; virtual void load(serialize::InputArchive archive) 0; };注意无状态 DataLoader 在 batch 装配时会把 sampler 给出的索引数组交给dataset.get_batch()因此对Dataset随机访问型来说RandomSampler每次只需维护一个打乱后的索引池即可。5.2 各类采样器选型采样器行为典型用途SequentialSampler按 0..N-1 顺序访问评估/测试集或对顺序敏感的场景见 samplers/sequential.hRandomSampler每个 epoch 以随机顺序访问保证各 epoch 次序不同训练集也是make_data_loader不带 sampler 调用时的默认选择见 samplers/random.hDistributedSampler基类DistributedRandomSampler分布式训练中为每个进程划分互不重叠的数据子集多机/多卡数据并行见 samplers/distributed.hDistributedSequentialSampler分布式下的顺序分片变体分布式评估StreamSampler面向流式无限数据源配合StreamDataset使用见 samplers/stream.h若需要完全自定义遍历策略可以直接继承Sampler并实现reset、next、save、load四个方法再把它作为第二个参数传给三参版make_data_loader(dataset, sampler, options)。六、Transform预处理与数据增广Transform 对样本施加归一化、增广等预处理并通过数据集的.map()自由串联。transforms 文档页 docs/cpp/source/api/data/transforms.md 覆盖了基础类与内置变换。6.1 两条变换抽象Transform 与 BatchTransform源码 transforms/base.h 清晰定义了两层抽象BatchTransformInputBatch, OutputBatch作用于整批数据核心是纯虚apply_batch(input_batch)TransformInput, Output继承BatchTransformstd::vectorInput, std::vectorOutput只需实现逐样本的apply(Input)基类默认apply_batch会遍历整批逐样本调用applystd::vectorOutput apply_batch(std::vectorInput input_batch) override { std::vectorOutput output_batch; output_batch.reserve(input_batch.size()); for (auto input : input_batch) { output_batch.push_back(apply(std::move(input))); } return output_batch; }这种“逐样本变换默认即逐批变换”的设计与Dataset/BatchDataset的关系完全对称需要整批协同的变换如Stack则直接针对批实现。6.2 内置变换一览变换作用实现位置NormalizeScalar用给定均值与标准差对张量做(x - mean) / std标准化transforms/tensor.hStack把批内的多个张量沿新维堆叠成一个张量如 N 个[1,28,28]→[N,1,28,28]transforms/stack.hLambda用任意可调用对象包装成变换transforms/lambda.hTensorTransform面向张量级变换的基类transforms/tensor.hTensorLambda/BatchLambdaLambda 的张量/整批变体transforms/lambda.hCollate相关批装配collation辅助transforms/collate.htorch/data/transforms.h聚合入口见 torch/csrc/api/include/torch/data/transforms.h。6.3 链式组合与实战要点文档反复强调的用法是变换之间、以及变换与数据集的组合auto dataset torch::data::datasets::MNIST(./data) .map(torch::data::transforms::Normalize(0.1307, 0.3081)) .map(torch::data::transforms::Stack());链路语义如下MNIST的get(i)先取出第 i 张原始图像与标签MapDatasetMNIST, Normalize的get(i)内部先取原始样本再调用Normalize的apply完成标准化外层再套Stackget_batch拿到一批已归一化的样本后Stack::apply_batch将其堆叠成单个四维张量因此 DataLoader 产出的batch.data才具有[64, 1, 28, 28]的形状。变换链的顺序很重要Normalize逐样本操作放在Stack之前既可节省内存也符合其“逐样本张量”输入约定把二者调换则Normalize需要作用在堆叠后的批量张量上语义与实现都会改变。同样地随机水平翻转、随机裁剪等增广变换作为自定义Transform应接在Normalize之前或之后取决于你的数据语义——文档中所有示例都遵循“先 Normalize、再 Stack”这一约定。七、四类组件的组合关系小结把索引页的组件概览与各子文档合并可以得到一张贯穿始终的组合图原始数据源MNIST 等内置数据集 / 自定义 Dataset │ Dataset::get(i) / get_batch(indices) ▼ . map(Normalize) → MapDataset 逐样本预处理 . map(Stack) → MapDataset 整批堆叠产出单个张量 │ ▼ make_data_loader(dataset, DataLoaderOptions{batch_size, workers}) │ 内部sampler默认 RandomSampler产出索引 → get_batch → 多 worker 预取 ▼ for (auto batch : *data_loader) → batch.data / batch.targetDataset定义“数据长什么样、如何取一个样本”get/sizeTransform定义“取到的样本如何加工”逐样本apply或整批apply_batch用.map()织入数据集Sampler定义“按什么顺序取”训练随机 / 评估顺序 / 分布式分片DataLoader负责把它们组装起来做索引→取数→加工→装配→多 worker 并行的流水线调度成为训练循环中唯一的数据入口。八、进一步阅读与源码索引官方 API 文档目录docs/cpp/source/api/data/index.md及其四个子页面 datasets.md、dataloader.md、samplers.md、transforms.md其中每个类都以 Doxygen 块列出全部公开成员。头文件聚合入口torch/csrc/api/include/torch/data.h。想深入源码torch/csrc/api/include/torch/data/目录下datasets/、dataloader/、samplers/、transforms/、detail/子目录分别对应四类组件及其底层队列/传输实现如 detail/data_shuttle.h、detail/queue.h。按文档给出的头文件清单任何一项能力只需一行 include 即可获得而把它接入现有训练工程时最关键的三步永远是定义/选用Dataset→ 用.map()串起Transform流水线 → 用make_data_loader加上batch_size与workers交给训练循环。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考