Text Generation
Transformers
Safetensors
English
Chinese
Russian
yue2
music-generation
orbitquant
quantization
4-bit precision
custom-code
8-bit precision
Instructions to use WaveCut/YuE2-3B-OrbitQuant-W4A4 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use WaveCut/YuE2-3B-OrbitQuant-W4A4 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="WaveCut/YuE2-3B-OrbitQuant-W4A4")# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("WaveCut/YuE2-3B-OrbitQuant-W4A4", device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use WaveCut/YuE2-3B-OrbitQuant-W4A4 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "WaveCut/YuE2-3B-OrbitQuant-W4A4" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "WaveCut/YuE2-3B-OrbitQuant-W4A4", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/WaveCut/YuE2-3B-OrbitQuant-W4A4
- SGLang
How to use WaveCut/YuE2-3B-OrbitQuant-W4A4 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "WaveCut/YuE2-3B-OrbitQuant-W4A4" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "WaveCut/YuE2-3B-OrbitQuant-W4A4", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "WaveCut/YuE2-3B-OrbitQuant-W4A4" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "WaveCut/YuE2-3B-OrbitQuant-W4A4", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use WaveCut/YuE2-3B-OrbitQuant-W4A4 with Docker Model Runner:
docker model run hf.co/WaveCut/YuE2-3B-OrbitQuant-W4A4
| # Vendored from vLLM: | |
| # | |
| # https://github.com/vllm-project/vllm/blob/main/cmake/utils.cmake | |
| # | |
| # Attempt to find the python package that uses the same python executable as | |
| # `EXECUTABLE` and is one of the `SUPPORTED_VERSIONS`. | |
| # | |
| macro (find_python_from_executable EXECUTABLE SUPPORTED_VERSIONS) | |
| file(REAL_PATH ${EXECUTABLE} EXECUTABLE) | |
| set(Python3_EXECUTABLE ${EXECUTABLE}) | |
| find_package(Python3 COMPONENTS Interpreter Development.Module Development.SABIModule) | |
| if (NOT Python3_FOUND) | |
| message(FATAL_ERROR "Unable to find python matching: ${EXECUTABLE}.") | |
| endif() | |
| set(_VER "${Python3_VERSION_MAJOR}.${Python3_VERSION_MINOR}") | |
| set(_SUPPORTED_VERSIONS_LIST ${SUPPORTED_VERSIONS} ${ARGN}) | |
| if (NOT _VER IN_LIST _SUPPORTED_VERSIONS_LIST) | |
| message(FATAL_ERROR | |
| "Python version (${_VER}) is not one of the supported versions: " | |
| "${_SUPPORTED_VERSIONS_LIST}.") | |
| endif() | |
| message(STATUS "Found python matching: ${EXECUTABLE}.") | |
| endmacro() | |
| # | |
| # Run `EXPR` in python. The standard output of python is stored in `OUT` and | |
| # has trailing whitespace stripped. If an error is encountered when running | |
| # python, a fatal message `ERR_MSG` is issued. | |
| # | |
| function (run_python OUT EXPR ERR_MSG) | |
| if(Python3_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python3_EXECUTABLE}") | |
| elseif(Python_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python_EXECUTABLE}") | |
| else() | |
| message(FATAL_ERROR "No Python executable found. Set Python3_EXECUTABLE or Python_EXECUTABLE.") | |
| endif() | |
| execute_process( | |
| COMMAND | |
| "${_PYTHON_EXECUTABLE}" "-c" "${EXPR}" | |
| OUTPUT_VARIABLE PYTHON_OUT | |
| RESULT_VARIABLE PYTHON_ERROR_CODE | |
| ERROR_VARIABLE PYTHON_STDERR | |
| OUTPUT_STRIP_TRAILING_WHITESPACE) | |
| if(NOT PYTHON_ERROR_CODE EQUAL 0) | |
| message(FATAL_ERROR "${ERR_MSG}: ${PYTHON_STDERR}") | |
| endif() | |
| set(${OUT} ${PYTHON_OUT} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # Run `SCRIPT_PATH` in Python. The standard output of Python is stored in | |
| # `OUT` and has trailing whitespace stripped. If the script exits with a | |
| # non-zero code, a fatal message `ERR_MSG` is issued. | |
| # | |
| function (run_python_script OUT SCRIPT_PATH ERR_MSG) | |
| if(Python3_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python3_EXECUTABLE}") | |
| elseif(Python_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python_EXECUTABLE}") | |
| else() | |
| message(FATAL_ERROR "No Python executable found. Set Python3_EXECUTABLE or Python_EXECUTABLE.") | |
| endif() | |
| execute_process( | |
| COMMAND | |
| "${_PYTHON_EXECUTABLE}" "${SCRIPT_PATH}" | |
| OUTPUT_VARIABLE PYTHON_OUT | |
| RESULT_VARIABLE PYTHON_ERROR_CODE | |
| ERROR_VARIABLE PYTHON_STDERR | |
| OUTPUT_STRIP_TRAILING_WHITESPACE) | |
| if(NOT PYTHON_ERROR_CODE EQUAL 0) | |
| message(FATAL_ERROR "${ERR_MSG}: ${PYTHON_STDERR}") | |
| endif() | |
| set(${OUT} ${PYTHON_OUT} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # Run `EXPR` in python. The standard output of python is stored in `OUT` and | |
| # has trailing whitespace stripped. If an error is encountered when running | |
| # python, `SUCCESS` is set to FALSE. If successful, `SUCCESS` is set to TRUE. | |
| # | |
| function (try_run_python OUT SUCCESS EXPR) | |
| if(Python3_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python3_EXECUTABLE}") | |
| elseif(Python_EXECUTABLE) | |
| set(_PYTHON_EXECUTABLE "${Python_EXECUTABLE}") | |
| else() | |
| message(FATAL_ERROR "No Python executable found. Set Python3_EXECUTABLE or Python_EXECUTABLE.") | |
| endif() | |
| execute_process( | |
| COMMAND | |
| "${_PYTHON_EXECUTABLE}" "-c" "${EXPR}" | |
| OUTPUT_VARIABLE PYTHON_OUT | |
| RESULT_VARIABLE PYTHON_ERROR_CODE | |
| ERROR_QUIET | |
| OUTPUT_STRIP_TRAILING_WHITESPACE) | |
| if(NOT PYTHON_ERROR_CODE EQUAL 0) | |
| set(${SUCCESS} FALSE PARENT_SCOPE) | |
| set(${OUT} "" PARENT_SCOPE) | |
| else() | |
| set(${SUCCESS} TRUE PARENT_SCOPE) | |
| set(${OUT} ${PYTHON_OUT} PARENT_SCOPE) | |
| endif() | |
| endfunction() | |
| # Run `EXPR` in python after importing `PKG`. Use the result of this to extend | |
| # `CMAKE_PREFIX_PATH` so the torch cmake configuration can be imported. | |
| macro (append_cmake_prefix_path PKG EXPR) | |
| run_python(_PREFIX_PATH | |
| "import ${PKG}; print(${EXPR})" "Failed to locate ${PKG} path") | |
| list(APPEND CMAKE_PREFIX_PATH ${_PREFIX_PATH}) | |
| endmacro() | |
| # | |
| # Add a target named `hipify${NAME}` that runs the hipify preprocessor on a set | |
| # of CUDA source files. The names of the corresponding "hipified" sources are | |
| # stored in `OUT_SRCS`. | |
| # | |
| function (hipify_sources_target OUT_SRCS NAME ORIG_SRCS) | |
| # | |
| # Split into C++ and non-C++ (i.e. CUDA) sources. | |
| # | |
| set(NODUP_SRCS ${ORIG_SRCS}) | |
| list(REMOVE_DUPLICATES NODUP_SRCS) | |
| set(SRCS ${NODUP_SRCS}) | |
| set(CXX_SRCS ${NODUP_SRCS}) | |
| list(FILTER SRCS INCLUDE REGEX "\.cu$") | |
| list(FILTER CXX_SRCS EXCLUDE REGEX "\.cu$") | |
| # | |
| # Generate ROCm/HIP source file names from CUDA file names. | |
| # Since HIP files are generated code, they will appear in the build area | |
| # `CMAKE_CURRENT_BINARY_DIR` directory rather than the original csrc dir. | |
| # | |
| set(HIP_SRCS) | |
| foreach (SRC ${SRCS}) | |
| get_source_file_property(include_dirs "${SRC}" INCLUDE_DIRECTORIES) | |
| get_source_file_property(compile_options "${SRC}" COMPILE_OPTIONS) | |
| string(REGEX REPLACE "\.cu$" "\.hip" SRC ${SRC}) | |
| string(REGEX REPLACE "cuda" "hip" SRC ${SRC}) | |
| if(include_dirs) | |
| # Copy over include directories from the original CUDA file. | |
| set_source_files_properties( | |
| ${SRC} | |
| PROPERTIES INCLUDE_DIRECTORIES "${include_dirs}") | |
| endif() | |
| if(compile_options) | |
| set_source_files_properties( | |
| ${SRC} | |
| PROPERTIES COMPILE_OPTIONS "${compile_options}") | |
| endif() | |
| list(APPEND HIP_SRCS "${CMAKE_CURRENT_BINARY_DIR}/${SRC}") | |
| endforeach() | |
| add_custom_target( | |
| hipify${NAME} | |
| COMMAND "${Python3_EXECUTABLE}" ${CMAKE_SOURCE_DIR}/cmake/hipify.py -p ${CMAKE_SOURCE_DIR} -o ${CMAKE_CURRENT_BINARY_DIR} ${SRCS} | |
| DEPENDS ${CMAKE_SOURCE_DIR}/cmake/hipify.py ${SRCS} | |
| BYPRODUCTS ${HIP_SRCS} | |
| COMMENT "Running hipify on ${NAME} extension source files.") | |
| # Swap out original extension sources with hipified sources. | |
| list(APPEND HIP_SRCS ${CXX_SRCS}) | |
| set(${OUT_SRCS} ${HIP_SRCS} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # Get additional GPU compiler flags from torch. | |
| # | |
| function (get_torch_gpu_compiler_flags OUT_GPU_FLAGS GPU_LANG) | |
| if (${GPU_LANG} STREQUAL "CUDA") | |
| # | |
| # Get common NVCC flags from torch. | |
| # | |
| run_python(GPU_FLAGS | |
| "from torch.utils.cpp_extension import COMMON_NVCC_FLAGS; print(';'.join(COMMON_NVCC_FLAGS))" | |
| "Failed to determine torch nvcc compiler flags") | |
| if (CUDA_VERSION VERSION_GREATER_EQUAL 11.8) | |
| list(APPEND GPU_FLAGS "-DENABLE_FP8") | |
| list(REMOVE_ITEM GPU_FLAGS | |
| "-D__CUDA_NO_HALF_OPERATORS__" | |
| "-D__CUDA_NO_HALF_CONVERSIONS__" | |
| "-D__CUDA_NO_BFLOAT16_CONVERSIONS__" | |
| "-D__CUDA_NO_HALF2_OPERATORS__") | |
| endif() | |
| elseif(${GPU_LANG} STREQUAL "HIP") | |
| # | |
| # Get common HIP/HIPCC flags from torch. | |
| # | |
| run_python(GPU_FLAGS | |
| "import torch.utils.cpp_extension as t; print(';'.join(t.COMMON_HIP_FLAGS + t.COMMON_HIPCC_FLAGS))" | |
| "Failed to determine torch nvcc compiler flags") | |
| list(APPEND GPU_FLAGS | |
| "-DUSE_ROCM" | |
| "-DENABLE_FP8" | |
| "-U__HIP_NO_HALF_CONVERSIONS__" | |
| "-U__HIP_NO_HALF_OPERATORS__" | |
| "-fno-gpu-rdc") | |
| endif() | |
| set(${OUT_GPU_FLAGS} ${GPU_FLAGS} PARENT_SCOPE) | |
| endfunction() | |
| # Macro for converting a `gencode` version number to a cmake version number. | |
| macro(string_to_ver OUT_VER IN_STR) | |
| string(REGEX REPLACE "\([0-9]+\)\([0-9]\)" "\\1.\\2" ${OUT_VER} ${IN_STR}) | |
| endmacro() | |
| # | |
| # Clear all `-gencode` flags from `CMAKE_CUDA_FLAGS`. | |
| # | |
| # Example: | |
| # CMAKE_CUDA_FLAGS="-Wall -gencode arch=compute_70,code=sm_70 -gencode arch=compute_75,code=sm_75" | |
| # clear_gencode_flags() | |
| # CMAKE_CUDA_FLAGS="-Wall" | |
| # | |
| macro(clear_gencode_flags) | |
| # Remove all `-gencode` flags from `CMAKE_CUDA_FLAGS` since they will be modified | |
| # and passed back via the `CUDA_ARCHITECTURES` property. | |
| string(REGEX REPLACE "-gencode arch=[^ ]+ *" "" CMAKE_CUDA_FLAGS | |
| ${CMAKE_CUDA_FLAGS}) | |
| endmacro() | |
| # | |
| # Extract unique CUDA architectures from a list of compute capabilities codes in | |
| # the form `<major><minor>[<letter>]`, convert them to the form sort | |
| # `<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" | |
| 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]+a?\)" _COMPUTE ${_ARCH}) | |
| if (_COMPUTE) | |
| set(_COMPUTE ${CMAKE_MATCH_1}) | |
| endif() | |
| string_to_ver(_COMPUTE_VER ${_COMPUTE}) | |
| list(APPEND _CUDA_ARCHES ${_COMPUTE_VER}) | |
| endforeach() | |
| list(REMOVE_DUPLICATES _CUDA_ARCHES) | |
| list(SORT _CUDA_ARCHES COMPARE NATURAL ORDER ASCENDING) | |
| set(${OUT_ARCHES} ${_CUDA_ARCHES} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # For a specific file set the `-gencode` flag in compile options conditionally | |
| # for the CUDA language. | |
| # | |
| # Example: | |
| # set_gencode_flag_for_srcs( | |
| # SRCS "foo.cu" | |
| # ARCH "compute_75" | |
| # CODE "sm_75") | |
| # adds: "-gencode arch=compute_75,code=sm_75" to the compile options for | |
| # `foo.cu` (only for the CUDA language). | |
| # | |
| macro(set_gencode_flag_for_srcs) | |
| set(options) | |
| set(oneValueArgs ARCH CODE) | |
| set(multiValueArgs SRCS) | |
| cmake_parse_arguments(arg "${options}" "${oneValueArgs}" | |
| "${multiValueArgs}" ${ARGN} ) | |
| set(_FLAG -gencode arch=${arg_ARCH},code=${arg_CODE}) | |
| set_property( | |
| SOURCE ${arg_SRCS} | |
| APPEND PROPERTY | |
| COMPILE_OPTIONS "$<$<COMPILE_LANGUAGE:CUDA>:${_FLAG}>" | |
| ) | |
| message(DEBUG "Setting gencode flag for ${arg_SRCS}: ${_FLAG}") | |
| endmacro(set_gencode_flag_for_srcs) | |
| # | |
| # For a list of source files set the `-gencode` flags in the files specific | |
| # compile options (specifically for the CUDA language). | |
| # | |
| # arguments are: | |
| # SRCS: list of source files | |
| # CUDA_ARCHS: list of CUDA architectures in the form `<major>.<minor>[letter]` | |
| # BUILD_PTX_FOR_ARCH: if set to true, then the PTX code will be built | |
| # for architecture `BUILD_PTX_FOR_ARCH` if there is a CUDA_ARCH in CUDA_ARCHS | |
| # that is larger than BUILD_PTX_FOR_ARCH. | |
| # | |
| macro(set_gencode_flags_for_srcs) | |
| set(options) | |
| set(oneValueArgs BUILD_PTX_FOR_ARCH) | |
| set(multiValueArgs SRCS CUDA_ARCHS) | |
| cmake_parse_arguments(arg "${options}" "${oneValueArgs}" | |
| "${multiValueArgs}" ${ARGN} ) | |
| foreach(_ARCH ${arg_CUDA_ARCHS}) | |
| # handle +PTX suffix: generate both sm and ptx codes if requested | |
| string(FIND "${_ARCH}" "+PTX" _HAS_PTX) | |
| if(NOT _HAS_PTX EQUAL -1) | |
| string(REPLACE "+PTX" "" _BASE_ARCH "${_ARCH}") | |
| string(REPLACE "." "" _STRIPPED_ARCH "${_BASE_ARCH}") | |
| set_gencode_flag_for_srcs( | |
| SRCS ${arg_SRCS} | |
| ARCH "compute_${_STRIPPED_ARCH}" | |
| CODE "sm_${_STRIPPED_ARCH}") | |
| set_gencode_flag_for_srcs( | |
| SRCS ${arg_SRCS} | |
| ARCH "compute_${_STRIPPED_ARCH}" | |
| CODE "compute_${_STRIPPED_ARCH}") | |
| else() | |
| string(REPLACE "." "" _STRIPPED_ARCH "${_ARCH}") | |
| set_gencode_flag_for_srcs( | |
| SRCS ${arg_SRCS} | |
| ARCH "compute_${_STRIPPED_ARCH}" | |
| CODE "sm_${_STRIPPED_ARCH}") | |
| endif() | |
| endforeach() | |
| if (${arg_BUILD_PTX_FOR_ARCH}) | |
| list(SORT arg_CUDA_ARCHS COMPARE NATURAL ORDER ASCENDING) | |
| list(GET arg_CUDA_ARCHS -1 _HIGHEST_ARCH) | |
| if (_HIGHEST_ARCH VERSION_GREATER_EQUAL ${arg_BUILD_PTX_FOR_ARCH}) | |
| string(REPLACE "." "" _PTX_ARCH "${arg_BUILD_PTX_FOR_ARCH}") | |
| set_gencode_flag_for_srcs( | |
| SRCS ${arg_SRCS} | |
| ARCH "compute_${_PTX_ARCH}" | |
| CODE "compute_${_PTX_ARCH}") | |
| endif() | |
| endif() | |
| endmacro() | |
| # | |
| # For the given `SRC_CUDA_ARCHS` list of gencode versions in the form | |
| # `<major>.<minor>[letter]` compute the "loose intersection" with the | |
| # `TGT_CUDA_ARCHS` list of gencodes. We also support the `+PTX` suffix in | |
| # `SRC_CUDA_ARCHS` which indicates that the PTX code should be built when there | |
| # is a CUDA_ARCH in `TGT_CUDA_ARCHS` that is equal to or larger than the | |
| # architecture in `SRC_CUDA_ARCHS`. | |
| # The loose intersection is defined as: | |
| # { max{ x \in tgt | x <= y } | y \in src, { x \in tgt | x <= y } != {} } | |
| # where `<=` is the version comparison operator. | |
| # In other words, for each version in `TGT_CUDA_ARCHS` find the highest version | |
| # in `SRC_CUDA_ARCHS` that is less or equal to the version in `TGT_CUDA_ARCHS`. | |
| # We have special handling for x.0a, if x.0a is in `SRC_CUDA_ARCHS` and x.0 is | |
| # in `TGT_CUDA_ARCHS` then we should remove x.0a from `SRC_CUDA_ARCHS` and add | |
| # x.0a to the result (and remove x.0 from TGT_CUDA_ARCHS). | |
| # The result is stored in `OUT_CUDA_ARCHS`. | |
| # | |
| # Example: | |
| # SRC_CUDA_ARCHS="7.5;8.0;8.6;9.0;9.0a" | |
| # TGT_CUDA_ARCHS="8.0;8.9;9.0" | |
| # cuda_archs_loose_intersection(OUT_CUDA_ARCHS SRC_CUDA_ARCHS TGT_CUDA_ARCHS) | |
| # OUT_CUDA_ARCHS="8.0;8.6;9.0;9.0a" | |
| # | |
| # Example With PTX: | |
| # SRC_CUDA_ARCHS="8.0+PTX" | |
| # TGT_CUDA_ARCHS="9.0" | |
| # cuda_archs_loose_intersection(OUT_CUDA_ARCHS SRC_CUDA_ARCHS TGT_CUDA_ARCHS) | |
| # OUT_CUDA_ARCHS="8.0+PTX" | |
| # | |
| function(cuda_archs_loose_intersection OUT_CUDA_ARCHS SRC_CUDA_ARCHS TGT_CUDA_ARCHS) | |
| set(_SRC_CUDA_ARCHS "${SRC_CUDA_ARCHS}") | |
| set(_TGT_CUDA_ARCHS ${TGT_CUDA_ARCHS}) | |
| # handle +PTX suffix: separate base arch for matching, record PTX requests | |
| set(_PTX_ARCHS) | |
| foreach(_arch ${_SRC_CUDA_ARCHS}) | |
| if(_arch MATCHES "\\+PTX$") | |
| string(REPLACE "+PTX" "" _base "${_arch}") | |
| list(APPEND _PTX_ARCHS "${_base}") | |
| list(REMOVE_ITEM _SRC_CUDA_ARCHS "${_arch}") | |
| list(APPEND _SRC_CUDA_ARCHS "${_base}") | |
| endif() | |
| endforeach() | |
| list(REMOVE_DUPLICATES _PTX_ARCHS) | |
| list(REMOVE_DUPLICATES _SRC_CUDA_ARCHS) | |
| # If x.0a or x.0f is in SRC_CUDA_ARCHS and x.0 is in CUDA_ARCHS then we should | |
| # remove x.0a or x.0f from SRC_CUDA_ARCHS and add x.0a or x.0f to _CUDA_ARCHS | |
| set(_CUDA_ARCHS) | |
| foreach(_arch ${_SRC_CUDA_ARCHS}) | |
| if(_arch MATCHES "[af]$") | |
| list(REMOVE_ITEM _SRC_CUDA_ARCHS "${_arch}") | |
| string(REGEX REPLACE "[af]$" "" _base "${_arch}") | |
| if ("${_base}" IN_LIST TGT_CUDA_ARCHS) | |
| list(REMOVE_ITEM _TGT_CUDA_ARCHS "${_base}") | |
| list(APPEND _CUDA_ARCHS "${_arch}") | |
| endif() | |
| endif() | |
| endforeach() | |
| list(SORT _SRC_CUDA_ARCHS COMPARE NATURAL ORDER ASCENDING) | |
| # for each ARCH in TGT_CUDA_ARCHS find the highest arch in SRC_CUDA_ARCHS that | |
| # is less or equal to ARCH (but has the same major version since SASS binary | |
| # compatibility is only forward compatible within the same major version). | |
| foreach(_ARCH ${_TGT_CUDA_ARCHS}) | |
| set(_TMP_ARCH) | |
| # Extract the major version of the target arch | |
| string(REGEX REPLACE "^([0-9]+)\\..*$" "\\1" TGT_ARCH_MAJOR "${_ARCH}") | |
| foreach(_SRC_ARCH ${_SRC_CUDA_ARCHS}) | |
| # Extract the major version of the source arch | |
| string(REGEX REPLACE "^([0-9]+)\\..*$" "\\1" SRC_ARCH_MAJOR "${_SRC_ARCH}") | |
| # Check version-less-or-equal, and allow PTX arches to match across majors | |
| if (_SRC_ARCH VERSION_LESS_EQUAL _ARCH) | |
| if (_SRC_ARCH IN_LIST _PTX_ARCHS OR SRC_ARCH_MAJOR STREQUAL TGT_ARCH_MAJOR) | |
| set(_TMP_ARCH "${_SRC_ARCH}") | |
| endif() | |
| else() | |
| # If we hit a version greater than the target, we can break | |
| break() | |
| endif() | |
| endforeach() | |
| # If we found a matching _TMP_ARCH, append it to _CUDA_ARCHS | |
| if (_TMP_ARCH) | |
| list(APPEND _CUDA_ARCHS "${_TMP_ARCH}") | |
| endif() | |
| endforeach() | |
| list(REMOVE_DUPLICATES _CUDA_ARCHS) | |
| # reapply +PTX suffix to architectures that requested PTX | |
| set(_FINAL_ARCHS) | |
| foreach(_arch ${_CUDA_ARCHS}) | |
| if(_arch IN_LIST _PTX_ARCHS) | |
| list(APPEND _FINAL_ARCHS "${_arch}+PTX") | |
| else() | |
| list(APPEND _FINAL_ARCHS "${_arch}") | |
| endif() | |
| endforeach() | |
| set(_CUDA_ARCHS ${_FINAL_ARCHS}) | |
| list(SORT _CUDA_ARCHS COMPARE NATURAL ORDER ASCENDING) | |
| set(${OUT_CUDA_ARCHS} ${_CUDA_ARCHS} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # For the given `SRC_ROCM_ARCHS` list of architecture versions in the form | |
| # `<name>` compute the "loose intersection" with the `TGT_ROCM_ARCHS` list. | |
| # The loose intersection is defined as: | |
| # { max{ x \in tgt | x <= y } | y \in src, { x \in tgt | x <= y } != {} } | |
| # where `<=` is the version comparison operator. | |
| # In other words, for each version in `TGT_ROCM_ARCHS` find the highest version | |
| # in `SRC_ROCM_ARCHS` that is less or equal to the version in `TGT_ROCM_ARCHS`. | |
| # The result is stored in `OUT_ROCM_ARCHS`. | |
| # | |
| # Example: | |
| # SRC_ROCM_ARCHS="gfx900;gfx906;gfx908;gfx90a" | |
| # TGT_ROCM_ARCHS="gfx906;gfx908;gfx1030" | |
| # hip_archs_loose_intersection(OUT_ROCM_ARCHS SRC_ROCM_ARCHS TGT_ROCM_ARCHS) | |
| # OUT_ROCM_ARCHS="gfx906;gfx908" | |
| # | |
| function(hip_archs_loose_intersection OUT_ROCM_ARCHS SRC_ROCM_ARCHS TGT_ROCM_ARCHS) | |
| list(REMOVE_DUPLICATES SRC_ROCM_ARCHS) | |
| # ROCm architectures are typically in format gfxNNN or gfxNNNx where N is a digit | |
| # and x is a letter. We can sort them by string comparison which works for this format. | |
| list(SORT SRC_ROCM_ARCHS COMPARE STRING ORDER ASCENDING) | |
| set(_ROCM_ARCHS) | |
| # Find the intersection of supported architectures | |
| foreach(_SRC_ARCH ${SRC_ROCM_ARCHS}) | |
| if(_SRC_ARCH IN_LIST TGT_ROCM_ARCHS) | |
| list(APPEND _ROCM_ARCHS ${_SRC_ARCH}) | |
| endif() | |
| endforeach() | |
| list(REMOVE_DUPLICATES _ROCM_ARCHS) | |
| set(${OUT_ROCM_ARCHS} ${_ROCM_ARCHS} PARENT_SCOPE) | |
| endfunction() | |
| function(cuda_remove_ptx_suffixes OUT_CUDA_ARCHS CUDA_ARCHS) | |
| set(_CUDA_ARCHS "${CUDA_ARCHS}") | |
| # handle +PTX suffix: separate base arch for matching, record PTX requests | |
| foreach(_arch ${CUDA_ARCHS}) | |
| if(_arch MATCHES "\\+PTX$") | |
| string(REPLACE "+PTX" "" _base "${_arch}") | |
| list(REMOVE_ITEM _CUDA_ARCHS "${_arch}") | |
| list(APPEND _CUDA_ARCHS "${_base}") | |
| endif() | |
| endforeach() | |
| list(REMOVE_DUPLICATES _CUDA_ARCHS) | |
| list(SORT _CUDA_ARCHS COMPARE NATURAL ORDER ASCENDING) | |
| set(${OUT_CUDA_ARCHS} ${_CUDA_ARCHS} PARENT_SCOPE) | |
| endfunction() | |
| # | |
| # Define a target named `GPU_MOD_NAME` for a single extension. The | |
| # arguments are: | |
| # | |
| # DESTINATION <dest> - Module destination directory. | |
| # LANGUAGE <lang> - The GPU language for this module, e.g CUDA, HIP, | |
| # etc. | |
| # SOURCES <sources> - List of source files relative to CMakeLists.txt | |
| # directory. | |
| # | |
| # Optional arguments: | |
| # | |
| # ARCHITECTURES <arches> - A list of target GPU architectures in cmake | |
| # format. | |
| # Refer `CMAKE_CUDA_ARCHITECTURES` documentation | |
| # and `CMAKE_HIP_ARCHITECTURES` for more info. | |
| # ARCHITECTURES will use cmake's defaults if | |
| # not provided. | |
| # COMPILE_FLAGS <flags> - Extra compiler flags passed to NVCC/hip. | |
| # INCLUDE_DIRECTORIES <dirs> - Extra include directories. | |
| # LIBRARIES <libraries> - Extra link libraries. | |
| # WITH_SOABI - Generate library with python SOABI suffix name. | |
| # USE_SABI <version> - Use python stable api <version> | |
| # | |
| # Note: optimization level/debug info is set via cmake build type. | |
| # | |
| function (define_gpu_extension_target GPU_MOD_NAME) | |
| cmake_parse_arguments(PARSE_ARGV 1 | |
| GPU | |
| "WITH_SOABI" | |
| "DESTINATION;LANGUAGE;USE_SABI" | |
| "SOURCES;COMPILE_FLAGS;INCLUDE_DIRECTORIES;LIBRARIES") | |
| # Add hipify preprocessing step when building with HIP/ROCm. | |
| if (GPU_LANGUAGE STREQUAL "HIP") | |
| hipify_sources_target(GPU_SOURCES ${GPU_MOD_NAME} "${GPU_SOURCES}") | |
| endif() | |
| if (GPU_WITH_SOABI) | |
| set(GPU_WITH_SOABI WITH_SOABI) | |
| else() | |
| set(GPU_WITH_SOABI) | |
| endif() | |
| if (GPU_USE_SABI) | |
| Python3_add_library(${GPU_MOD_NAME} MODULE USE_SABI ${GPU_USE_SABI} ${GPU_WITH_SOABI} "${GPU_SOURCES}") | |
| else() | |
| Python3_add_library(${GPU_MOD_NAME} MODULE ${GPU_WITH_SOABI} "${GPU_SOURCES}") | |
| endif() | |
| if (GPU_LANGUAGE STREQUAL "HIP") | |
| # Make this target dependent on the hipify preprocessor step. | |
| add_dependencies(${GPU_MOD_NAME} hipify${GPU_MOD_NAME}) | |
| # Clear target architectures, we are passing arch flags per source file. | |
| set_property(TARGET ${GPU_MOD_NAME} PROPERTY HIP_ARCHITECTURES off) | |
| endif() | |
| if (TORCH_VERSION VERSION_LESS 2.12.0) | |
| set_property(TARGET ${GPU_MOD_NAME} PROPERTY CXX_STANDARD 17) | |
| else() | |
| set_property(TARGET ${GPU_MOD_NAME} PROPERTY CXX_STANDARD 20) | |
| endif() | |
| target_compile_options(${GPU_MOD_NAME} PRIVATE | |
| $<$<COMPILE_LANGUAGE:${GPU_LANGUAGE}>:${GPU_COMPILE_FLAGS}>) | |
| target_compile_definitions(${GPU_MOD_NAME} PRIVATE | |
| "-DTORCH_EXTENSION_NAME=${GPU_MOD_NAME}") | |
| target_include_directories(${GPU_MOD_NAME} PRIVATE csrc | |
| ${GPU_INCLUDE_DIRECTORIES}) | |
| target_link_libraries(${GPU_MOD_NAME} PRIVATE torch ${GPU_LIBRARIES}) | |
| # Don't use `TORCH_LIBRARIES` for CUDA since it pulls in a bunch of | |
| # dependencies that are not necessary and may not be installed. | |
| if (GPU_LANGUAGE STREQUAL "CUDA") | |
| target_link_libraries(${GPU_MOD_NAME} PRIVATE CUDA::cudart) | |
| else() | |
| target_link_libraries(${GPU_MOD_NAME} PRIVATE ${TORCH_LIBRARIES}) | |
| endif() | |
| install(TARGETS ${GPU_MOD_NAME} LIBRARY DESTINATION ${GPU_DESTINATION} COMPONENT ${GPU_MOD_NAME}) | |
| endfunction() | |
| # Map a GPU language to its backend name. | |
| # | |
| # Arguments: | |
| # OUT_BACKEND - Output variable name for the backend string | |
| # GPU_LANG - The GPU language (CPU, CUDA, HIP, METAL, SYCL) | |
| # | |
| function(gpu_lang_to_backend OUT_BACKEND GPU_LANG) | |
| if (${GPU_LANG} STREQUAL "CPU") | |
| set(_BACKEND "cpu") | |
| elseif (${GPU_LANG} STREQUAL "CUDA") | |
| set(_BACKEND "cuda") | |
| elseif (${GPU_LANG} STREQUAL "HIP") | |
| set(_BACKEND "rocm") | |
| elseif (${GPU_LANG} STREQUAL "METAL") | |
| set(_BACKEND "metal") | |
| elseif (${GPU_LANG} STREQUAL "SYCL") | |
| set(_BACKEND "xpu") | |
| else() | |
| message(FATAL_ERROR "Unsupported GPU_LANG: ${GPU_LANG}") | |
| endif() | |
| set(${OUT_BACKEND} "${_BACKEND}" PARENT_SCOPE) | |
| endfunction() | |