set(CMAKE_CXX_STANDARD_REQUIRED ON)
set(CMAKE_CXX_EXTENSIONS ON)
set(CMAKE_EXPORT_COMPILE_COMMANDS ON)

# Check if running on Linux
if(NOT CMAKE_SYSTEM_NAME STREQUAL "Linux")
    message(FATAL_ERROR "This CPU extension only supports Linux. Current system: ${CMAKE_SYSTEM_NAME}")
endif()

# Define environment variables for special configuration
if(DEFINED ENV{VLLM_CPU_AVX512BF16})
    set(ENABLE_AVX512BF16 ON)
endif()

# include directories
include_directories("${CMAKE_SOURCE_DIR}/csrc")

# enable NUMA
set(ENABLE_NUMA TRUE)

# === 定义核心源文件 ===
set(CPU_CORE_SRC
    "csrc/cpu/cpu_inference_manager.cpp"
    "csrc/cpu/model_weight_base.cpp"
    "csrc/cpu/tensor_base.cpp"
    "csrc/cpu/tensor.cpp"
    "csrc/cpu/quantization_fp16.cpp"
    "csrc/cpu/quantization_q4_0.cpp"
    "csrc/cpu/quantization_q8_0.cpp"
    "csrc/cpu/cpu_utils.cpp"
    "csrc/cpu/matmul_fp16.cpp"
    "csrc/cpu/matmul_q8.cpp"
    "csrc/cpu/memory_manager.cpp"
)

# === 绑定层源文件 ===
set(CPU_BINDING_SRC
    "csrc/cpu/cpu_bindings.cpp"
    "csrc/cpu/cpu_inference.cpp"
)

# === CPU feature detection functions ===
function (find_isa CPUINFO TARGET OUT)
    string(FIND ${CPUINFO} ${TARGET} ISA_FOUND)
    if(NOT ISA_FOUND EQUAL -1)
        set(${OUT} ON PARENT_SCOPE)
    else()
        set(${OUT} OFF PARENT_SCOPE)
    endif()
endfunction()

# AVX512 extension (x86_64 only)
function (is_avx512_disabled OUT)
    set(DISABLE_AVX512 $ENV{VLLM_CPU_DISABLE_AVX512})
    if(DISABLE_AVX512 AND DISABLE_AVX512 STREQUAL "true")
        set(${OUT} ON PARENT_SCOPE)
    else()
        set(${OUT} OFF PARENT_SCOPE)
    endif()
endfunction()

# read CPU info from /proc/cpuinfo (Linux only)
execute_process(COMMAND cat /proc/cpuinfo
                RESULT_VARIABLE CPUINFO_RET
                OUTPUT_VARIABLE CPUINFO)
if (NOT CPUINFO_RET EQUAL 0)
    message(FATAL_ERROR "Failed to check CPU features via /proc/cpuinfo")
endif()

# === 初始化并构建完整的 CXX_COMPILE_FLAGS ===
list(APPEND CXX_COMPILE_FLAGS "-fopenmp" "-DENABLE_CPU_MP")

# detect CPU ISA features for all architectures
# x86_64 ISA features
find_isa(${CPUINFO} "avx2" AVX2_FOUND)
find_isa(${CPUINFO} "avx512f" AVX512_FOUND)
is_avx512_disabled(AVX512_DISABLED)

# ARM ISA features
find_isa(${CPUINFO} "asimd" ASIMD_FOUND) # Check for ARM NEON support
find_isa(${CPUINFO} "bf16" ARM_BF16_FOUND) # Check for ARM BF16 support
find_isa(${CPUINFO} "i8mm" ARM_I8MM_FOUND) # Check for ARM Int8 Matrix Multiply support
find_isa(${CPUINFO} "sve" ARM_SVE_FOUND)

