首页
/ PyTorch C++ 前端 Data API Transforms 完全指南:预处理、归一化与链式数据管线

PyTorch C++ 前端 Data API Transforms 完全指南:预处理、归一化与链式数据管线

2026-09-07 20:22:49作者:宣聪麟

导读

torch::data::transforms 是 PyTorch C++ 前端数据加载体系中的"预处理层",负责对数据集产出的样本(Example)或整个 batch 施加归一化、增强、堆叠(Stack)等变换,并允许通过 .map() 把多个变换链式组合成一个干净的数据管线。本文以 docs/cpp/source/api/data/transforms.md 的 API 文档为骨架,结合仓库中 transforms 头文件的真实实现,逐层拆解 TransformBatchTransformTensorTransformNormalizeStackLambda 家族与 Collate 的继承关系、批处理语义和广播规则,并给出可直接照搬的 MNIST 归一化+堆叠示例,帮助你编写自己的自定义变换并在 DataLoader 中正确使用。

一、Transforms 在 C++ Data API 中的定位

在 Python 侧,torchvision.transforms 提供了成熟的图像预处理工具箱;在 C++ 前端中,等价的能力由 torch::data::transforms 命名空间提供,源码集中在 torch/csrc/api/include/torch/data/transforms/。它的职责可以用官方文档的一句话概括:

Transforms apply preprocessing to data samples, such as normalization or augmentation. They can be chained using the .map() method on datasets.

即两类能力:

  1. 单样本级预处理:归一化、缩放、颜色抖动等,作用于单个 Example
  2. 整批聚合:把一批零散样本 stack 成单个张量(collation),供矩阵式算子消费。

transform 本身不直接执行,而是通过数据集上的 .map(transform) 返回一个包装数据集(MapDataset),真正取数据时才执行变换。这种"惰性组合"让流水线既可以复用又可组合,是 C++ Data API 的一个核心设计。测试 test/cpp/api/integration.cpp 中即可看到真实代码在数据集上链式使用 transforms::Stack<>() 的写法。

二、两层抽象基类:Transform 与 BatchTransform

torch/csrc/api/include/torch/data/transforms/base.h 中定义了两层最基础的抽象:

BatchTransform:面向"批"的变换

/// A transformation of a batch to a new batch.
template <typename InputBatch, typename OutputBatch>
class BatchTransform {
 public:
  using InputBatchType = InputBatch;
  using OutputBatchType = OutputBatch;

  virtual ~BatchTransform() = default;

  /// Applies the transformation to the given `input_batch`.
  virtual OutputBatch apply_batch(InputBatch input_batch) = 0;
};

BatchTransform 是整个体系的根接口,模板参数分别是输入批与输出批的类型。子类只需实现唯一的纯虚函数 apply_batch()。注意 InputBatchType / OutputBatchType 这两个嵌套 typedef 并非可有可无——MapDataset 与自由函数 map() 依赖它们做静态类型检查(见下文第五节)。

Transform:把"单样本变换"包装成批变换

/// A transformation of individual input examples to individual output examples.
template <typename Input, typename Output>
class Transform
    : public BatchTransform<std::vector<Input>, std::vector<Output>> {
 public:
  using InputType = Input;
  using OutputType = Output;

  /// Applies the transformation to the given `input`.
  virtual OutputType apply(InputType input) = 0;

  /// Applies the `transformation` over the entire `input_batch`.
  std::vector<Output> apply_batch(std::vector<Input> input_batch) override {
    std::vector<Output> output_batch;
    output_batch.reserve(input_batch.size());
    for (auto&& input : input_batch) {
      output_batch.push_back(apply(std::move(input)));
    }
    return output_batch;
  }
};

从源码结构看,Transform 是对"单样本 → 单样本"逻辑的建模:你只需实现 apply(Input),而基类已经用默认的 apply_batch() 把它映射到整批——遍历输入向量、逐个调用 apply、结果收集进输出向量。这正是批处理语义统一在 BatchTransform::apply_batch 接口下的原因:数据加载器(DataLoader)永远只面对批量接口,Transform 只是把单样本逻辑翻译成批逻辑的适配器

