modify the sgl-kernel to be compatible with transformers 5.x. (#14625)
This commit is contained in:
@@ -106,6 +106,16 @@ $PIP_CMD install -e "python[${EXTRAS}]" --extra-index-url https://download.pytor
|
||||
# Install router for pd-disagg test
|
||||
$PIP_CMD install sglang-router $PIP_INSTALL_SUFFIX
|
||||
|
||||
PYTHON_LIB_PATH=$(python3 -c "import site; print(site.getsitepackages()[0])")
|
||||
FLASH_ATTN_PATH="${PYTHON_LIB_PATH}/flash_attn"
|
||||
|
||||
if [ -d "$FLASH_ATTN_PATH" ]; then
|
||||
echo "Directory $FLASH_ATTN_PATH exists. Removing..."
|
||||
rm -rf "$FLASH_ATTN_PATH"
|
||||
else
|
||||
echo "Directory $FLASH_ATTN_PATH does not exist."
|
||||
fi
|
||||
|
||||
# Install sgl-kernel
|
||||
SGL_KERNEL_VERSION_FROM_KERNEL=$(grep -Po '(?<=^version = ")[^"]*' sgl-kernel/pyproject.toml)
|
||||
SGL_KERNEL_VERSION_FROM_SRT=$(grep -Po -m1 '(?<=sgl-kernel==)[0-9A-Za-z\.\-]+' python/pyproject.toml)
|
||||
@@ -147,3 +157,14 @@ python3 -c "import torch; print(torch.version.cuda)"
|
||||
|
||||
# Prepare the CI runner (cleanup HuggingFace cache, etc.)
|
||||
bash "${SCRIPT_DIR}/prepare_runner.sh"
|
||||
|
||||
PYTHON_LIB_PATH=$(python3 -c "import site; print(site.getsitepackages()[0])")
|
||||
FLASH_ATTN_PATH="${PYTHON_LIB_PATH}/flash_attn"
|
||||
|
||||
if [ -d "$FLASH_ATTN_PATH" ]; then
|
||||
echo "Directory $FLASH_ATTN_PATH exists. Removing..."
|
||||
rm -rf "$FLASH_ATTN_PATH"
|
||||
echo "error: this should not happen"
|
||||
else
|
||||
echo "Directory $FLASH_ATTN_PATH does not exist."
|
||||
fi
|
||||
|
||||
@@ -593,9 +593,40 @@ install(DIRECTORY "${repo-triton_SOURCE_DIR}/python/triton_kernels/triton_kernel
|
||||
# ============================ Extra Install: FA4 ============================= #
|
||||
# TODO: find a better install condition.
|
||||
if ("${CUDA_VERSION}" VERSION_GREATER_EQUAL "12.8" OR SGL_KERNEL_ENABLE_SM100A)
|
||||
# flash_attn/cute
|
||||
install(DIRECTORY "${repo-flash-attention_SOURCE_DIR}/flash_attn/cute/"
|
||||
DESTINATION "flash_attn/cute"
|
||||
|
||||
set(FLASH_ATTN_CUTE_SRC "${repo-flash-attention_SOURCE_DIR}/flash_attn/cute")
|
||||
set(FLASH_ATTN_CUTE_DST "${CMAKE_CURRENT_BINARY_DIR}/flash_attn_origin/cute")
|
||||
|
||||
file(MAKE_DIRECTORY "${FLASH_ATTN_CUTE_DST}")
|
||||
|
||||
file(COPY "${FLASH_ATTN_CUTE_SRC}/"
|
||||
DESTINATION "${FLASH_ATTN_CUTE_DST}"
|
||||
PATTERN ".git*" EXCLUDE
|
||||
PATTERN "__pycache__" EXCLUDE)
|
||||
|
||||
file(GLOB_RECURSE FLASH_ATTN_CUTE_DST_PY
|
||||
"${FLASH_ATTN_CUTE_DST}/*.py")
|
||||
|
||||
foreach(FILE_PATH IN LISTS FLASH_ATTN_CUTE_DST_PY)
|
||||
file(READ "${FILE_PATH}" FILE_CONTENT)
|
||||
|
||||
set(MODIFIED_CONTENT "${FILE_CONTENT}")
|
||||
|
||||
# The main goal is to avoid using "flash_attn" so that other libraries (such as transformers) do not mistakenly assume that "flash_attn" is already installed.
|
||||
|
||||
string(REPLACE "flash_attn.cute"
|
||||
"flash_attn_origin.cute"
|
||||
MODIFIED_CONTENT "${MODIFIED_CONTENT}")
|
||||
|
||||
if (NOT FILE_CONTENT STREQUAL MODIFIED_CONTENT)
|
||||
file(WRITE "${FILE_PATH}" "${MODIFIED_CONTENT}")
|
||||
message(STATUS " - [FA4 Patch] Patched: ${FILE_PATH}")
|
||||
endif()
|
||||
endforeach()
|
||||
|
||||
install(DIRECTORY "${FLASH_ATTN_CUTE_DST}/"
|
||||
DESTINATION "flash_attn_origin/cute"
|
||||
PATTERN ".git*" EXCLUDE
|
||||
PATTERN "__pycache__" EXCLUDE)
|
||||
endif()
|
||||
|
||||
endif()
|
||||
|
||||
@@ -19,9 +19,9 @@ import cutlass
|
||||
import cutlass.cute as cute
|
||||
import torch
|
||||
from cutlass.cute.runtime import from_dlpack
|
||||
from flash_attn.cute import utils
|
||||
from flash_attn.cute.flash_fwd import FlashAttentionForwardSm90
|
||||
from flash_attn.cute.flash_fwd_sm100 import FlashAttentionForwardSm100
|
||||
from flash_attn_origin.cute import utils
|
||||
from flash_attn_origin.cute.flash_fwd import FlashAttentionForwardSm90
|
||||
from flash_attn_origin.cute.flash_fwd_sm100 import FlashAttentionForwardSm100
|
||||
|
||||
|
||||
def maybe_contiguous(x):
|
||||
|
||||
Reference in New Issue
Block a user