modify the sgl-kernel to be compatible with transformers 5.x. (#14625)

This commit is contained in:
Yuhao Yang
2025-12-08 16:39:00 +08:00
committed by GitHub
parent aeff0d386b
commit f72a77038f
3 changed files with 59 additions and 7 deletions

View File

@@ -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

View File

@@ -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()

View File

@@ -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):