三、面向张量样本的 TensorTransform

真实数据集中最常见的样本形态是 Example<Tensor, Target>:一张图片张量加一个标签。文档中的 torch/csrc/api/include/torch/data/example.h 定义了:

template <typename Data = at::Tensor, typename Target = at::Tensor>
struct Example {
  Data data;
  Target target;
};

torch/csrc/api/include/torch/data/transforms/tensor.h 中的 TensorTransform 则是专门面向这种形态的派生基类:

template <typename Target = Tensor>
class TensorTransform
    : public Transform<Example<Tensor, Target>, Example<Tensor, Target>> {
 public:
  using E = Example<Tensor, Target>;
  ...
  /// Transforms a single input tensor to an output tensor.
  virtual Tensor operator()(Tensor input) = 0;

  /// Implementation of `Transform::apply` that calls `operator()`.
  OutputType apply(InputType input) override {
    input.data = (*this)(std::move(input.data));
    return input;
  }
};

两个关键设计值得注意:

  • operator() 契约:子类不再实现 apply(Example),而是实现 operator()(Tensor),把关注点收敛到"一个张量如何处理"。
  • 目标保留语义:基类 apply 只改写 input.datainput.target 原样不动。这让归一化、几何增强等只针对输入的变换天然无需关心标签,避免样板代码。

Example 还有一个针对无标签场景的约定:example::NoTarget(即 void),配套的特化 Example<Data, NoTarget> 提供了到 Data& 的隐式转换,其别名 TensorExample 表示"只有数据没有标签"的样本。TensorTransform<Target> 的模板参数因此默认是 Tensor,若数据无标签可指定 Target = example::NoTarget

四、内建变换:Normalize 及其广播语义

torch/csrc/api/include/torch/data/transforms/tensor.h 尾部,Normalize 是继承 TensorTransform<Target> 的内建归一化变换,其完整实现为:

template <typename Target = Tensor>
struct Normalize : public TensorTransform<Target> {
  Normalize(ArrayRef<double> mean, ArrayRef<double> stddev)
      : mean(torch::tensor(mean, torch::kFloat32)
                 .unsqueeze(/*dim=*/1)
                 .unsqueeze(/*dim=*/2)),
        stddev(torch::tensor(stddev, torch::kFloat32)
                   .unsqueeze(/*dim=*/1)
                   .unsqueeze(/*dim=*/2)) {}

  torch::Tensor operator()(Tensor input) override {
    return input.sub(mean).div(stddev);
  }
  torch::Tensor mean, stddev;
};

源码揭示了三个实现细节:

  1. 公式out = (input - mean) / stddev,直接由 sub/div 张量算子完成,运算后 target 不受影响(由 TensorTransform::apply 保证)。
  2. 均值/方差被预转为 kFloat32 张量并 reshape 为 [C, 1, 1]:构造函数对 torch::tensor(mean, torch::kFloat32) 连续做 unsqueeze(1).unsqueeze(2)。这使传入的逐通道 mean/std(长度为通道数)能够按通道方向与 [N, C, H, W] 的输入自动广播。头文件注释也特别强调:mean 与 stddev 可以是任何能与输入张量广播的形状,"like single scalars"。
  3. 输入类型是 ArrayRef<double>:因此既支持多通道数组 {0.485, 0.456, 0.406},也支持单个标量。

Normalize 的实战用法

文档给出的 MNIST 示例同时展示了这两种传参形态:

// 简单形态:对单通道灰度图用单一标量 mean/std
auto dataset = torch::data::datasets::MNIST("./data")
    .map(torch::data::transforms::Normalize<>(0.5, 0.5))
    .map(torch::data::transforms::Stack<>());
// 更贴近 MNIST 官方统计值的形态
auto dataset = torch::data::datasets::MNIST("./data")
    .map(torch::data::transforms::Normalize<>(0.1307, 0.3081))
    .map(torch::data::transforms::Stack<>());

