自定义实现C++拓展pytorch功能

ncrelu.cpp

cpp 复制代码
#include <torch/extension.h>					// 头文件引用部分

namespace py = pybind11;

torch::Tensor ncrelu_forward(torch::Tensor input) {
    auto pos = input.clamp_min(0);				       // 具体实现部分
    auto neg = input.clamp_max(0);
    return torch::cat({pos, neg}, 1);
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {	// 绑定部分
    m.def("forward", &ncrelu_forward, py::arg("input"), "NCReLU forward");
}

setup.py

python 复制代码
from setuptools import setup
from torch.utils import cpp_extension


setup(
    name='ncrelu_cpp',
    version='1.0',# 编译后的链接库名称
    py_modules=['ncrelu_cpp'],
    ext_modules=[
        cpp_extension.CppExtension(
            'ncrelu_cpp', ['ncrelu.cpp'],
            extra_compile_args={'cxx': ['-O2']}
            # 待编译文件,及编译函数
        )
    ],
    cmdclass={						       # 执行编译命令设置
        'build_ext': cpp_extension.BuildExtension
    }
)

test.py

python 复制代码
import torch
import ncrelu_cpp
import sys
print(sys.path)
a = torch.randn(4,3)
print(a)
b = ncrelu_cpp.forward(a)

python setup.py install

或pip install .

但是在Windows平台下不知道为什么会报错找不到包,或者找不到函数,很奇怪,但是正常运行没有任何问题

相关推荐
东莞市云毅网络有限公司6 小时前
AI 引用句逐条回指原文:让回答可回溯的校验实现
python·数据清洗·rag·企业知识库·文档解析
数字融合7 小时前
透明化视频三维矿山井下照明重建技术
人工智能·python·数码相机
fpcc8 小时前
c++编程实践—堆和栈越界调试
c++
yi0118 小时前
LeetCode 219:存在重复元素 II——哈希表记录“最近一次出现的位置”
数据结构·人工智能·笔记·python·算法·leetcode·哈希表
Marst Code9 小时前
上位机开发日记 · 第 2 篇 · 架构先行:六层分层与边界
python
weiabc9 小时前
MSYS2 UCRT64 + g++ C++ 中文输出乱码 完整全过程
c++
李日华大战鸡红10 小时前
FOC状态空间方程模型推导(学习记录)
python·学习·线性代数
灯澜忆梦10 小时前
【面向对象编程C++】| 基础语法
java·c++·算法
aqiu11111111 小时前
【C++算法打怪专栏】LeetCode 165. 比较版本号(双指针与字符串流法)
c++·算法·leetcode
外收内放12 小时前
Python基础语法练习题(57-58)
开发语言·python