cmake_minimum_required(VERSION 3.11)
project(qv100 LANGUAGES CXX)

set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)

# ============================================================
# 平台检测
# ============================================================
if(MSVC)
    set(WUYUE_MSVC ON)
    message(STATUS "Compiler: MSVC (Windows)")
else()
    set(WUYUE_GCC_CLANG ON)
    message(STATUS "Compiler: GCC/Clang (Linux/macOS)")
endif()

# ============================================================
# 全局编译选项
# ============================================================
if(CMAKE_BUILD_TYPE STREQUAL "Debug")
    # Debug: 带调试信息，无优化
    if(WUYUE_MSVC)
        set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} /Od /Zi /MDd")
    else()
        set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -g -O0 -fno-omit-frame-pointer")
    endif()
else()
    # Release/其他: 优化
    if(WUYUE_MSVC)
        # MSVC: /O2 全局优化, /fp:fast 快速浮点, /MD 多线程DLL
        set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} /O2 /fp:fast /MD /DNDEBUG /DEIGEN_NO_DEBUG")
    else()
        set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -O3 -march=native -ffast-math")
        set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -DNDEBUG -DEIGEN_NO_DEBUG")
    endif()
endif()

# 确保在未指定时，默认使用 Release 模式
if(NOT CMAKE_BUILD_TYPE AND NOT CMAKE_CONFIGURATION_TYPES)
    if(WUYUE_MSVC)
        # MSVC 多配置生成器不设置 CMAKE_BUILD_TYPE
    else()
        set(CMAKE_BUILD_TYPE "RelWithDebInfo" CACHE STRING "Choose the type of build." FORCE)
    endif()
endif()

# ============================================================
# 依赖项检测
# ============================================================
# 自动检测 Python torch 包路径，实现跨环境兼容
if(DEFINED Python3_EXECUTABLE)
    set(PYTHON_EXE "${Python3_EXECUTABLE}")
elseif(DEFINED PYTHON_EXECUTABLE)
    set(PYTHON_EXE "${PYTHON_EXECUTABLE}")
else()
    find_program(PYTHON_EXE NAMES python3 python)
    if(NOT PYTHON_EXE)
        message(FATAL_ERROR "Python executable not found. Please set Python3_EXECUTABLE or PYTHON_EXECUTABLE.")
    endif()
endif()
message(STATUS "Using Python: ${PYTHON_EXE}")

# 获取 torch 包的 base 目录
execute_process(
    COMMAND "${PYTHON_EXE}" -c "import torch; import os; print(os.path.dirname(torch.__file__))"
    OUTPUT_VARIABLE TORCH_BASE_DIR
    OUTPUT_STRIP_TRAILING_WHITESPACE
    ERROR_QUIET
    RESULT_VARIABLE TORCH_IMPORT_RESULT
)
if(NOT TORCH_IMPORT_RESULT EQUAL 0 OR NOT TORCH_BASE_DIR)
    message(FATAL_ERROR "Cannot import torch via ${PYTHON_EXE}. Please ensure torch is installed.")
endif()
message(STATUS "Torch base dir: ${TORCH_BASE_DIR}")

# 设置 Torch_DIR
set(Torch_DIR "${TORCH_BASE_DIR}/share/cmake/Torch")
if(NOT EXISTS "${Torch_DIR}/TorchConfig.cmake")
    message(FATAL_ERROR "TorchConfig.cmake not found at ${Torch_DIR}. Please check torch installation.")
endif()

find_package(Torch REQUIRED)
set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} ${TORCH_CXX_FLAGS}")
find_package(Eigen3 REQUIRED)

# 从 pip 安装的 pybind11 查找 cmake 配置
execute_process(
    COMMAND "${PYTHON_EXE}" -c "import pybind11; import os; print(os.path.join(os.path.dirname(pybind11.__file__), 'share', 'cmake', 'pybind11'))"
    OUTPUT_VARIABLE PYBIND11_DIR
    OUTPUT_STRIP_TRAILING_WHITESPACE
    ERROR_QUIET
    RESULT_VARIABLE PYBIND11_RESULT
)
if(PYBIND11_RESULT EQUAL 0 AND EXISTS "${PYBIND11_DIR}/pybind11Config.cmake")
    message(STATUS "Found pybind11 at: ${PYBIND11_DIR}")
    set(pybind11_DIR "${PYBIND11_DIR}")
    find_package(pybind11 REQUIRED)
else()
    message(FATAL_ERROR "Cannot find pybind11 cmake config. Please install: pip install pybind11")