其中 Normalize<> 的空尖括号取默认模板参数 Target = Tensor,对应 MNIST 输出的带标签 Example<>(data 与 target 均为张量)。0.5 是 MNIST 训练集常用的简单归一化值;0.1307 / 0.3081 是 MNIST 官方数据集统计出的像素均值与标准差,是文档中直接给出的可复现数值。MNIST 数据集本体定义在 torch/csrc/api/include/torch/data/datasets/mnist.h

五、整批聚合:Stack、Collation 与 Collate

Stack:把样本批拼成单个张量

Normalize 处理的是"每个样本内",而 Stack 处理的是"样本之间"。它的实现位于 torch/csrc/api/include/torch/data/transforms/stack.h,针对不同样本形态提供了两个特化:

template <typename T = Example<>>
struct Stack;

/// Stacks all data tensors into one tensor, and all target tensors into one tensor.
template <>
struct Stack<Example<>> : public Collation<Example<>> {
  Example<> apply_batch(std::vector<Example<>> examples) override {
    std::vector<torch::Tensor> data, targets;
    data.reserve(examples.size());
    targets.reserve(examples.size());
    for (auto& example : examples) {
      data.push_back(std::move(example.data));
      targets.push_back(std::move(example.target));
    }
    return {torch::stack(data), torch::stack(targets)};
  }
};

/// Stacks all data tensors into one tensor (no targets).
template <>
struct Stack<TensorExample>
    : public Collation<Example<Tensor, example::NoTarget>> {
  TensorExample apply_batch(std::vector<TensorExample> examples) override {
    std::vector<torch::Tensor> data;
    data.reserve(examples.size());
    for (auto& example : examples) {
      data.push_back(std::move(example.data));
    }
    return torch::stack(data);
  }
};

实现要点:

  • 带标签版本把 data 与 target 分开收集,分别调用 torch::stack,返回一个"聚合后的单一 Example";
  • 无标签版本(TensorExample)只堆叠 data,直接返回聚合张量;
  • 两个特化都利用了 std::move 转移样本,避免深拷贝开销。

需要强调的是:Stack 的输入是整批 std::vector<Example<>>,输出是一个 Example<>,因此它本质上是一个 collation(聚合)而非逐样本变换。这也解释了为什么示例代码里它总是排在 .map() 链的最末端——先把样本逐张归一化,最后统一堆叠成一个 mini-batch。

Collation 与 Collate

聚合语义在 torch/csrc/api/include/torch/data/transforms/collate.h 中被形式化为两个别名:

/// A `Collation` is a transform that reduces a batch into a single value.
template <typename T, typename BatchType = std::vector<T>>
using Collation = BatchTransform<BatchType, T>;

/// `Collate` is the lambda version of `Collation`.
template <typename T, typename BatchType = std::vector<T>>
using Collate = BatchLambda<BatchType, T>;
  • Collation<Example<>> 就是把 std::vector<Example<>> 归约为一个 Example<> 的批变换,Stack 正是它的特例;
  • Collate 则允许直接传入一个 lambda 完成自定义聚合,无需写一个类。collate.h 自带的最小示例是:
using namespace torch::data;

auto dataset = datasets::MNIST("path/to/mnist")
    .map(transforms::Collate<Example<>>([](std::vector<Example<>> e) {
      return std::move(e.front());   // 自定义聚合:只取批中第一个样本
    }));

如果你需要按自己规则拼装 batch(例如先 pad 到相同长度再堆叠),Collate 是比手写子类更轻量的入口。

六、Lambda 家族:把自由函数变成变换

日常中很多预处理只是一行计算,不值得为此定义一个类。 torch/csrc/api/include/torch/data/transforms/lambda.h 提供了三个基于 std::function 的"函数对象适配器":

继承自 包装的函数签名 作用层级
Lambda<Input, Output> Transform<Input, Output> Output(Input) 逐个样本
TensorLambda<Target> TensorTransform<Target> Tensor(Tensor) 张量样本的 data
BatchLambda<Input, Output> BatchTransform<Input, Output> OutputBatch(InputBatch) 整个批

它们的实现都是"存一个 std::function,在 apply/apply_batch/operator() 里调用它":

