diff --git a/.github/workflows/_build.yml b/.github/workflows/_build.yml index ca9aa246d09..bb433878a73 100644 --- a/.github/workflows/_build.yml +++ b/.github/workflows/_build.yml @@ -168,7 +168,7 @@ jobs: # Limit MAX_JOBS otherwise the github runner goes OOM # nvcc 11.8 can compile with 2 jobs, but nvcc 12.3 goes OOM - export MAX_JOBS=$([ "$MATRIX_CUDA_VERSION" == "129" ] || [ "$MATRIX_CUDA_VERSION" == "130" ] && echo 1 || echo 2) + export MAX_JOBS=$([ "$MATRIX_CUDA_VERSION" == "129" ] || [ "$MATRIX_CUDA_VERSION" == "130" ] || [ "$MATRIX_CUDA_VERSION" == "132" ] && echo 1 || echo 2) export NVCC_THREADS=2 export FLASH_ATTENTION_FORCE_BUILD="TRUE" export FLASH_ATTENTION_FORCE_CXX11_ABI=${{ inputs.cxx11_abi }} diff --git a/setup.py b/setup.py index 50f4b2fc79e..2fa1a0bbc5d 100644 --- a/setup.py +++ b/setup.py @@ -562,9 +562,14 @@ def get_wheel_url(): # We're using the CUDA version used to build torch, not the one currently installed # _, cuda_version_raw = get_cuda_bare_metal_version(CUDA_HOME) torch_cuda_version = parse(torch.version.cuda) - # For CUDA 11, we only compile for CUDA 11.8, and for CUDA 12 we only compile for CUDA 12.3 + # For CUDA 11 we compile for 11.8, for CUDA 12 for 12.3, and for CUDA 13 for 13.0 # to save CI time. Minor versions should be compatible. - torch_cuda_version = parse("11.8") if torch_cuda_version.major == 11 else parse("12.3") + if torch_cuda_version.major == 11: + torch_cuda_version = parse("11.8") + elif torch_cuda_version.major == 12: + torch_cuda_version = parse("12.3") + else: + torch_cuda_version = parse("13.0") # cuda_version = f"{cuda_version_raw.major}{cuda_version_raw.minor}" cuda_version = f"{torch_cuda_version.major}"