aikenml/data_mining
0
1import os2from setuptools import setup3from torch.utils.cpp_extension import BuildExtension, CUDAExtension, CppExtension4from os.path import join5 6CPU_ONLY = False7project_root = 'Correlation_Module'8 9source_files = ['correlation.cpp', 'correlation_sampler.cpp']10 11cxx_args = ['-std=c++17', '-fopenmp']12 13def generate_nvcc_args(gpu_archs):14 nvcc_args = []15 for arch in gpu_archs:16 nvcc_args.extend(['-gencode', f'arch=compute_{arch},code=sm_{arch}'])17 return nvcc_args18 19gpu_arch = os.environ.get('GPU_ARCH', '').split()20nvcc_args = generate_nvcc_args(gpu_arch)21 22with open("README.md", "r") as fh:23 long_description = fh.read()24 25 26def launch_setup():27 if CPU_ONLY:28 Extension = CppExtension29 macro = []30 else:31 Extension = CUDAExtension32 source_files.append('correlation_cuda_kernel.cu')33 macro = [("USE_CUDA", None)]34 35 sources = [join(project_root, file) for file in source_files]36 37 setup(38 name='spatial_correlation_sampler',39 version="0.4.0",40 author="Clément Pinard",41 author_email="clement.pinard@ensta-paristech.fr",42 description="Correlation module for pytorch",43 long_description=long_description,44 long_description_content_type="text/markdown",45 url="https://github.com/ClementPinard/Pytorch-Correlation-extension",46 install_requires=['torch>=1.1', 'numpy'],47 ext_modules=[48 Extension('spatial_correlation_sampler_backend',49 sources,50 define_macros=macro,51 extra_compile_args={'cxx': cxx_args, 'nvcc': nvcc_args},52 extra_link_args=['-lgomp'])53 ],54 package_dir={'': project_root},55 packages=['spatial_correlation_sampler'],56 cmdclass={57 'build_ext': BuildExtension58 },59 classifiers=[60 "Programming Language :: Python :: 3",61 "License :: OSI Approved :: MIT License",62 "Operating System :: POSIX :: Linux",63 "Intended Audience :: Science/Research",64 "Topic :: Scientific/Engineering :: Artificial Intelligence"65 ])66 67 68if __name__ == '__main__':69 launch_setup()70 