template <typename Input, typename Output = Input>
class Lambda : public Transform<Input, Output> {
 public:
  using FunctionType = std::function<Output(Input)>;
  explicit Lambda(FunctionType function) : function_(std::move(function)) {}
  OutputType apply(InputType input) override {
    return function_(std::move(input));
  }
 private:
  FunctionType function_;
};
template <typename Target = Tensor>
class TensorLambda : public TensorTransform<Target> {
 public:
  using FunctionType = std::function<Tensor(Tensor)>;
  explicit TensorLambda(FunctionType function)
      : function_(std::move(function)) {}
  Tensor operator()(Tensor input) override {
    return function_(std::move(input));
  }
 private:
  FunctionType function_;
};

选型建议:

  • 对"普通样本结构"做逐样本变换(如把一个 Example<int, int> 映射为浮点样本)用 Lambda
  • 对典型"张量 in → 张量 out"的图像增强用 TensorLambda(它能无缝接在 Normalize 同一条链上,因为二者都是 TensorTransform);
  • 对必须整体看待批的算子(如动态 pad、重排、batch 内归一化统计)用 BatchLambda,它直接操作 std::vector<Example<>> 这类批结构。

七、链式组合的底层机制:.map() 与 MapDataset

文档强调"Transforms can be chained together using .map()",而这一机制在 torch/csrc/api/include/torch/data/datasets/map.h 中有明确的源码支撑。

MapDataset 保存一个源数据集与一个变换,重写 get_batch() 为"先取源 batch,再对 batch 施加 transform_.apply_batch(...)":

template <typename SourceDataset, typename AppliedTransform>
class MapDataset : public BatchDataset<...> {
  OutputBatchType get_batch(BatchRequestType indices) override {
    return get_batch_impl(std::move(indices));
  }
  ...
 private:
  // stateless 版本:直接变换取出的 batch
  OutputBatchType get_batch_impl(BatchRequestType indices) {
    return transform_.apply_batch(dataset_.get_batch(std::move(indices)));
  }
  SourceDataset dataset_;
  AppliedTransform transform_;
};

对应的自由函数 map(dataset, transform) 在返回前做了一道编译期防线:

static_assert(
    std::is_same_v<... /* 源数据集 BatchType */, typename TransformType::InputBatchType>,
    "BatchType type of dataset does not match input type of transform");

也就是说,如果前一个环节的输出类型与下一个 transform 的输入类型不匹配,会在编译期报错而不是运行期崩溃——这是链式组合能安全成立的根本保证。

进一步看,map.h 还处理了有状态数据集的差异:若源数据集是有状态的(SourceDataset::is_stateful == true),MapDatasetBatchType 会变成 std::optional<...>get_batch 遵循"Optional.map"语义——源返回空 optional 则透传空,否则才施加变换;同时暴露 reset() 将状态透传给底层数据集。这也是为何 .map() 同时适用于顺序(stateful)与随机访问(stateless)两类数据集。

正是"每 .map() 一次就包一层 MapDataset"的递归结构,让文档中的 MNIST 三段式写法成立:

auto dataset = torch::data::datasets::MNIST("./data")
    .map(torch::data::transforms::Normalize<>(0.1307, 0.3081))   // 逐样本归一化
    .map(torch::data::transforms::Stack<>());                     // 整批堆叠成 mini-batch

执行语义为:MNIST 取出一批原始样本 → Normalize(作为 Transform)在批内逐个样本对 data 做 (x-0.1307)/0.3081Stack 把 data 与 target 分别 torch::stack 成一个 batch 张量。

八、端到端实战:把 transform 链接入 DataLoader

将上面链条作为 BatchDataset 传给 C++ 的 make_data_loader,就能得到可迭代的数据源;此处给出一个完整的 MNIST 预处理组合作为参考:

#include <torch/torch.h>

// 1) 数据变换链:归一化 + 堆叠(亦可把数据增强的 TensorLambda 插在二者之间)
auto dataset = torch::data::datasets::MNIST("./data")
    .map(torch::data::transforms::Normalize<>(0.1307, 0.3081))
    .map(torch::data::transforms::Stack<>());

