yifeng99/nl2sql_14b
024
1from torch.utils import cpp_extension2import pathlib3import os4import subprocess5 6def _get_cuda_bare_metal_version(cuda_dir):7 raw_output = subprocess.check_output([cuda_dir + "/bin/nvcc", "-V"],8 universal_newlines=True)9 output = raw_output.split()10 release_idx = output.index("release") + 111 release = output[release_idx].split(".")12 bare_metal_major = release[0]13 bare_metal_minor = release[1][0]14 15 return raw_output, bare_metal_major, bare_metal_minor16 17def _create_build_dir(buildpath):18 try:19 os.mkdir(buildpath)20 except OSError:21 if not os.path.isdir(buildpath):22 print(f"Creation of the build directory {buildpath} failed")23 24# Check if cuda 11 is installed for compute capability 8.025cc_flag = []26_, bare_metal_major, bare_metal_minor = _get_cuda_bare_metal_version(cpp_extension.CUDA_HOME)27if int(bare_metal_major) >= 11:28 cc_flag.append('-gencode')29 cc_flag.append('arch=compute_80,code=sm_80')30 if int(bare_metal_minor) >= 7:31 cc_flag.append('-gencode')32 cc_flag.append('arch=compute_90,code=sm_90')33 34# Build path35srcpath = pathlib.Path(__file__).parent.absolute()36buildpath = srcpath / 'build'37_create_build_dir(buildpath)38 39def _cpp_extention_load_helper(name, sources, extra_cuda_flags):40 return cpp_extension.load(41 name=name,42 sources=sources,43 build_directory=buildpath,44 extra_cflags=['-O3', ],45 extra_cuda_cflags=['-O3',46 '-gencode', 'arch=compute_70,code=sm_70',47 '--use_fast_math'] + extra_cuda_flags + cc_flag,48 verbose=149 )50 51extra_flags = []52 53cache_autogptq_cuda_256_sources = ["./cache_autogptq_cuda_256.cpp",54 "./cache_autogptq_cuda_kernel_256.cu"]55cache_autogptq_cuda_256 = _cpp_extention_load_helper("cache_autogptq_cuda_256", cache_autogptq_cuda_256_sources, extra_flags)56 