PyTorch源码中torch::col2im的定义位置与关联逻辑疑问
PyTorch中torch::col2im接口来源疑问
近期查阅PyTorch源码时,我发现在fold.h文件(对应路径torch/csrc/api/include/torch/nn/functional/fold.h,第18行)中使用了torch::col2im接口,对应代码片段如下:
#include <torch/nn/options/fold.h> namespace torch { namespace nn { namespace functional { ... if (input.dim() == 3 || input.dim() == 2) { return torch::col2im( input, output_size, kernel_size, dilation, padding, stride); } ...
但我在源码的torch命名空间下并未找到col2im的定义,仅在im2col.h文件(对应路径aten/src/ATen/native/im2col.h,第54行)中找到了at::native::col2im的实现,对应代码片段如下:
#include <ATen/ATen.h> #include <ATen/TensorUtils.h> #include <ATen/Utils.h> #include <c10/util/irange.h> #include <algorithm> namespace at { namespace native { ... template <typename T> static void col2im( const T* data_col, const int64_t channels, const int64_t height, const int64_t width, const int64_t output_height, const int64_t output_width, const int64_t kernel_h, const int64_t kernel_w, const int64_t pad_h, const int64_t pad_w, const int64_t stride_h, const int64_t stride_w, const int64_t dilation_h, const int64_t dilation_w, T* data_im) { ...
令我困惑的是,我在fold.h的#include引用中未找到二者的关联逻辑,但该源码显然可以正常编译成功。想请问torch::col2im是如何生成的?是否我遗漏了编译、链接环节的相关逻辑?
问题解答
你找到的at::native::col2im是底层数值计算的模板kernel,和上层调用的torch::col2im不是同一个实体,二者的关联来自PyTorch的编译期代码生成机制,不需要通过头文件显式#include做关联:
torch命名空间下的所有算子接口,都是编译阶段根据native_functions.yaml中的算子声明自动生成的。只要算子在YAML配置中完成注册,就会自动生成对应C++对外接口,不需要手动在源码树的头文件中写声明。- 你看到的
at::native::col2im是静态内部实现,不会直接暴露到外层命名空间。PyTorch会为该kernel生成统一的算子包装函数,注册到at命名空间;而torch命名空间通过全局using namespace at;的引入规则,会把所有ATen层的算子全部纳入自身命名空间,这是fold.h中可以直接调用torch::col2im的核心原因。 fold.h会间接引入ATen核心头文件,编译期生成的所有算子声明会在引入ATen核心头时被自动加载。你在原始源码树的torch命名空间下找不到col2im的手写定义,是因为这部分生成代码存放在编译输出的构建目录中,不会出现在拉取的原始源码目录里。
内容的提问来源于stack exchange,提问作者Zion
相关产品推荐
相关产品推荐