# x86_64 architecture settings
if (CMAKE_SYSTEM_PROCESSOR MATCHES "x86_64")
    list(APPEND CXX_COMPILE_FLAGS "-mf16c")
    
    if (AVX512_FOUND AND NOT AVX512_DISABLED)
        list(APPEND CXX_COMPILE_FLAGS
            "-mavx512f"
            "-mavx512vl"
            "-mavx512bw"
            "-mavx512dq")
        
        find_isa(${CPUINFO} "avx512_bf16" AVX512BF16_FOUND)
        if (AVX512BF16_FOUND OR ENABLE_AVX512BF16)
            if (CMAKE_CXX_COMPILER_ID STREQUAL "GNU" AND CMAKE_CXX_COMPILER_VERSION VERSION_GREATER_EQUAL 12.3)
                list(APPEND CXX_COMPILE_FLAGS "-mavx512bf16")
            else()
                message(WARNING "Disable AVX512-BF16 ISA support, requires gcc/g++ >= 12.3")
            endif()
        else()
            message(WARNING "Disable AVX512-BF16 ISA support, no avx512_bf16 found in local CPU flags." " If cross-compilation is required, please set env VLLM_CPU_AVX512BF16=1.")
        endif()
    elseif (AVX2_FOUND)
        list(APPEND CXX_COMPILE_FLAGS "-mavx2")
        message(WARNING "vLLM CPU backend using AVX2 ISA")
    endif()
endif()

# ARM architecture settings
if (ASIMD_FOUND)
    message(STATUS "ARMv8 or later architecture detected")
    
    if(ARM_BF16_FOUND AND ARM_I8MM_FOUND)
        message(STATUS "BF16 and I8MM extension detected")
        set(MARCH_FLAGS "-march=armv8.2-a+bf16+dotprod+fp16+i8mm")
        add_compile_definitions(ARM_BF16_SUPPORT ARM_I8MM_SUPPORT __ARM_FEATURE_MATMUL_INT8)
    elseif(ARM_BF16_FOUND)
        message(STATUS "BF16 extension detected")
        set(MARCH_FLAGS "-march=armv8.2-a+bf16+dotprod+fp16")
        add_compile_definitions(ARM_BF16_SUPPORT)
    elseif(ARM_I8MM_FOUND)
        message(STATUS "I8MM extension detected")
        set(MARCH_FLAGS "-march=armv8.2-a+dotprod+fp16+i8mm")
        add_compile_definitions(ARM_I8MM_SUPPORT)
        add_compile_definitions(__ARM_FEATURE_MATMUL_INT8)
    elseif(ARM_SVE2_FOUND)
        message(STATUS "SVE2 extension detected (or forced)")
        set(MARCH_FLAGS "-march=armv9-a+sve2+fp16")
    elseif(ARM_SVE_FOUND)
        message(STATUS "SVE extension detected (or forced)")
        set(MARCH_FLAGS "-march=armv8.2-a+sve+fp16")
    else()
        message(WARNING "BF16 and I8MM functionality is not available")
        set(MARCH_FLAGS "-march=armv8.2-a+dotprod+fp16")
    endif()
    
    list(APPEND CXX_COMPILE_FLAGS ${MARCH_FLAGS}
        "-fpermissive"
        "-O3"
        "-funroll-loops"
        "-fomit-frame-pointer"
        "-ffast-math"
        "-finline-functions"
        "-fno-math-errno"
        "-flto"
        "-ftree-vectorize"
        "-funsafe-math-optimizations"
        "-falign-functions=16"
        "-falign-loops=16"
        "-fno-unwind-tables")
endif()

# Validate that at least one supported architecture was detected
if (NOT (AVX512_FOUND OR AVX2_FOUND OR ASIMD_FOUND))
    message(FATAL_ERROR "vLLM CPU backend requires AVX512, AVX2 or ARMv8 support.")
endif()

message(STATUS "CPU extension compile flags: ${CXX_COMPILE_FLAGS}")

# === NOW create cpu_core with full flags ===
add_library(cpu_core STATIC ${CPU_CORE_SRC})
target_include_directories(cpu_core PUBLIC "${CMAKE_SOURCE_DIR}/csrc" ${TORCH_INCLUDE_DIRS})
target_compile_options(cpu_core PRIVATE ${CXX_COMPILE_FLAGS})

# NUMA
if(ENABLE_NUMA)
    target_link_libraries(cpu_core PRIVATE numa)
    list(APPEND LIBS numa)
else()
    message(STATUS "NUMA is disabled")
    add_compile_definitions(-DVLLM_NUMA_DISABLED)
endif()

# === Extension target ===
define_gpu_extension_target(
    _cpu_inference
    DESTINATION sysHAX_adapter
    LANGUAGE CXX
    SOURCES ${CPU_BINDING_SRC}
    LIBRARIES ${LIBS} cpu_core
    COMPILE_FLAGS ${CXX_COMPILE_FLAGS}
    USE_SABI 3
    WITH_SOABI
)

message(STATUS "Enabling CPU LLM C extension.")