求助:PyTorch指定卷积方法及提取MKL/cuDNN调用流程实现自定义算子
解决方案:PyTorch指定卷积实现及封装MKL/cuDNN自定义算子
一、在PyTorch中强制指定卷积实现方法
PyTorch的卷积算子会根据设备、输入维度、数据类型自动选择最优backend(如GPU的cuDNN、CPU的MKL-DNN,或自研native实现),要强制指定特定方法,可通过以下方式操作:
CPU环境(指定MKL-DNN或Native实现)
- 强制禁用MKL-DNN,使用PyTorch自研native实现:
import torch torch.backends.mkldnn.enabled = False # 执行卷积 input = torch.randn(1, 3, 224, 224) conv = torch.nn.Conv2d(3, 64, kernel_size=3) output = conv(input) - 强制启用MKL-DNN(默认已启用,若被禁用可重新开启):
torch.backends.mkldnn.enabled = True
GPU环境(指定cuDNN或Native实现)
- 强制禁用cuDNN,使用PyTorch自研CUDA实现:
import torch torch.backends.cudnn.enabled = False input = torch.randn(1, 3, 224, 224).cuda() conv = torch.nn.Conv2d(3, 64, kernel_size=3).cuda() output = conv(input) - 固定cuDNN算法(避免自动选择带来的性能波动):
注:torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = Falsebenchmark=True会预先测试所有可用cuDNN算法并选最优,deterministic=True会选择确定性算法,两者不可同时启用。
二、封装MKL/cuDNN卷积为自定义算子(消除对比不确定性)
要直接调用MKL/cuDNN的卷积实现,需绕过PyTorch的自动调度,提取底层库调用逻辑封装为独立自定义算子,以下分CPU(MKL-DNN)和GPU(cuDNN)场景说明:
核心逻辑定位
PyTorch对MKL/cuDNN的调用集中在底层代码:
- cuDNN调用路径:
torch/csrc/cuda/Conv.cpp(创建卷积描述符、选择算法、执行卷积) - MKL-DNN调用路径:
torch/csrc/aten/src/ATen/native/mkldnn/Conv.cpp(构建MKL-DNN primitive并执行)
步骤1:封装cuDNN卷积自定义算子(GPU)
- 编写CUDA扩展代码(示例框架):
#include <torch/extension.h> #include <cudnn.h> torch::Tensor cudnn_conv2d(torch::Tensor input, torch::Tensor weight, torch::Tensor bias, int stride, int padding) { // 初始化cuDNN handle cudnnHandle_t handle; cudnnCreate(&handle); // 创建tensor描述符 cudnnTensorDescriptor_t input_desc, weight_desc, output_desc; cudnnCreateTensorDescriptor(&input_desc); cudnnCreateTensorDescriptor(&weight_desc); cudnnCreateTensorDescriptor(&output_desc); // 设置描述符参数(NCHW格式、float类型) int n = input.size(0), c_in = input.size(1), h = input.size(2), w = input.size(3); int c_out = weight.size(0), k_h = weight.size(2), k_w = weight.size(3); int h_out = (h + 2*padding - k_h) / stride + 1; int w_out = (w + 2*padding - k_w) / stride + 1; cudnnSetTensor4dDescriptor(input_desc, CUDNN_TENSOR_NCHW, CUDNN_DATA_FLOAT, n, c_in, h, w); cudnnSetFilter4dDescriptor(weight_desc, CUDNN_DATA_FLOAT, CUDNN_TENSOR_NCHW, c_out, c_in, k_h, k_w); cudnnSetTensor4dDescriptor(output_desc, CUDNN_TENSOR_NCHW, CUDNN_DATA_FLOAT, n, c_out, h_out, w_out); // 创建卷积描述符 cudnnConvolutionDescriptor_t conv_desc; cudnnCreateConvolutionDescriptor(&conv_desc); cudnnSetConvolution2dDescriptor(conv_desc, padding, padding, stride, stride, 1, 1, CUDNN_CONVOLUTION, CUDNN_DATA_FLOAT); // 选择最优卷积算法 cudnnConvolutionFwdAlgo_t algo; cudnnFindConvolutionForwardAlgorithm(handle, input_desc, weight_desc, conv_desc, output_desc, 1, &algo); // 分配输出tensor auto output = torch::empty({n, c_out, h_out, w_out}, input.options()); // 执行卷积前向传播 const float alpha = 1.0f, beta = 0.0f; cudnnConvolutionForward(handle, &alpha, input_desc, input.data_ptr(), weight_desc, weight.data_ptr(), conv_desc, algo, nullptr, 0, &beta, output_desc, output.data_ptr()); // 释放资源 cudnnDestroyTensorDescriptor(input_desc); cudnnDestroyTensorDescriptor(weight_desc); cudnnDestroyTensorDescriptor(output_desc); cudnnDestroyConvolutionDescriptor(conv_desc); cudnnDestroy(handle); // 添加偏置 if (bias.defined()) { output += bias.view({1, c_out, 1, 1}); } return output; } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("cudnn_conv2d", &cudnn_conv2d, "cuDNN based 2D convolution"); } - 编译扩展:
from torch.utils.cpp_extension import load cudnn_conv = load(name='cudnn_conv', sources=['cudnn_conv.cpp'], verbose=True) - 调用自定义算子:
input = torch.randn(1, 3, 224, 224).cuda() weight = torch.randn(64, 3, 3, 3).cuda() bias = torch.randn(64).cuda() output = cudnn_conv.cudnn_conv2d(input, weight, bias, stride=1, padding=1)
步骤2:封装MKL-DNN卷积自定义算子(CPU)
核心逻辑类似GPU场景,调用MKL-DNN的C++ API:
- 编写C++扩展代码(核心逻辑):
#include <torch/extension.h> #include <mkldnn.hpp> using namespace mkldnn; torch::Tensor mkldnn_conv2d(torch::Tensor input, torch::Tensor weight, torch::Tensor bias, int stride, int padding) { // 初始化MKL-DNN engine和stream engine eng(engine::kind::cpu, 0); stream s(eng); // 转换PyTorch tensor为MKL-DNN memory对象 auto input_md = memory::desc({input.size(0), input.size(1), input.size(2), input.size(3)}, memory::data_type::f32, memory::format::nchw); auto input_mem = memory(input_md, eng, input.data_ptr()); auto weight_md = memory::desc({weight.size(0), weight.size(1), weight.size(2), weight.size(3)}, memory::data_type::f32, memory::format::oihw); auto weight_mem = memory(weight_md, eng, weight.data_ptr()); // 计算输出形状 int n = input.size(0), c_out = weight.size(0); int h_out = (input.size(2) + 2*padding - weight.size(2)) / stride + 1; int w_out = (input.size(3) + 2*padding - weight.size(3)) / stride + 1; auto output_md = memory::desc({n, c_out, h_out, w_out}, memory::data_type::f32, memory::format::nchw); auto output_mem = memory(output_md, eng); // 创建卷积primitive描述符 auto conv_desc = convolution_forward::desc(prop_kind::forward_inference, algorithm::convolution_direct, input_md, weight_md, output_md, {stride, stride}, {padding, padding}, {padding, padding}); auto conv_pd = convolution_forward::primitive_desc(conv_desc, eng); // 执行卷积 convolution_forward conv(conv_pd); conv.execute(s, {{MKLDNN_ARG_SRC, input_mem}, {MKLDNN_ARG_WEIGHT, weight_mem}, {MKLDNN_ARG_DST, output_mem}}); s.wait(); // 转换回PyTorch tensor auto output = torch::empty({n, c_out, h_out, w_out}, input.options()); output_mem.get_data_handle(output.data_ptr()); // 添加偏置 if (bias.defined()) { output += bias.view({1, c_out, 1, 1}); } return output; } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("mkldnn_conv2d", &mkldnn_conv2d, "MKL-DNN based 2D convolution"); } - 编译扩展(需系统已安装MKL-DNN):
from torch.utils.cpp_extension import load mkldnn_conv = load(name='mkldnn_conv', sources=['mkldnn_conv.cpp'], verbose=True) - 调用自定义算子:
input = torch.randn(1, 3, 224, 224) weight = torch.randn(64, 3, 3, 3) bias = torch.randn(64) output = mkldnn_conv.mkldnn_conv2d(input, weight, bias, stride=1, padding=1)
验证正确性
对比自定义算子与PyTorch原生卷积的输出,确保误差在浮点精度范围内:
import torch # 原生卷积结果 conv = torch.nn.Conv2d(3, 64, kernel_size=3, padding=1) native_output = conv(input) # 自定义算子结果 custom_output = cudnn_conv.cudnn_conv2d(input, conv.weight, conv.bias, stride=1, padding=1) # 误差检查 print(torch.allclose(native_output, custom_output, atol=1e-6)) # 应输出True
内容的提问来源于stack exchange,提问作者Sen Zhang
相关产品推荐
相关产品推荐

