CoolFace
Apppublic

aikenml/data_mining

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
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