endif()

find_package(OpenMP) # 尝试查找 OpenMP

# ============================================================
# 定义扩展模块
# ============================================================
pybind11_add_module(types src/types.cpp src/types.h)
pybind11_add_module(qv100 src/qv100.cpp src/types.h)
pybind11_add_module(n_local_accelerate src/n_local_accelerate.cpp src/types.h)
pybind11_add_module(full_amplitude_accelerate src/full_amplitude_accelerate.cpp src/types.h src/gates.h src/gates.cpp)
pybind11_add_module(gradient_simulator src/gradient_simulator.cpp src/types.h src/gates.h src/gates.cpp)

# ============================================================
# RPATH 配置（仅 Unix 需要；Windows 通过 PATH 或同目录查找 DLL）
# ============================================================
if(NOT WIN32)
    foreach(mod types qv100 n_local_accelerate full_amplitude_accelerate gradient_simulator)
        set_target_properties(${mod} PROPERTIES
            BUILD_RPATH "${TORCH_BASE_DIR}/lib"
            INSTALL_RPATH "${TORCH_BASE_DIR}/lib"
        )
    endforeach()
endif()

# ============================================================
# 查找 libtorch_python 库（跨平台：.so / .dll / .dylib）
# 提供 type_caster<at::Tensor> 实现
# ============================================================
if(WIN32)
    set(_TORCH_PYTHON_CANDIDATES
        "${TORCH_BASE_DIR}/lib/torch_python.dll"
        "${TORCH_BASE_DIR}/lib/torch_python.lib"
    )
else()
    set(_TORCH_PYTHON_CANDIDATES
        "${TORCH_BASE_DIR}/lib/libtorch_python.so"
        "${TORCH_BASE_DIR}/lib/libtorch_python.dylib"
    )
endif()

set(TORCH_PYTHON_LIB "")
foreach(_candidate ${_TORCH_PYTHON_CANDIDATES})
    if(EXISTS "${_candidate}" AND NOT TORCH_PYTHON_LIB)
        set(TORCH_PYTHON_LIB "${_candidate}")
    endif()
endforeach()

# Windows 上 find_library 默认查找 .lib 导入库；这里也尝试一下
if(NOT TORCH_PYTHON_LIB)
    find_library(TORCH_PYTHON_LIB
        NAMES torch_python
        PATHS "${TORCH_INSTALL_PREFIX}/lib" "${TORCH_BASE_DIR}/lib" "${TORCH_BASE_DIR}/../lib"
        NO_DEFAULT_PATH
    )
endif()

if(NOT TORCH_PYTHON_LIB)
    message(WARNING "Cannot find torch_python library. Modules may not work with torch::Tensor parameters.")
else()
    message(STATUS "Found torch_python: ${TORCH_PYTHON_LIB}")
endif()

# Windows 上链接 .dll 需要 .lib 导入库；如果只找到 .dll，尝试找对应的 .lib
if(WIN32 AND TORCH_PYTHON_LIB AND TORCH_PYTHON_LIB MATCHES "\\.dll$")
    get_filename_component(_tp_dir "${TORCH_PYTHON_LIB}" DIRECTORY)
    if(EXISTS "${_tp_dir}/torch_python.lib")
        set(TORCH_PYTHON_LINK_LIB "${_tp_dir}/torch_python.lib")
    else()
        # 没找到 .lib，直接用 .dll（MSVC 也能链接，但不推荐）
        set(TORCH_PYTHON_LINK_LIB "${TORCH_PYTHON_LIB}")
    endif()
else()
    set(TORCH_PYTHON_LINK_LIB "${TORCH_PYTHON_LIB}")
endif()

# 所有模块统一链接 torch 库
foreach(mod types qv100 n_local_accelerate full_amplitude_accelerate gradient_simulator)
    target_link_libraries(${mod} PRIVATE
        Eigen3::Eigen
        ${TORCH_LIBRARIES}
    )
    if(TORCH_PYTHON_LINK_LIB)
        target_link_libraries(${mod} PRIVATE ${TORCH_PYTHON_LINK_LIB})
    endif()
endforeach()

# ============================================================
# OpenMP 配置
# ============================================================
if(OpenMP_CXX_FOUND)
    message(STATUS "OpenMP found via find_package. Linking using INTERFACE target.")
    target_link_libraries(qv100 PRIVATE OpenMP::OpenMP_CXX)
    target_link_libraries(full_amplitude_accelerate PRIVATE OpenMP::OpenMP_CXX)
    target_link_libraries(gradient_simulator PRIVATE OpenMP::OpenMP_CXX)
