Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 9 additions & 4 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -219,10 +219,8 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
# the set of architectures we want to compile for and remove the from the
# CMAKE_CUDA_FLAGS so that they are not applied globally.
#
# `+PTX` in TORCH_CUDA_ARCH_LIST is not preserved here. It is emitted by torch
# as `code=compute_*`, while extract_unique_cuda_archs_ascending() records only
# `arch=compute_*`. If a kernel really needs PTX, add `+PTX` to that kernel's
# component-specific arch list below.
# `+PTX` in TORCH_CUDA_ARCH_LIST is not preserved here. If a kernel really
# needs PTX, add `+PTX` to that kernel's component-specific arch list below.
#
clear_cuda_arches(CUDA_ARCH_FLAGS)
extract_unique_cuda_archs_ascending(CUDA_ARCHS "${CUDA_ARCH_FLAGS}")
Expand All @@ -232,6 +230,13 @@ if(VLLM_GPU_LANG STREQUAL "CUDA")
cuda_archs_loose_intersection(CUDA_ARCHS
"${CUDA_SUPPORTED_ARCHS}" "${CUDA_ARCHS}")
message(STATUS "CUDA supported target architectures: ${CUDA_ARCHS}")
if(NOT CUDA_ARCHS)
message(FATAL_ERROR
"No supported CUDA architectures; the build would produce a binary "
"with no usable kernels. Detected gencode flags: ${CUDA_ARCH_FLAGS}; "
"supported: ${CUDA_SUPPORTED_ARCHS}. "
"Set TORCH_CUDA_ARCH_LIST for your GPU (e.g. 12.0).")
endif()
else()
#
# For other GPU targets override the GPU architectures detected by cmake/torch
Expand Down
11 changes: 6 additions & 5 deletions cmake/utils.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -241,14 +241,15 @@ endmacro()
# `<major>.<minor>`, dedupes them and then sorts them in ascending order and
# stores them in `OUT_ARCHES`.
#
# Example:
# CUDA_ARCH_FLAGS="-gencode arch=compute_75,code=sm_75;...;-gencode arch=compute_90a,code=sm_90a"
# extract_unique_cuda_archs_ascending(OUT_ARCHES CUDA_ARCH_FLAGS)
# OUT_ARCHES="7.5;...;9.0"
# Prefer `code=sm_*`; fall back to `arch=compute_*` for PTX-only flags.
# This handles mismatches such as `arch=compute_20,code=sm_121`.
function(extract_unique_cuda_archs_ascending OUT_ARCHES CUDA_ARCH_FLAGS)
set(_CUDA_ARCHES)
foreach(_ARCH ${CUDA_ARCH_FLAGS})
string(REGEX MATCH "arch=compute_\([0-9]+[af]?\)" _COMPUTE ${_ARCH})
string(REGEX MATCH "code=sm_\([0-9]+[af]?\)" _COMPUTE ${_ARCH})
if (NOT _COMPUTE)
string(REGEX MATCH "arch=compute_\([0-9]+[af]?\)" _COMPUTE ${_ARCH})
endif()
if (_COMPUTE)
set(_COMPUTE ${CMAKE_MATCH_1})
endif()
Expand Down
25 changes: 25 additions & 0 deletions tests/test_cmake_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,3 +21,28 @@ def test_exact_family_arch_precedes_generic_family_fallback(tmp_path: Path):
)

subprocess.run(["cmake", "-P", script], check=True)


def test_extract_archs_prefers_sass_target_over_corrupted_virtual_arch(
tmp_path: Path,
):
"""torch's autodetection can emit a bogus arch=compute_* half (e.g.
capability 12.1 corrupted to arch=compute_20,code=sm_121); the SASS
target must win, while PTX-only entries keep the virtual arch."""
repo_root = Path(__file__).parents[1]
script = tmp_path / "test_extract_archs.cmake"
script.write_text(
f"""
cmake_minimum_required(VERSION 3.26)
include("{repo_root / "cmake" / "utils.cmake"}")
extract_unique_cuda_archs_ascending(actual
"-gencode arch=compute_20,code=sm_121;\
-gencode arch=compute_80,code=sm_80;\
-gencode arch=compute_80,code=compute_80")
if(NOT "${{actual}}" STREQUAL "8.0;12.1")
message(FATAL_ERROR "Expected '8.0;12.1', got '${{actual}}'")
endif()
"""
)

subprocess.run(["cmake", "-P", script], check=True)
Loading