Skip to content
Open
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
3 changes: 3 additions & 0 deletions .gitmodules
Original file line number Diff line number Diff line change
Expand Up @@ -10,3 +10,6 @@
[submodule "third_party/googletest"]
path = third_party/googletest
url = git@github.com:google/googletest.git
[submodule "third_party/flash-attention"]
path = third_party/flash-attention
url = git@github.com:Dao-AILab/flash-attention.git
76 changes: 76 additions & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ option(PROFILE_MODE "ENABLE PROFILE MODE" OFF)
option(USE_OMP "Use OpenMP as backend for Eigen" ON)
option(USE_NCCL "Build project for distributed running" ON)
option(BUILD_TEST "Build InfiniTrain tests" OFF)
option(USE_FLASH_ATTENTION "Enable FlashAttention 2 CUDA backend" ON)

project(infini_train VERSION 0.6.0 LANGUAGES CXX)

Expand Down Expand Up @@ -97,11 +98,70 @@ if(USE_CUDA)
find_package(CUDAToolkit REQUIRED)
include_directories(${CUDAToolkit_INCLUDE_DIRS})

if(USE_FLASH_ATTENTION)
set(FLASH_ATTN_SOURCE_DIR "${PROJECT_SOURCE_DIR}/third_party/flash-attention" CACHE PATH
"Path to vendored FlashAttention source tree")
if(NOT EXISTS "${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash.h")
message(FATAL_ERROR "FlashAttention submodule not found at ${FLASH_ATTN_SOURCE_DIR}. "
"Run: git submodule update --init --recursive third_party/flash-attention")
endif()
if(NOT EXISTS "${FLASH_ATTN_SOURCE_DIR}/csrc/cutlass/include/cutlass/cutlass.h")
message(FATAL_ERROR "FlashAttention CUTLASS dependency not found. "
"Run: git -C third_party/flash-attention submodule update --init csrc/cutlass")
endif()

# Minimal training subset: causal FP16/BF16 kernels for GPT-2/LLaMA head dims.
# TODO: add other head dims and architectures when InfiniTrain models require them.
set(FLASH_ATTN_CUDA_SOURCES
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_hdim64_fp16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_hdim64_bf16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_bwd_hdim64_fp16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_bwd_hdim64_bf16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_hdim128_fp16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_fwd_hdim128_bf16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_bwd_hdim128_fp16_causal_sm80.cu"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src/flash_bwd_hdim128_bf16_causal_sm80.cu"
)
add_library(flash_attn_native STATIC ${FLASH_ATTN_CUDA_SOURCES})
set_target_properties(flash_attn_native PROPERTIES
CUDA_STANDARD 17
CUDA_STANDARD_REQUIRED ON
CUDA_ARCHITECTURES "80"
POSITION_INDEPENDENT_CODE ON
)
target_compile_definitions(flash_attn_native PRIVATE
FLASHATTENTION_DISABLE_DROPOUT
FLASHATTENTION_DISABLE_ALIBI
FLASHATTENTION_DISABLE_SOFTCAP
FLASHATTENTION_DISABLE_LOCAL
)
target_include_directories(flash_attn_native PRIVATE
"${PROJECT_SOURCE_DIR}/infini_train/src/kernels/cuda/flash_attention_compat"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src"
"${FLASH_ATTN_SOURCE_DIR}/csrc/cutlass/include"
)
target_compile_options(flash_attn_native PRIVATE
$<$<COMPILE_LANGUAGE:CUDA>:-O3>
$<$<COMPILE_LANGUAGE:CUDA>:-U__CUDA_NO_HALF_OPERATORS__>
$<$<COMPILE_LANGUAGE:CUDA>:-U__CUDA_NO_HALF_CONVERSIONS__>
$<$<COMPILE_LANGUAGE:CUDA>:-U__CUDA_NO_HALF2_OPERATORS__>
$<$<COMPILE_LANGUAGE:CUDA>:-U__CUDA_NO_BFLOAT16_CONVERSIONS__>
$<$<COMPILE_LANGUAGE:CUDA>:--expt-relaxed-constexpr>
$<$<COMPILE_LANGUAGE:CUDA>:--expt-extended-lambda>
$<$<COMPILE_LANGUAGE:CUDA>:--use_fast_math>
)
target_link_libraries(flash_attn_native PUBLIC CUDA::cudart)
endif()

# CUDA compilation options
set(CMAKE_CUDA_FLAGS "${CMAKE_CUDA_FLAGS} --expt-extended-lambda --expt-relaxed-constexpr")

# Only compile CUDA kernels / cuda sources here (your original used src/*.cu)
file(GLOB_RECURSE CUDA_KERNELS ${PROJECT_SOURCE_DIR}/infini_train/src/*.cu)
if(NOT USE_FLASH_ATTENTION)
list(FILTER CUDA_KERNELS EXCLUDE REGEX ".*/kernels/cuda/flash_attention\\.cu$")
endif()

add_library(infini_train_cuda_kernels STATIC ${CUDA_KERNELS})
set_target_properties(infini_train_cuda_kernels PROPERTIES CUDA_ARCHITECTURES "75;80;90")
Expand All @@ -114,6 +174,21 @@ if(USE_CUDA)
CUDA::cuda_driver
)

if(USE_FLASH_ATTENTION)
target_compile_definitions(infini_train_cuda_kernels PRIVATE
FLASHATTENTION_DISABLE_DROPOUT
FLASHATTENTION_DISABLE_ALIBI
FLASHATTENTION_DISABLE_SOFTCAP
FLASHATTENTION_DISABLE_LOCAL
)
target_include_directories(infini_train_cuda_kernels PRIVATE
"${PROJECT_SOURCE_DIR}/infini_train/src/kernels/cuda/flash_attention_compat"
"${FLASH_ATTN_SOURCE_DIR}/csrc/flash_attn/src"
"${FLASH_ATTN_SOURCE_DIR}/csrc/cutlass/include"
)
target_link_libraries(infini_train_cuda_kernels PUBLIC flash_attn_native)
endif()

if(USE_NCCL)
message(STATUS "Add USE_NCCL, use NCCL with CUDA")
list(APPEND CMAKE_MODULE_PATH ${PROJECT_SOURCE_DIR}/cmake)
Expand Down Expand Up @@ -151,6 +226,7 @@ if(USE_CUDA)
# keep this. Otherwise it's harmless.
target_link_libraries(infini_train PUBLIC nccl)
endif()

endif()

# ------------------------------------------------------------------------------
Expand Down
Loading
Loading