PyTorch C++ 神经网络模块(torch::nn)API 指南:从 Module 基类到自定义模型构建
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
torch::nn是 PyTorch 为 C++ 提供的神经网络构建模块命名空间,与 Python 侧torch.nn一一对应,可用于在 C++ 中搭建模型、自定义层、将 Python 模型移植到 C++ 做生产级推理,乃至纯 C++ 完成端到端训练。本文将围绕 torch::nn 官方 C++ API 文档 的核心内容展开,结合本仓库头文件实现(如 module.h 与 pimpl.h),深入讲解 PIMPL 设计模式、Module 基类注册与遍历机制、头文件组织、模块分类清单,并给出可编译的自定义模型示例,帮助读者在 C++ 侧写出与 Python 等价且易于移植的模型代码。
torch::nn 概览:C++ 侧的神经网络积木
torch::nn命名空间提供了与 Pythontorch.nn模块镜像对应的神经网络构建单元。它采用 PIMPL(Pointer to Implementation,指向实现的指针)设计:用户直接面对的类型(如Conv2d)本质上是“句柄”,内部包装了对应的Conv2dImpl实现类。这种设计使外层句柄可以被安全复制、按值存储于容器或作为类成员,而真正持有参数和状态的是通过std::shared_ptr管理的Impl对象。
从源码看,这套实现位于仓库的 torch/csrc/api/include/torch/nn 目录:其中 module.h 定义基类Module,pimpl.h 定义ModuleHolder(PIMPL 包装器)与Module之间按std::shared_ptr共享语义协作;modules.h 汇总所有模块实现,options.h 汇总所有 Options 配置结构体。
何时使用 torch::nn
原文档明确列出了四类典型场景:
- 在 C++ 中构建神经网络模型:无需 Python 解释器参与,纯 C++ 即可定义前向计算;
- 创建自定义层与模块:通过继承
Module并重写forward(),可把任意算子封装为可复用、可组合的层; - 将 Python 模型移植到 C++ 进行生产推理:Python 训练 + C++ 部署是常见的工程路径,
torch::nn提供与torch.nn对齐的语义,便于平移; - 完全在 C++ 中训练模型:结合 torch::optim 与数据加载 API(torch::data),可构建不依赖 Python 的训练流程。
基本用法示例
原文档给出了一个从“定义模型”到“前向计算”的最小闭环:
#include <torch/torch.h> // Define a simple model struct Net : torch::nn::Module { torch::nn::Conv2d conv1{nullptr}; torch::nn::Linear fc1{nullptr}; Net() { conv1 = register_module("conv1", torch::nn::Conv2d( torch::nn::Conv2dOptions(1, 32, 3).stride(1).padding(1))); fc1 = register_module("fc1", torch::nn::Linear(32 * 28 * 28, 10)); } torch::Tensor forward(torch::Tensor x) { x = torch::relu(conv1->forward(x)); x = x.view({-1, 32 * 28 * 28}); return fc1->forward(x); } }; // Create and use the model auto model = std::make_shared<Net>(); auto input = torch::randn({1, 1, 28, 28}); auto output = model->forward(input);该示例有四个值得注意的编码惯例:
- 成员声明为“默认空句柄”:
torch::nn::Conv2d conv1{nullptr};。因为Conv2d是ModuleHolder类型,可用nullptr先占位,待构造函数中真正创建后再赋值; - 子模块必须
register_module:register_module返回std::shared_ptr<Conv2d>,可直接赋值回成员(隐式转换为句柄),同时把子模块登记进父模块的children_表,保证后续parameters()、to()、save()能递归生效; - 通过
->调用成员:句柄重载了operator->,等价于操作内部Impl,因此写conv1->forward(x)而非conv1.forward(x); - 输入输出为
torch::Tensor:forward的签名与 Python 中Module.forward对应,但类型显式、维度以view显式拉平。
PIMPL 模式与 ModuleHolder:理解“句柄 + Impl”的设计
要真正用好torch::nn,必须先理解它“为什么长这样”。PIMPL 的核心是把类的实现细节藏到不透明指针后面,从而获得稳定的 ABI 与更好的封装性。
源码层面,ModuleHolder<Contained>承担这一职责。在 pimpl.h 中可以看到它声明了类型别名using ContainedType = Contained;,即句柄所指的实现类型;并重载了operator->与const operator->,让句柄在用法上像裸指针一样透明。基类 module.h 中以Module为核心对象,而register_module的模板重载同时接受裸std::shared_ptr<ModuleType>与ModuleHolder<ModuleType>,二者最终都登记为shared_ptr形式存入children_(一个OrderedDict<std::string, std::shared_ptr<Module>>)。
由此形成三层关系:
| 层次 | 类型 | 作用 |
|---|---|---|
| 句柄层 | torch::nn::Conv2d等ModuleHolder | 值语义外壳,可复制、可空初始化 |
| 实现层 | Conv2dImpl等*Impl子类 | 继承torch::nn::Module,真正持有weight/bias参数并实现forward |
| 基类层 | torch::nn::Module | 参数/缓冲区/子模块注册与递归遍历、to、train/eval、序列化等通用能力 |
一个直观后果是:当你在自定义Module中声明torch::nn::Linear fc1{nullptr};时,成员是“句柄”而非“Impl”,因此调用需要fc1->forward(x);当你通过register_module返回句柄对应的shared_ptr时,所有权由父模块统一管理,避免悬垂引用。
Module 基类:一切模块的“递归树”根节点
官方头文件对 torch::nn::Module 的定位是“PyTorch 中所有模块的基类”,其设计“主要基于 Python API”(源码注释明确写道The design and implementation of this class is largely based on the Python API)。Module表示某个函数或算法的实现抽象,可能携带持久化数据;模块可以递归嵌套子模块,构成一棵递归树。
Module区分三类持久化数据(见 module.h 源码注释):
- Parameters(参数):记录梯度、通常在反向阶段被更新的张量,例如
Linear的weight; - Buffers(缓冲区):不记录梯度、通常在前向阶段被更新的状态,例如
BatchNorm中的running_mean与running_var; - 其它任意状态:非张量、供实现或配置使用的额外数据。
注册 API:register_module / register_parameter / register_buffer
这三个方法构成模块“告知基类自己有什么”的核心入口,通常在各模块构造函数中调用:
register_parameter(std::string name, Tensor tensor, bool requires_grad = true):登记学习参数,返回该参数的引用。源码示例weight_ = register_parameter("weight", torch::randn({A, B}));展示其典型用法;也允许注册未定义张量(等价于 Python 侧的None参数);register_buffer(std::string name, Tensor tensor):登记不参与梯度更新的状态,如mean_ = register_buffer("mean", torch::empty({num_features_}));;register_module(...):登记子模块。源码实现会对名称做两项硬校验(见 module.h):名称不能为空,且不能包含点号.(Submodule name must not contain a dot),因为点号被保留用作named_modules的层级分隔符。
register_module返回std::shared_ptr<ModuleType>,好处是可以把返回值直接赋给句柄成员,保证“注册表中持有的指针”与“类成员持有的句柄”指向同一个对象——这正是示例中conv1 = register_module(...)的写法依据。
除注册外,基类还提供replace_module(name, module)(替换已注册子模块,常用于微调时替换结构)与unregister_module(name)(注销子模块,不存在则抛异常),细节同样在 module.h。
遍历 API:parameters / named_parameters / modules / children
注册之后,模块树上的所有状态都可被统一遍历:
std::vector<Tensor> parameters(bool recurse = true):返回参数列表;OrderedDict<std::string, Tensor> named_parameters(bool recurse = true):返回带键名(如conv1.weight)的参数表;buffers()/named_buffers():对缓冲区做同样的遍历;std::vector<std::shared_ptr<Module>> modules(bool include_self = true):返回整个子模块层级;当include_self = true时会把this的shared_ptr放在首位——源码用warning强调:只有模块本身存于shared_ptr中时才可传true,否则抛异常;children()/named_children():仅返回直接子模块。
这些遍历能力是后续to()、save()等递归操作的基础。
设备与类型转换:to()
Module提供了三个to()重载,递归作用于全部已注册参数与缓冲区:
to(torch::Device device, torch::Dtype dtype, bool non_blocking = false)to(torch::Dtype dtype, bool non_blocking = false)to(torch::Device device, bool non_blocking = false)
其实现逻辑清晰可见于 module.h:先对所有子模块递归调用to(),再对本模块(recurse=false)的参数与缓冲区分别执行set_data(tensor.to(...))。这种“先孩子、后自身”的顺序保证整棵模块树一致迁移。non_blocking在源为锁页内存且目标为 GPU(或相反)时使拷贝相对主机异步执行,其余情况无效果。
典型调用如module->to(torch::kCUDA)(全部参数移到 GPU)或module->to(torch::kFloat32)(统一 dtype),与 Pythonmodule.to()对齐。
训练与评估模式:train() / eval() / is_training()
每个模块内部有一个布尔状态is_training_{true}(见 module.h),决定模块处于训练模式还是评估(推理)模式:
virtual void train(bool on = true):进入训练模式(递归作用于子模块);void eval():等价于train(false),源码注释明确“不要重写eval(),应重写train()本身”;virtual bool is_training() const noexcept:查询当前模式。
BatchNorm与Dropout是依赖该状态切换行为路径的典型模块(源码注释明确点名二者)。因此推理前调用model->eval()、训练中调用model->train()是必须养成的习惯,否则归一化统计与随机失活行为将与预期不符。
梯度清零:zero_grad()
virtual void zero_grad(bool set_to_none = true)递归将每个已注册参数的grad置零。set_to_none = true时直接置为 None 而非零张量(与 Python 侧行为一致),有利于节省内存、加快后续反向。
序列化:save() / load()
Module通过序列化归档对象完成状态存取:
virtual void save(serialize::OutputArchive& archive) const;virtual void load(serialize::InputArchive& archive);
源码注释指出:若模块含不可序列化的子模块(例如nn::Functional),保存时会跳过它;加载时同样不检查该类子模块是否存在于归档中。与之配套,命名空间级还重载了流操作符operator<<(OutputArchive&, const std::shared_ptr<nn::Module>&)与operator>>(...),便于把模型状态写入 torch 归档文件。C++ 序列化的更多用法可参考 serialize API 文档。
递归深拷贝:clone() 与 Cloneable
Module声明了virtual std::shared_ptr<Module> clone(const std::optional<Device>& device = std::nullopt) const;,实现模块及所有已注册参数、缓冲区、子模块的递归深拷贝,可附带目标设备。
但源码给出一个重要提醒:直接调用从基类继承的clone()会失败。要获得真正的clone()实现,必须让自定义模块继承模板基类 Cloneable,它会基于具体模块类型生成正确的拷贝逻辑;基类上保留该虚方法仅为提供易用的多态接口。
递归遍历工具:apply() 与 as()
apply()系列方法把函数递归施加到模块自身及每个子模块,且提供多种回调签名变体:接收Module&、const Module&、带键名的变体(键可加name_prefix前缀)、以及shared_ptr变体。典型用法来自 module.h 内嵌示例——统一初始化权重:
void initialize_weights(nn::Module& module) { torch::NoGradGuard no_grad; if (auto* linear = module.as<nn::Linear>()) { linear->weight.normal_(0.0, 0.02); } } MyModule module; module->apply(initialize_weights);其中template <typename ModuleType> ContainedType* as()(及其const版本)做类型安全下转换:对ModuleHolder类型传入时自动取其ContainedType,对裸Impl类型则直接dynamic_cast。配合apply(),可在不修改模块源码的前提下批量执行初始化、Hook 注册或诊断打印。
名称与打印:name() 与 pretty_print()
Module关联一个字符串名(如"Linear"),大多数情况下由运行时类型信息(RTTI)自动推断;若禁用 RTTI,可把显式名称传给基类构造函数explicit Module(std::string name);。pretty_print(std::ostream&)输出模块的易读表示,默认递归打印自身名称及所有子模块;重写该方法可定制打印格式。operator<<(std::ostream&, const nn::Module&)已被声明为友元,便于直接用流输出模块。
头文件导航:从哪里 include 什么
原文档列出的头文件都位于 torch/csrc/api/include/torch 下,在代码中按需 include:
| 头文件 | 内容 | 仓库对应路径 |
|---|---|---|
torch/nn.h | 神经网络主头文件,聚合 include 全部内容 | nn.h |
torch/nn/module.h | Module 基类声明与实现 | module.h |
torch/nn/modules.h | 所有模块实现类汇总 | modules.h |
torch/nn/options.h | 各模块的 Options 配置结构体 | options.h |
torch/nn/functional.h | 函数式 API(无状态算子调用) | functional.h |
在实际工程中,最简单的方式是只写#include <torch/torch.h>,它聚合了torch::nn、torch::optim、torch::data等常用模块。若追求更快的编译,可按上述细粒度头文件裁剪 include。
模块分类全景:13 大类能力总览
torch::nn提供的模块按功能分为多个子页,原文档在 docs/cpp/source/api/nn 目录中以 toctree 组织如下:
- containers(容器模块):
Sequential、ModuleList、ModuleDict、ParameterList、ParameterDict等组合与管理工具,对应源码 torch/nn/modules/container; - convolution(卷积层):
Conv1d/2d/3d与ConvTranspose1d/2d/3d,源码见 conv.h; - pooling(池化层):各类最大/平均/自适应池化,源码见 pooling.h 与 adaptive.h;
- linear(线性层):
Linear、Bilinear、Identity、Flatten、Unflatten,源码见 linear.h; - activation(激活函数):
ReLU、GELU等可带参激活模块,源码见 activation.h; - normalization(归一化层):
BatchNorm、LayerNorm、GroupNorm、InstanceNorm等,源码见 normalization.h、batchnorm.h 与 instancenorm.h; - dropout(随机失活):
Dropout、Dropout2d/3d、AlphaDropout等,源码见 dropout.h; - embedding(嵌入层):
Embedding、EmbeddingBag等,源码见 embedding.h; - recurrent(循环层):
RNN、LSTM、GRU及其多层层级,源码见 rnn.h; - transformer(Transformer 组件):
Transformer、TransformerEncoder/Decoder、TransformerEncoderLayer/DecoderLayer、MultiheadAttention等,源码见 transformer.h、transformerlayer.h 与 transformercoder.h; - loss(损失函数):
L1Loss、MSELoss、CrossEntropyLoss、BCELoss、NLLLoss等,源码见 loss.h; - functional(函数式接口):无状态地调用算子(如
F::relu),不维护参数,源码见 torch/nn/functional; - utilities(工具模块):
ReflectionPad、ZeroPad、PixelShuffle、Upsample等实用层,对应 padding.h、pixelshuffle.h、upsampling.h 等。
Options 配置结构体:模块构造的标准姿势
多数模块不是靠“一堆裸参数”构造,而是通过*Options结构体链式配置。以卷积为例(详见 卷积层文档):
// Create Conv2d: 3 input channels, 64 output channels, 3x3 kernel auto conv = torch::nn::Conv2d( torch::nn::Conv2dOptions(3, 64, 3) .stride(1) .padding(1) .bias(true)); auto output = conv->forward(input); // input: [N, 3, H, W]卷积层Options的核心字段(源码定义于 conv.h,其中ConvOptions统一服务一维/二维/三维卷积)包括:
| 字段 | 含义 | 默认值 |
|---|---|---|
in_channels | 输入通道数 | 必填 |
out_channels | 输出通道数(滤波器个数) | 必填 |
kernel_size | 卷积核尺寸 | 必填 |
stride | 滑动步长 | 1 |
padding | 输入补零量 | 0 |
dilation | 卷积核元素间距 | 1 |
groups | 分组连接数;设为in_channels即深度可分离卷积 | 1 |
构造时先以必填项调用Conv2dOptions(3, 64, 3),随后链式.stride(...)、.padding(...)、.bias(...)覆盖默认值,语义与 Pythontorch.nn.Conv2d(3, 64, 3, stride=1, padding=1, bias=True)完全一致。
转置卷积用于上采样,配置方式相同:
auto conv_transpose = torch::nn::ConvTranspose2d( torch::nn::ConvTranspose2dOptions(64, 32, 4) .stride(2) .padding(1));线性层与之类似(详见 线性层文档):Linear计算仿射变换y = xW^T + b,Bilinear对两个输入做双线性变换,Identity常用于残差直连,Flatten/Unflatten负责卷积特征与全连接输入之间的形状变换。示例:
auto linear = torch::nn::Linear(torch::nn::LinearOptions(784, 256).bias(true)); auto output = linear->forward(input); // input: [N, 784]需要说明的是,各子文档(卷积、池化、线性、激活、归一化等)在仓库中进一步提供了按类展开的 Doxygen 成员明细,阅读某一具体模块的完整 API 时可直接查阅对应子页,例如 docs/cpp/source/api/nn/linear.md、docs/cpp/source/api/nn/convolution.md。
编译与工程实践要点
在真实工程中使用torch::nn通常还需要注意以下几点:
- 链接正确的库与头文件路径:本仓库为源码形态,编译需先按仓库 README 与 docs/cpp 流程构建 libtorch;工程中使用
torch/torch.h即可获得 nn/optim/data 全量 API; - 用
shared_ptr管理模块所有权:示例中auto model = std::make_shared<Net>();并非随意为之——modules(include_self=true)、clone()、register_module返回类型都基于shared_ptr;模块树按共享所有权组织更安全; - 句柄先置空再赋值:自定义模块成员句柄用
{nullptr}初始化,避免默认构造开销或未注册导致的参数丢失; - 别忘
register_*:任何希望被parameters()/to()/save()捕获的参数、缓冲区或子模块都必须显式注册,否则会“悄悄丢失”; - 推理前
eval()、训练前train():涉及Dropout、BatchNorm等行为随模式变化的模块时尤其重要; - Python 模型移植逐层对照:将 Python 侧
torch.nn模型迁移时,可逐层把nn.Conv2d(in, out, k, ...)改写为torch::nn::Conv2d(torch::nn::Conv2dOptions(in, out, k)...),结构、参数命名(点号分隔)与 state dict 语义一致,便于复用权重文件。
小结
torch::nn在 C++ 侧复刻了 Pythontorch.nn的模块化心智模型:ModuleHolder(句柄)+*Impl(实现)的 PIMPL 结构兼顾值语义与 ABI 稳定;Module基类统一承担参数/缓冲区/子模块的注册、递归遍历、设备与 dtype 转换、训练态切换、序列化与深拷贝;*Options结构体让卷积、线性、归一化等各色模块以链式配置的方式构造。基于本仓库 nn API 文档 以及 module.h、pimpl.h 的实现源码,开发者可以在 C++ 中构建与 Python 语义对齐、可移植可训练的自定义神经网络模型。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考