elseif(WUYUE_GCC_CLANG)
    message(STATUS "OpenMP not found via find_package. Manually setting -fopenmp flags.")
    target_compile_options(qv100 PRIVATE "-fopenmp")
    target_link_options(qv100 PRIVATE "-fopenmp")
    target_compile_options(full_amplitude_accelerate PRIVATE "-fopenmp")
    target_link_options(full_amplitude_accelerate PRIVATE "-fopenmp")
    target_compile_options(gradient_simulator PRIVATE "-fopenmp")
    target_link_options(gradient_simulator PRIVATE "-fopenmp")
else()
    # MSVC 上 OpenMP 未找到
    message(WARNING "OpenMP not found. Modules will be built without OpenMP support (single-threaded).")
endif()

# ============================================================
# 模块级极致优化编译选项
# ============================================================
if(WUYUE_MSVC)
    # MSVC 优化选项
    # /O2 最大化速度, /fp:fast 快速浮点, /arch:AVX2 启用 AVX2 指令集,
    # /openmp 已通过 OpenMP::OpenMP_CXX 处理（若找到）, /Ob2 内联展开,
    # /GS- 关闭缓冲区安全检查（可选, 提速）, /EHa- 关闭异常处理加速
    # /wd4267: 抑制 size_t→int 转换警告（torch/STL 内部头文件大量触发）
    # /wd4819: 抑制源文件编码警告
    set(_MSVC_OPTS /O2 /fp:fast /arch:AVX2 /Ob2 /DNDEBUG /DEIGEN_NO_DEBUG /wd4267 /wd4819)

    # 若 OpenMP 未通过 find_package 找到，手动加 /openmp
    if(NOT OpenMP_CXX_FOUND)
        list(APPEND _MSVC_OPTS /openmp)
    endif()

    target_compile_options(qv100 PRIVATE ${_MSVC_OPTS})
    target_compile_options(full_amplitude_accelerate PRIVATE ${_MSVC_OPTS})
    target_compile_options(gradient_simulator PRIVATE ${_MSVC_OPTS})

    # n_local_accelerate 和 types 不需要 OpenMP，只加基础优化
    target_compile_options(types PRIVATE /O2 /fp:fast /arch:AVX2 /Ob2 /DNDEBUG /DEIGEN_NO_DEBUG /wd4267 /wd4819)
    target_compile_options(n_local_accelerate PRIVATE /O2 /fp:fast /arch:AVX2 /Ob2 /DNDEBUG /DEIGEN_NO_DEBUG /wd4267 /wd4819)
else()
    # GCC/Clang 优化选项（保持原有行为）
    target_compile_options(qv100 PRIVATE
        -O3
        -march=native
        -mtune=native
        -ffast-math
        -funroll-loops
        -finline-functions
        -fomit-frame-pointer
        -DNDEBUG
        -DEIGEN_NO_DEBUG
        -fopenmp
    )
    target_link_options(qv100 PRIVATE -fopenmp)

    target_compile_options(full_amplitude_accelerate PRIVATE
        -O3
        -march=native
        -mtune=native
        -ffast-math
        -funroll-loops
        -finline-functions
        -fomit-frame-pointer
        -DNDEBUG
        -fopenmp
    )

    target_compile_options(gradient_simulator PRIVATE
        -O3
        -march=native
        -mtune=native
        -ffast-math
        -funroll-loops
        -finline-functions
        -fomit-frame-pointer
        -DNDEBUG
        -fopenmp
    )
endif()

# ============================================================
# 设置输出模块名（移除 lib 前缀；Windows 上 pybind11_add_module 自动生成 .pyd）
# ============================================================
set_target_properties(types PROPERTIES PREFIX "")
set_target_properties(qv100 PROPERTIES PREFIX "")
set_target_properties(n_local_accelerate PROPERTIES PREFIX "")
set_target_properties(full_amplitude_accelerate PROPERTIES PREFIX "")
set_target_properties(gradient_simulator PROPERTIES PREFIX "")

# Windows 上确保输出为 .pyd（pybind11_add_module 默认就是 .pyd，这里显式确认）
if(WIN32)
    foreach(mod types qv100 n_local_accelerate full_amplitude_accelerate gradient_simulator)
        set_target_properties(${mod} PROPERTIES SUFFIX ".pyd")
    endforeach()
endif()