// 2) 交给 DataLoader 迭代训练
//    (C++ Data API 中通过 make_data_loader 与 DataLoaderOptions 配置 batch_size 等)
// auto loader = torch::data::make_data_loader(std::move(dataset),
//     torch::data::DataLoaderOptions().batch_size(64));
// for (auto& batch : *loader) {
//   auto data = batch.data;   // shape: [64, 1, 28, 28]
//   auto target = batch.target; // shape: [64]
// }

工程上的注意点:

  • 变换顺序决定语义:先做逐样本归一化再做 Stack。若反序,Stack 输出的聚合张量会被当成一个样本传入 Normalize,广播结果将不符合预期。
  • 无标签数据用 TensorExample 分支:若数据没有标签,优先让数据集产出 Example<Tensor, example::NoTarget>,从而走 Stack<TensorExample> 特化,最后返回聚合张量而非 Example
  • 批大小与维度一致性torch::stack 要求批内各样本形状一致;对变长序列需先用 Collate 自定义 pad 再堆叠。
  • 变换可以继续叠加:因为 .map() 的产物仍是 BatchDataset,可以在 Stack 之后再接面向 batch 张量的 BatchLambda,完成 batch 级后处理。

九、编写自定义 Transform 的决策路径

结合以上源码分析,自定义变换只需回答"作用在什么粒度":

你要做的变换 继承/使用的类型 必须实现
图像归一化/增强:Tensor → Tensor,样本是 Example<Tensor, Target> TensorTransform<Target> operator()(Tensor)
任意样本结构逐样本映射 Transform<Input, Output> apply(Input)
用自由函数实现上述两者 Lambda / TensorLambda 传入 std::function
整批聚合/重排/pad BatchTransform<InputBatch, OutputBatch>Collate<T> apply_batch(...) 或传入 lambda
堆叠张量批 Stack<> / Stack<TensorExample> 直接复用内建实现

依据源码 base.htensor.h,一个自定义归一化变换的最小形态可以这样书写并直接接入 .map() 链:

// 自定义:把像素范围缩放到 [-1, 1]
struct ScaleToUnit : torch::data::transforms::TensorTransform<> {
  torch::Tensor operator()(torch::Tensor input) override {
    return input.mul(2.0).sub(1.0);
  }
};

auto dataset = torch::data::datasets::MNIST("./data")
    .map(ScaleToUnit{})   // 与内建 Normalize 同属 TensorTransform,可互换、可串接
    .map(torch::data::transforms::Stack<>());

编写时务必记住 TensorTransform::apply 默认只改 data 保留 target(见 tensor.h),因此若你的变换需要连带修改标签(如标签平滑、样本重采样),就不应继承 TensorTransform,而应改用直接操纵 ExampleTransformLambda

十、小结

  • C++ Data API 的 transforms 采用"单样本抽象 + 批量执行"双层设计:Transform(含 TensorTransform)定义样本级语义,BatchTransform 提供统一的批接口,二者在 base.h 中被桥接。
  • Normalize 把 mean/std 预转换为 [C,1,1] 浮点张量后执行 sub/div,天然支持标量与逐通道两种广播用法(tensor.h)。
  • Stack 是典型的 collation:把整批 Example 的 data/target 分别 torch::stack,是数据管线的收尾一环(stack.h)。
  • Lambda / TensorLambda / BatchLambda 让函数对象无缝融入流水线(lambda.h),Collate 提供自定义聚合的最短路径(collate.h)。
  • .map() 通过 MapDataset 逐层包装数据集并在 get_batch() 时调用 apply_batch(),配合 static_assert 保证链上类型自洽(map.h)。

把归一化、增强、堆叠这类"样板预处理"固化为可组合的 transform,并把类型检查前置到编译期,正是这套 C++ Data API 让数据管线既简洁又可维护的关键所在。进一步查阅 API 总览可回到 docs/cpp/source/api/data/ 下的相关文档,结合仓库头文件逐行对照阅读效果最佳。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.13 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
897
5.8 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
529
593
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
915
1.83 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.58 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.35 K
1.46 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.01 K
515
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
547
388