如何通过PyTorch C++ Extension加载外部静态或共享库?
我在PyTorch中使用自行编写的C++函数,通过torch.utils.cpp_extension.load加载JIT扩展,当前wrapper.py代码如下:
import os from torch.utils.cpp_extension import load dir_path = os.path.dirname(os.path.realpath(__file__)) my_func = load(name='my_func', sources=[os.path.join(dir_path, 'my_func.cpp')], extra_cflags=['-fopenmp', '-O2'], extra_ldflags=['-lgomp','-lrt'])
因为my_func.cpp用到OpenMP,所以加了上述编译链接参数。现在要在my_func.cpp中集成zstd库函数,已经克隆编译好zstd仓库,生成了libzstd.so共享库和libzstd.a静态库,代码里也加了#include <zstd.h>并调用相关函数。
命令行编译两种库都能成功:
g++ -fopenmp -O2 -lgomp -lrt -o my_func my_func.cpp lib/libzstd.so.1.5.3 g++ -fopenmp -O2 -lgomp -lrt -o my_func my_func.cpp lib/libzstd.a
但不知道怎么通过torch.utils.cpp_extension.load实现相同的编译链接,需要修改哪些参数?是否支持加载外部静态或共享库?
没问题,torch.utils.cpp_extension.load完全支持链接外部静态库或共享库,针对zstd的两种库文件,你可以按以下方式修改参数:
一、链接zstd共享库(libzstd.so)
需要补充两个关键参数:
- 告诉编译器zstd头文件的位置(如果头文件不在系统默认搜索路径),在
extra_cflags中添加-I<zstd_include_dir> - 告诉链接器zstd库文件的位置和库名,在
extra_ldflags中添加-L<zstd_lib_dir>和-lzstd
修改后的wrapper.py示例:
import os from torch.utils.cpp_extension import load dir_path = os.path.dirname(os.path.realpath(__file__)) # 替换成你的zstd实际安装路径 zstd_include_dir = "/path/to/zstd/include" zstd_lib_dir = "/path/to/zstd/lib" my_func = load( name='my_func', sources=[os.path.join(dir_path, 'my_func.cpp')], extra_cflags=['-fopenmp', '-O2', f'-I{zstd_include_dir}'], extra_ldflags=['-lgomp', '-lrt', f'-L{zstd_lib_dir}', '-lzstd'] )
二、链接zstd静态库(libzstd.a)
有两种简单的实现方式:
方式1:直接将静态库路径加入sources列表
静态库本质是目标文件的归档包,直接加到sources里,编译器会自动处理链接逻辑:
import os from torch.utils.cpp_extension import load dir_path = os.path.dirname(os.path.realpath(__file__)) # 替换成你的zstd静态库实际路径 zstd_static_lib = "/path/to/zstd/lib/libzstd.a" my_func = load( name='my_func', sources=[os.path.join(dir_path, 'my_func.cpp'), zstd_static_lib], extra_cflags=['-fopenmp', '-O2', '-I/path/to/zstd/include'], extra_ldflags=['-lgomp', '-lrt'] )
方式2:通过链接参数指定
和共享库逻辑类似,用-L指定库目录,-lzstd指定库名。如果目录下同时存在.so和.a,默认优先链接共享库,若要强制选静态库,可以用-Bstatic -lzstd -Bdynamic(链接完zstd后切回动态链接模式):
import os from torch.utils.cpp_extension import load dir_path = os.path.dirname(os.path.realpath(__file__)) my_func = load( name='my_func', sources=[os.path.join(dir_path, 'my_func.cpp')], extra_cflags=['-fopenmp', '-O2', '-I/path/to/zstd/include'], extra_ldflags=['-lgomp', '-lrt', '-L/path/to/zstd/lib', '-Bstatic', '-lzstd', '-Bdynamic'] )
注意事项
如果zstd的头文件已经安装到系统默认路径(比如/usr/include),可以去掉-I<zstd_include_dir>参数;如果库文件在系统默认路径(比如/usr/lib),可以去掉-L<zstd_lib_dir>参数,只保留-lzstd即可。
内容的提问来源于stack exchange,提问作者SHM

