news 2026/9/8 17:00:06

PyTorch C++ 神经网络模块(torch::nn)API 指南:从 Module 基类到自定义模型构建

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch C++ 神经网络模块(torch::nn)API 指南:从 Module 基类到自定义模型构建

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);

该示例有四个值得注意的编码惯例:

  1. 成员声明为“默认空句柄”torch::nn::Conv2d conv1{nullptr};。因为Conv2dModuleHolder类型,可用nullptr先占位,待构造函数中真正创建后再赋值;
  2. 子模块必须register_moduleregister_module返回std::shared_ptr<Conv2d>,可直接赋值回成员(隐式转换为句柄),同时把子模块登记进父模块的children_表,保证后续parameters()to()save()能递归生效;
  3. 通过->调用成员:句柄重载了operator->,等价于操作内部Impl,因此写conv1->forward(x)而非conv1.forward(x)
  4. 输入输出为torch::Tensorforward的签名与 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::Conv2dModuleHolder值语义外壳,可复制、可空初始化
实现层Conv2dImpl*Impl子类继承torch::nn::Module,真正持有weight/bias参数并实现forward
基类层torch::nn::Module参数/缓冲区/子模块注册与递归遍历、totrain/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 源码注释):

  1. Parameters(参数):记录梯度、通常在反向阶段被更新的张量,例如Linearweight
  2. Buffers(缓冲区):不记录梯度、通常在前向阶段被更新的状态,例如BatchNorm中的running_meanrunning_var
  3. 其它任意状态:非张量、供实现或配置使用的额外数据。

注册 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时会把thisshared_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:查询当前模式。

BatchNormDropout是依赖该状态切换行为路径的典型模块(源码注释明确点名二者)。因此推理前调用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.hModule 基类声明与实现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::nntorch::optimtorch::data等常用模块。若追求更快的编译,可按上述细粒度头文件裁剪 include。

模块分类全景:13 大类能力总览

torch::nn提供的模块按功能分为多个子页,原文档在 docs/cpp/source/api/nn 目录中以 toctree 组织如下:

  • containers(容器模块):SequentialModuleListModuleDictParameterListParameterDict等组合与管理工具,对应源码 torch/nn/modules/container;
  • convolution(卷积层):Conv1d/2d/3dConvTranspose1d/2d/3d,源码见 conv.h;
  • pooling(池化层):各类最大/平均/自适应池化,源码见 pooling.h 与 adaptive.h;
  • linear(线性层):LinearBilinearIdentityFlattenUnflatten,源码见 linear.h;
  • activation(激活函数):ReLUGELU等可带参激活模块,源码见 activation.h;
  • normalization(归一化层):BatchNormLayerNormGroupNormInstanceNorm等,源码见 normalization.h、batchnorm.h 与 instancenorm.h;
  • dropout(随机失活):DropoutDropout2d/3dAlphaDropout等,源码见 dropout.h;
  • embedding(嵌入层):EmbeddingEmbeddingBag等,源码见 embedding.h;
  • recurrent(循环层):RNNLSTMGRU及其多层层级,源码见 rnn.h;
  • transformer(Transformer 组件):TransformerTransformerEncoder/DecoderTransformerEncoderLayer/DecoderLayerMultiheadAttention等,源码见 transformer.h、transformerlayer.h 与 transformercoder.h;
  • loss(损失函数):L1LossMSELossCrossEntropyLossBCELossNLLLoss等,源码见 loss.h;
  • functional(函数式接口):无状态地调用算子(如F::relu),不维护参数,源码见 torch/nn/functional;
  • utilities(工具模块):ReflectionPadZeroPadPixelShuffleUpsample等实用层,对应 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 + bBilinear对两个输入做双线性变换,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通常还需要注意以下几点:

  1. 链接正确的库与头文件路径:本仓库为源码形态,编译需先按仓库 README 与 docs/cpp 流程构建 libtorch;工程中使用torch/torch.h即可获得 nn/optim/data 全量 API;
  2. shared_ptr管理模块所有权:示例中auto model = std::make_shared<Net>();并非随意为之——modules(include_self=true)clone()register_module返回类型都基于shared_ptr;模块树按共享所有权组织更安全;
  3. 句柄先置空再赋值:自定义模块成员句柄用{nullptr}初始化,避免默认构造开销或未注册导致的参数丢失;
  4. 别忘register_*:任何希望被parameters()/to()/save()捕获的参数、缓冲区或子模块都必须显式注册,否则会“悄悄丢失”;
  5. 推理前eval()、训练前train():涉及DropoutBatchNorm等行为随模式变化的模块时尤其重要;
  6. 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),仅供参考

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

three js 13 光照和阴影

文章目录1 灯光的类型2 材质3 如何场景有影子3 平行光4 聚光灯5 点光源1 灯光的类型 平行光 点光 面光 无阴影 射灯 2 材质 以下材质会接受光照 MeshStandardMaterial 标准PBR材质&#xff0c;主要用这个 MeshPhysicalMaterial 高级物理材质 没用 MeshLambertMaterial 兰伯…

作者头像 李华
网站建设 2026/9/8 16:54:57

冻融循环与氯离子侵蚀耦合下的混凝土耐久性数值模拟

这几年在北方沿海、盐渍土地区跑项目&#xff0c;混凝土耐久性病害里最让人头疼的组合就是冻融循环和氯离子侵蚀同时出现。单独做冻融试验或单测氯离子扩散&#xff0c;结果往往偏乐观&#xff0c;现场却早早出现顺筋裂缝和表层剥落。原因在于这两个过程根本就不是简单叠加&…

作者头像 李华
网站建设 2026/9/8 16:53:35

Windows消息机制详解:从硬件事件到窗口过程的完整链路

刚把Windows系统的启动流程和进程调度捋清楚没多久&#xff0c;我又一头扎进了消息机制。这个知识点我老早就想整理成笔记&#xff0c;但一直觉得它既抽象又琐碎&#xff1a;网上能找到的资料要么停留在“给你一段WinMain抄一下”&#xff0c;要么就直接上MFC/消息循环源码&…

作者头像 李华
网站建设 2026/9/8 16:52:33

传感器实战解析:从原理到选型与信号处理

传感器这东西&#xff0c;干我们这行的天天跟它打交道&#xff0c;但真要说“懂”它&#xff0c;很多人其实是懵的。你问一个刚入行的工程师“传感器是啥”&#xff0c;他能给你背出“将非电量转换为电量的器件”这种教科书答案&#xff1b;但你问他“为什么你的称重数据老是漂…

作者头像 李华