ihd*_*hdv 1 c++ python pytorch libtorch
我正在为 pytorch 编写一个 C++ 扩展,其中我需要通过索引访问张量的元素,并且还需要将元素转换为标准 C++ 类型。这是一个简短的例子。假设我有一个二维张量a,我需要访问a[i][j]它并将其转换为浮点数。
#include <torch/extension.h>\n\nfloat get(torch::Tensor a, int i, int j) {\n return a[i][j];\n}\nRun Code Online (Sandbox Code Playgroud)\n上面的内容被放入一个名为tensortest.cpp. 在另一个文件中setup.py我写
from setuptools import setup, Extension\nfrom torch.utils import cpp_extension\n\nsetup(name=\'tensortest\',\n ext_modules=[cpp_extension.CppExtension(\'tensortest_cpp\', [\'tensortest.cpp\'])],\n cmdclass={\'build_ext\': cpp_extension.BuildExtension})\nRun Code Online (Sandbox Code Playgroud)\n当我运行python setup.py install编译器时报告以下错误
running install\nrunning bdist_egg\nrunning egg_info\ncreating tensortest.egg-info\nwriting tensortest.egg-info/PKG-INFO\nwriting dependency_links to tensortest.egg-info/dependency_links.txt\nwriting top-level names to tensortest.egg-info/top_level.txt\nwriting manifest file \'tensortest.egg-info/SOURCES.txt\'\n/home/trisst/.local/lib/python3.8/site-packages/torch/utils/cpp_extension.py:335: UserWarning: Attempted to use ninja as the BuildExtension backend but we could not find ninja.. Falling back to using the slow distutils backend.\n warnings.warn(msg.format(\'we could not find ninja.\'))\nreading manifest file \'tensortest.egg-info/SOURCES.txt\'\nwriting manifest file \'tensortest.egg-info/SOURCES.txt\'\ninstalling library code to build/bdist.linux-x86_64/egg\nrunning install_lib\nrunning build_ext\nbuilding \'tensortest_cpp\' extension\ncreating build\ncreating build/temp.linux-x86_64-3.8\nx86_64-linux-gnu-gcc -pthread -Wno-unused-result -Wsign-compare -DNDEBUG -g -fwrapv -O2 -Wall -g -fstack-protector-strong -Wformat -Werror=format-security -g -fwrapv -O2 -g -fstack-protector-strong -Wformat -Werror=format-security -Wdate-time -D_FORTIFY_SOURCE=2 -fPIC -I/home/user/.local/lib/python3.8/site-packages/torch/include -I/home/user/.local/lib/python3.8/site-packages/torch/include/torch/csrc/api/include -I/home/user/.local/lib/python3.8/site-packages/torch/include/TH -I/home/user/.local/lib/python3.8/site-packages/torch/include/THC -I/usr/include/python3.8 -c tensortest.cpp -o build/temp.linux-x86_64-3.8/tensortest.o -DTORCH_API_INCLUDE_EXTENSION_H -DTORCH_EXTENSION_NAME=tensortest_cpp -D_GLIBCXX_USE_CXX11_ABI=0 -std=c++14\nIn file included from /home/user/.local/lib/python3.8/site-packages/torch/include/ATen/Parallel.h:149,\n from /home/user/.local/lib/python3.8/site-packages/torch/include/torch/csrc/api/include/torch/utils.h:3,\n from /home/user/.local/lib/python3.8/site-packages/torch/include/torch/csrc/api/include/torch/nn/cloneable.h:5,\n from /home/user/.local/lib/python3.8/site-packages/torch/include/torch/csrc/api/include/torch/nn.h:3,\n from /home/user/.local/lib/python3.8/site-packages/torch/include/torch/csrc/api/include/torch/all.h:7,\n from /home/user/.local/lib/python3.8/site-packages/torch/include/torch/extension.h:4,\n from tensortest.cpp:1:\n/home/user/.local/lib/python3.8/site-packages/torch/include/ATen/ParallelOpenMP.h:84: warning: ignoring #pragma omp parallel [-Wunknown-pragmas]\n 84 | #pragma omp parallel for if ((end - begin) >= grain_size)\n | \ntensortest.cpp: In function \xe2\x80\x98float get(at::Tensor, int, int)\xe2\x80\x99:\ntensortest.cpp:4:15: error: cannot convert \xe2\x80\x98at::Tensor\xe2\x80\x99 to \xe2\x80\x98float\xe2\x80\x99 in return\n 4 | return a[i][j];\n | ^\nerror: command \'x86_64-linux-gnu-gcc\' failed with exit status 1\nRun Code Online (Sandbox Code Playgroud)\n我能做些什么?
\n已编辑
#include <torch/extension.h>
float get(torch::Tensor a, int i, int j)
{
return a[i][j].item<float>();
}
Run Code Online (Sandbox Code Playgroud)
| 归档时间: |
|
| 查看次数: |
4168 次 |
| 最近记录: |