cmake_minimum_required(VERSION 3.24)

project(separable_compilation_test LANGUAGES CXX CUDA)

# Build device code for the GPU present on the build machine.
set(CMAKE_CUDA_ARCHITECTURES native)

set(NUM_KERNELS 1 CACHE STRING "Total number of kernel translation units")
if(NOT "${NUM_KERNELS}" MATCHES "^[1-9][0-9]*$")
    message(FATAL_ERROR "NUM_KERNELS must be a positive integer")
endif()

find_package(CUDAToolkit REQUIRED)

function(configure_kernel_library target kernel_namespace)
    target_include_directories(${target} PUBLIC include)
    target_link_libraries(${target} PUBLIC CUDA::cudart)
    target_compile_features(${target} PUBLIC cxx_std_17 cuda_std_17)
    target_compile_definitions(${target} PRIVATE
        KERNEL_NAMESPACE=${kernel_namespace}
    )
    target_compile_options(${target} PRIVATE
        $<$<COMPILE_LANGUAGE:CUDA>:-lineinfo>
        $<$<AND:$<COMPILE_LANGUAGE:CUDA>,$<CONFIG:Release>>:-O3>
    )
endfunction()

set(kernel_sources src/kernel.cu)

if(NUM_KERNELS GREATER 1)
    set(generated_kernel_dir "${CMAKE_CURRENT_BINARY_DIR}/generated")
    file(MAKE_DIRECTORY "${generated_kernel_dir}")
    math(EXPR last_extra_kernel "${NUM_KERNELS} - 1")

    foreach(EXTRA RANGE 1 ${last_extra_kernel})
        set(generated_kernel
            "${generated_kernel_dir}/extra_kernel_${EXTRA}.cu"
        )
        configure_file(
            src/extra_kernel.cu
            "${generated_kernel}"
            @ONLY
        )
        list(APPEND kernel_sources "${generated_kernel}")
    endforeach()
endif()

# All vec4 definitions are visible in every kernel translation unit. This
# target is intentionally compiled without relocatable device code or a device
# link.
add_library(kernel_direct STATIC ${kernel_sources})
configure_kernel_library(kernel_direct direct)
target_compile_definitions(kernel_direct PRIVATE
    INCLUDE_VEC4_IMPL_HEADER
)
set_target_properties(kernel_direct PROPERTIES
    CUDA_SEPARABLE_COMPILATION OFF
    CUDA_RESOLVE_DEVICE_SYMBOLS OFF
)

# The kernel sees only vec4 declarations. Its device calls are resolved against
# vec4_impl.cu by this library's device-link step.
add_library(kernel_separable STATIC
    ${kernel_sources}
    src/vec4_impl.cu
)
configure_kernel_library(kernel_separable separable_compilation)
set_target_properties(kernel_separable PROPERTIES
    CUDA_SEPARABLE_COMPILATION ON
    CUDA_RESOLVE_DEVICE_SYMBOLS ON
)

# This target has the same translation-unit boundary as kernel_separable, but
# preserves LTO IR during compilation and optimizes across it at device link.
add_library(kernel_dlto STATIC
    ${kernel_sources}
    src/vec4_impl.cu
)
configure_kernel_library(kernel_dlto dlto)
set_target_properties(kernel_dlto PROPERTIES
    CUDA_SEPARABLE_COMPILATION ON
    CUDA_RESOLVE_DEVICE_SYMBOLS ON
)
target_compile_options(kernel_dlto PRIVATE
    $<$<COMPILE_LANGUAGE:CUDA>:-dlto>
)
target_link_options(kernel_dlto PRIVATE
    $<DEVICE_LINK:-dlto>
)

add_executable(separable_compilation_test src/main.cpp)
target_compile_features(separable_compilation_test PRIVATE cxx_std_17)
target_link_libraries(separable_compilation_test PRIVATE
    kernel_direct
    kernel_separable
    kernel_dlto
)
